From b7d40aeb037e7e0f338603ebdddd58333c479550 Mon Sep 17 00:00:00 2001 From: Erwan Leboucher Date: Tue, 14 Jul 2026 18:29:14 +0200 Subject: [PATCH 1/2] chore: update matrix-js-sdk with the new oauth api --- package.json | 2 +- pnpm-lock.yaml | 26 +--- pnpm-workspace.yaml | 1 + src/app/components/ServerConfigsLoader.tsx | 13 +- src/app/hooks/useAuthFlows.ts | 4 +- src/app/pages/auth/login/OidcLogin.tsx | 4 +- .../pages/auth/login/oidcLoginUtil.test.ts | 27 ++-- src/app/pages/auth/login/oidcLoginUtil.ts | 115 ++++++++++++------ .../pages/client/BackgroundNotifications.tsx | 17 ++- src/app/state/sessions.ts | 3 +- src/client/initMatrix.ts | 53 ++++---- src/client/oidcTokenRefresher.ts | 63 +++++----- src/types/matrix-sdk.ts | 9 +- vitest.config.ts | 7 ++ 14 files changed, 172 insertions(+), 172 deletions(-) diff --git a/package.json b/package.json index df86181fed..50392d3ed2 100644 --- a/package.json +++ b/package.json @@ -90,7 +90,7 @@ "linkify-react": "^4.3.3", "linkifyjs": "^4.3.3", "marked": "^18.0.5", - "matrix-js-sdk": "^41.9.0", + "matrix-js-sdk": "42.0.0", "matrix-widget-api": "^1.17.0", "pdfjs-dist": "^6.1.200", "react": "^18.3.1", diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 1d4e55ab0b..da7ba51340 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -186,8 +186,8 @@ importers: specifier: ^18.0.5 version: 18.0.5 matrix-js-sdk: - specifier: ^41.9.0 - version: 41.9.0 + specifier: 42.0.0 + version: 42.0.0 matrix-widget-api: specifier: ^1.17.0 version: 1.17.0 @@ -4102,10 +4102,6 @@ packages: resolution: {integrity: sha512-p/nXbhSEcu3pZRdkW1OfJhpsVtW1gd4Wa1fnQc9YLiTfAjn0312eMKimbdIQzuZl9aa9xUGaRlP9T/CJE/ditQ==} engines: {node: '>=0.10.0'} - jwt-decode@4.0.0: - resolution: {integrity: sha512-+KJGIyHgkGuIq3IEBNftfhW/LfWhXUIY6OmyVWjliu5KH1y0fw7VQ8YndE2O4qZdMSd9SqbnC8GOcZEy0Om7sA==} - engines: {node: '>=18'} - katex@0.17.0: resolution: {integrity: sha512-Vdw0ATsQ9V+LuegM/BTwQqV/6cTl5lbGcIrU+BCgLxyf6bo38ybOr372tuSIxir3CN720flu1meYR6XzNMwQnw==} hasBin: true @@ -4275,8 +4271,8 @@ packages: matrix-events-sdk@0.0.1: resolution: {integrity: sha512-1QEOsXO+bhyCroIe2/A5OwaxHvBm7EsSQ46DEDn8RBIfQwN5HWBpFvyWWR4QY0KHPPnnJdI99wgRiAl7Ad5qaA==} - matrix-js-sdk@41.9.0: - resolution: {integrity: sha512-xRoaIxu8e7ECV8ctT+gxXvOA+KqPJdellihIeM7rZ5LPl9ZWMct+ltN/Z4f76PJtAazjIXd8A3edPu6OSKk9/g==} + matrix-js-sdk@42.0.0: + resolution: {integrity: sha512-i3SnhdkgnKA7O0iVlaZCo96GGaIZ0FGD93LQ34N2+Hj/kaC7Ffbch7OEkwYA//mYGCKRAMxduBacoy8WscfqVw==} engines: {node: '>=22.0.0'} matrix-widget-api@1.17.0: @@ -4373,10 +4369,6 @@ packages: resolution: {integrity: sha512-9miFgM2OFba7hB+pRgvtV84pYTBaoTHohvmIgiRt6dRIzbwEOIaNaP+dIlGs2fNFoB0SeISs0Jz5WFVRid6Xyg==} engines: {node: '>=12.20.0'} - oidc-client-ts@3.5.0: - resolution: {integrity: sha512-l2q8l9CTCTOlbX+AnK4p3M+4CEpKpyQhle6blQkdFhm0IsBqsxm15bYaSa11G7pWdsYr6epdsRZxJpCyCRbT8A==} - engines: {node: '>=18'} - own-keys@1.0.1: resolution: {integrity: sha512-qFOyK5PjiWZd+QQIh+1jhdb9LpxTF0qs7Pm8o5QHYZ0M3vKqSqzsZaEB6oWlxZ+q2sJBMI/Ktgd2N5ZwQoRHfg==} engines: {node: '>= 0.4'} @@ -8858,8 +8850,6 @@ snapshots: jsonpointer@5.0.1: {} - jwt-decode@4.0.0: {} - katex@0.17.0: dependencies: commander: 8.3.0 @@ -8998,18 +8988,16 @@ snapshots: matrix-events-sdk@0.0.1: {} - matrix-js-sdk@41.9.0: + matrix-js-sdk@42.0.0: dependencies: '@babel/runtime': 8.0.0 '@matrix-org/matrix-sdk-crypto-wasm': 18.3.1 another-json: 0.2.0 bs58: 6.0.0 content-type: 2.0.0 - jwt-decode: 4.0.0 loglevel: 1.9.2 matrix-events-sdk: 0.0.1 matrix-widget-api: 1.17.0 - oidc-client-ts: 3.5.0 p-retry: 8.0.0 sdp-transform: 3.0.0 unhomoglyph: 1.0.6 @@ -9100,10 +9088,6 @@ snapshots: obug@2.1.3: {} - oidc-client-ts@3.5.0: - dependencies: - jwt-decode: 4.0.0 - own-keys@1.0.1: dependencies: get-intrinsic: 1.3.0 diff --git a/pnpm-workspace.yaml b/pnpm-workspace.yaml index b616adb35d..dd46996060 100644 --- a/pnpm-workspace.yaml +++ b/pnpm-workspace.yaml @@ -10,6 +10,7 @@ allowBuilds: engineStrict: true minimumReleaseAge: 1440 minimumReleaseAgeExclude: + - 'matrix-js-sdk' - '@sableclient/sable-call-embedded' - '@sableclient/twemoji-font' diff --git a/src/app/components/ServerConfigsLoader.tsx b/src/app/components/ServerConfigsLoader.tsx index 833fdbe3c0..68776a7b66 100644 --- a/src/app/components/ServerConfigsLoader.tsx +++ b/src/app/components/ServerConfigsLoader.tsx @@ -1,12 +1,10 @@ import type { ReactNode } from 'react'; import { useCallback, useMemo } from 'react'; import type { Capabilities, ValidatedAuthMetadata } from '$types/matrix-sdk'; -import { validateAuthMetadata } from '$types/matrix-sdk'; import { AsyncStatus, useAsyncCallbackValue } from '$hooks/useAsyncCallback'; import { useMatrixClient } from '$hooks/useMatrixClient'; import type { MediaConfig } from '$hooks/useMediaConfig'; import { promiseFulfilledResult } from '$utils/common'; -import { createLogger } from '$utils/debug'; export type ServerConfigs = { capabilities?: Capabilities; @@ -18,8 +16,6 @@ type ServerConfigsLoaderProps = { children: (configs: ServerConfigs) => ReactNode; }; -const log = createLogger('ServerConfigsLoader'); - export function ServerConfigsLoader({ children }: ServerConfigsLoaderProps) { const mx = useMatrixClient(); const fallbackConfigs = useMemo(() => ({}), []); @@ -35,18 +31,11 @@ export function ServerConfigsLoader({ children }: ServerConfigsLoaderProps) { const capabilities = promiseFulfilledResult(result[0]); const mediaConfig = promiseFulfilledResult(result[1]); const authMetadata = promiseFulfilledResult(result[2]); - let validatedAuthMetadata: ValidatedAuthMetadata | undefined; - - try { - validatedAuthMetadata = validateAuthMetadata(authMetadata); - } catch (e) { - log.error('Failed to validate auth metadata:', e); - } return { capabilities, mediaConfig, - authMetadata: validatedAuthMetadata, + authMetadata, }; }, [mx]) ); diff --git a/src/app/hooks/useAuthFlows.ts b/src/app/hooks/useAuthFlows.ts index c046d86122..c7f9d82c66 100644 --- a/src/app/hooks/useAuthFlows.ts +++ b/src/app/hooks/useAuthFlows.ts @@ -3,7 +3,7 @@ import type { IAuthData, MatrixError, ILoginFlowsResponse, - OidcClientConfig, + ValidatedAuthMetadata, } from '$types/matrix-sdk'; export enum RegisterFlowStatus { @@ -48,7 +48,7 @@ export const parseRegisterErrResp = (matrixError: MatrixError): RegisterFlowsRes export type AuthFlows = { loginFlows: ILoginFlowsResponse; registerFlows: RegisterFlowsResponse; - authMetadata?: OidcClientConfig; + authMetadata?: ValidatedAuthMetadata; }; const AuthFlowsContext = createContext(null); diff --git a/src/app/pages/auth/login/OidcLogin.tsx b/src/app/pages/auth/login/OidcLogin.tsx index b64413a757..d6f71b34c2 100644 --- a/src/app/pages/auth/login/OidcLogin.tsx +++ b/src/app/pages/auth/login/OidcLogin.tsx @@ -1,7 +1,7 @@ import { Box, Button, Overlay, OverlayBackdrop, OverlayCenter, Spinner, Text } from 'folds'; import { useCallback, useEffect } from 'react'; import { useNavigate } from 'react-router-dom'; -import type { OidcClientConfig } from '$types/matrix-sdk'; +import type { ValidatedAuthMetadata } from '$types/matrix-sdk'; import { AsyncStatus, useAsyncCallback } from '$hooks/useAsyncCallback'; import { useAuthServer } from '$hooks/useAuthServer'; import { InfoCard } from '$components/info-card'; @@ -32,7 +32,7 @@ const oidcErrorMessage = (error: unknown): string => { }; type OidcLoginButtonProps = { - authMetadata: OidcClientConfig; + authMetadata: ValidatedAuthMetadata; homeserverUrl: string; redirectUri: string; label: string; diff --git a/src/app/pages/auth/login/oidcLoginUtil.test.ts b/src/app/pages/auth/login/oidcLoginUtil.test.ts index 3729a5fa64..aa5d0001ee 100644 --- a/src/app/pages/auth/login/oidcLoginUtil.test.ts +++ b/src/app/pages/auth/login/oidcLoginUtil.test.ts @@ -1,4 +1,4 @@ -import { describe, it, expect, vi, afterEach } from 'vitest'; +import { describe, it, expect } from 'vitest'; import type { BearerTokenResponse } from '$types/matrix-sdk'; import { deviceIdFromScope, expiresInMsFromToken } from './oidcLoginUtil'; @@ -20,23 +20,18 @@ describe('deviceIdFromScope', () => { }); describe('expiresInMsFromToken', () => { - afterEach(() => { - vi.useRealTimers(); - }); - - it('prefers expires_in (seconds → ms)', () => { - const token = { expires_in: 299 } as BearerTokenResponse; + it('converts expires_in (seconds → ms)', () => { + const token = { + access_token: 'x', + token_type: 'Bearer', + expires_in: 299, + } as BearerTokenResponse; expect(expiresInMsFromToken(token)).toBe(299_000); }); - it('derives ms from expires_at relative to now', () => { - vi.useFakeTimers(); - vi.setSystemTime(new Date(1_000_000_000_000)); - const token = { expires_at: 1_000_000_000 + 300 } as BearerTokenResponse; - expect(expiresInMsFromToken(token)).toBe(300_000); - }); - - it('returns undefined when neither field is present', () => { - expect(expiresInMsFromToken({} as BearerTokenResponse)).toBeUndefined(); + it('returns undefined when expires_in is absent', () => { + expect( + expiresInMsFromToken({ access_token: 'x', token_type: 'Bearer' } as BearerTokenResponse) + ).toBeUndefined(); }); }); diff --git a/src/app/pages/auth/login/oidcLoginUtil.ts b/src/app/pages/auth/login/oidcLoginUtil.ts index c16f63eb99..8b160aec96 100644 --- a/src/app/pages/auth/login/oidcLoginUtil.ts +++ b/src/app/pages/auth/login/oidcLoginUtil.ts @@ -1,11 +1,10 @@ import to from 'await-to-js'; -import type { OidcClientConfig, BearerTokenResponse } from '$types/matrix-sdk'; -import { - createClient, - registerOidcClient, - generateOidcAuthorizationUrl, - completeAuthorizationCodeGrant, +import type { + ValidatedAuthMetadata, + BearerTokenResponse, + OAuthRegistrationRequest, } from '$types/matrix-sdk'; +import { OAuth2, createClient } from '$types/matrix-sdk'; import { isTauri } from '@tauri-apps/api/core'; import { openUrl } from '@tauri-apps/plugin-opener'; import { createLogger } from '$utils/debug'; @@ -26,6 +25,7 @@ export enum OidcLoginError { UserCancelled = 'UserCancelled', CodeExchangeFailed = 'CodeExchangeFailed', MissingDeviceId = 'MissingDeviceId', + MissingOauthContext = 'MissingOauthContext', Unknown = 'Unknown', } @@ -36,8 +36,38 @@ export class OidcLoginFailure extends Error { } } +const OAUTH_CONTEXT_KEY = 'oauth_login_context'; + +type OauthLoginContext = { + metadata: ValidatedAuthMetadata; + clientId: string; + redirectUri: string; + deviceId?: string; + codeVerifier?: string; + homeserverUrl: string; +}; + +const persistOauthContext = (state: string, ctx: OauthLoginContext): void => { + sessionStorage.setItem(OAUTH_CONTEXT_KEY, JSON.stringify({ state, ...ctx })); +}; + +const consumeOauthContext = (state: string): OauthLoginContext | undefined => { + const raw = sessionStorage.getItem(OAUTH_CONTEXT_KEY); + if (!raw) return undefined; + try { + const stored = JSON.parse(raw) as OauthLoginContext & { state: string }; + if (stored.state !== state) return undefined; + sessionStorage.removeItem(OAUTH_CONTEXT_KEY); + const { state: _state, ...ctx } = stored; + return ctx; + } catch { + sessionStorage.removeItem(OAUTH_CONTEXT_KEY); + return undefined; + } +}; + export const startOidcLogin = async ( - authMetadata: OidcClientConfig, + authMetadata: ValidatedAuthMetadata, homeserverUrl: string, redirectUri: string, opts?: { prompt?: string; server?: string } @@ -45,31 +75,33 @@ export const startOidcLogin = async ( const tauri = isTauri(); const effectiveRedirectUri = tauri ? buildTauriOidcRedirectUrl() : redirectUri; - const [registerErr, clientId] = await to( - registerOidcClient(authMetadata, { - clientName: CLIENT_NAME, - clientUri: tauri ? TAURI_OIDC_CLIENT_URI : window.location.origin, - applicationType: tauri ? 'native' : 'web', - redirectUris: [effectiveRedirectUri], - contacts: undefined, - tosUri: undefined, - policyUri: undefined, - }) - ); + const clientMetadata: OAuthRegistrationRequest = { + client_name: CLIENT_NAME, + client_uri: tauri ? TAURI_OIDC_CLIENT_URI : window.location.origin, + application_type: tauri ? 'native' : 'web', + redirect_uris: [effectiveRedirectUri], + }; + + const [registerErr, clientId] = await to(OAuth2.registerClient(authMetadata, clientMetadata)); if (registerErr || !clientId) { - log.error('OIDC client registration failed', registerErr); + log.error('OAuth2 client registration failed', registerErr); throw new OidcLoginFailure(OidcLoginError.RegistrationFailed); } - const authUrl = await generateOidcAuthorizationUrl({ + const state = crypto.randomUUID(); + const oauth2 = new OAuth2(authMetadata, { clientId, redirectUri: effectiveRedirectUri }); + + persistOauthContext(state, { metadata: authMetadata, clientId, - homeserverUrl, redirectUri: effectiveRedirectUri, - nonce: crypto.randomUUID(), - prompt: opts?.prompt, + deviceId: oauth2.context.deviceId, + codeVerifier: oauth2.context.codeVerifier, + homeserverUrl, }); + const authUrl = await oauth2.generateAuthorizationCodeGrantUrl(state, 'query', opts?.prompt); + if (tauri) { rememberTauriOidcServer(opts?.server); await openUrl(authUrl); @@ -86,7 +118,6 @@ export const deviceIdFromScope = (scope: string): string | undefined => export const expiresInMsFromToken = (token: BearerTokenResponse): number | undefined => { if (typeof token.expires_in === 'number') return token.expires_in * 1000; - if (typeof token.expires_at === 'number') return token.expires_at * 1000 - Date.now(); return undefined; }; @@ -96,44 +127,54 @@ export type OidcLoginResult = { }; export const completeOidcLogin = async (code: string, state: string): Promise => { - const [grantErr, grant] = await to(completeAuthorizationCodeGrant(code, state)); - if (grantErr || !grant) { - log.error('OIDC code exchange failed', grantErr); - throw new OidcLoginFailure(OidcLoginError.CodeExchangeFailed); + const ctx = consumeOauthContext(state); + if (!ctx) { + log.error('No OAuth2 context found for state', state); + throw new OidcLoginFailure(OidcLoginError.MissingOauthContext); } - const { tokenResponse, homeserverUrl, oidcClientSettings, idTokenClaims } = grant; + const oauth2 = new OAuth2(ctx.metadata, { + clientId: ctx.clientId, + redirectUri: ctx.redirectUri, + deviceId: ctx.deviceId, + codeVerifier: ctx.codeVerifier, + }); + + const [grantErr, tokenResponse] = await to(oauth2.completeAuthorizationCodeGrant(code)); + if (grantErr || !tokenResponse) { + log.error('OAuth2 code exchange failed', grantErr); + throw new OidcLoginFailure(OidcLoginError.CodeExchangeFailed); + } const mx = createClient({ - baseUrl: homeserverUrl, + baseUrl: ctx.homeserverUrl, fetchFn: fetch, accessToken: tokenResponse.access_token, }); const [whoamiErr, whoami] = await to(mx.whoami()); if (whoamiErr || !whoami?.user_id) { - log.error('OIDC whoami failed', whoamiErr); + log.error('OAuth2 whoami failed', whoamiErr); throw new OidcLoginFailure(OidcLoginError.Unknown); } // Prefer the server-authoritative device id; fall back to the one carried in the granted scope. - const deviceId = whoami.device_id ?? deviceIdFromScope(tokenResponse.scope); + const deviceId = whoami.device_id ?? deviceIdFromScope(tokenResponse.scope ?? ''); if (!deviceId) { throw new OidcLoginFailure(OidcLoginError.MissingDeviceId); } const session: Session = { - baseUrl: homeserverUrl, + baseUrl: ctx.homeserverUrl, userId: whoami.user_id, deviceId, accessToken: tokenResponse.access_token, refreshToken: tokenResponse.refresh_token, expiresInMs: expiresInMsFromToken(tokenResponse), oidc: { - issuer: oidcClientSettings.issuer, - clientId: oidcClientSettings.clientId, - idTokenClaims, + issuer: ctx.metadata.issuer, + clientId: ctx.clientId, }, }; - return { baseUrl: homeserverUrl, session }; + return { baseUrl: ctx.homeserverUrl, session }; }; diff --git a/src/app/pages/client/BackgroundNotifications.tsx b/src/app/pages/client/BackgroundNotifications.tsx index d9b7d93463..36b662acee 100644 --- a/src/app/pages/client/BackgroundNotifications.tsx +++ b/src/app/pages/client/BackgroundNotifications.tsx @@ -44,7 +44,7 @@ import { } from '$utils/notificationStyle'; import * as Sentry from '@sentry/react'; import { startClient, stopClient } from '$client/initMatrix'; -import { SessionOidcTokenRefresher } from '$client/oidcTokenRefresher'; +import { createSessionTokenRefresher } from '$client/oidcTokenRefresher'; import { mobileOrTablet } from '$utils/user-agent'; import { isTauri } from '@tauri-apps/api/core'; import { type as osType } from '@tauri-apps/plugin-os'; @@ -83,8 +83,15 @@ const startBackgroundClient = async (session: Session): Promise => dbName: storeName.sync, }); - const tokenRefresher = - session.oidc && session.refreshToken ? new SessionOidcTokenRefresher(session) : undefined; + const tempClient = createClient({ + baseUrl: session.baseUrl, + fetchFn: fetch, + accessToken: session.accessToken, + refreshToken: session.refreshToken, + userId: session.userId, + deviceId: session.deviceId, + }); + const tokenRefresher = await createSessionTokenRefresher(session, tempClient); const mx = createClient({ baseUrl: session.baseUrl, @@ -95,9 +102,7 @@ const startBackgroundClient = async (session: Session): Promise => deviceId: session.deviceId, store: indexedDBStore, timelineSupport: false, - tokenRefreshFunction: tokenRefresher - ? (refreshToken) => tokenRefresher.doRefreshAccessToken(refreshToken) - : undefined, + tokenRefreshFunction: tokenRefresher?.tokenRefreshFunction, }); const startOpts = { diff --git a/src/app/state/sessions.ts b/src/app/state/sessions.ts index 2637616dfa..4bedb03ef0 100644 --- a/src/app/state/sessions.ts +++ b/src/app/state/sessions.ts @@ -1,5 +1,5 @@ import { atom } from 'jotai'; -import type { IdTokenClaims, MatrixEvent, Room } from '$types/matrix-sdk'; +import type { MatrixEvent, Room } from '$types/matrix-sdk'; import { createLogger } from '$utils/debug'; import { atomWithLocalStorage, @@ -17,7 +17,6 @@ const notifySessionChanged = (): void => { export type OidcSessionInfo = { issuer: string; clientId: string; - idTokenClaims?: IdTokenClaims; }; export type Session = { diff --git a/src/client/initMatrix.ts b/src/client/initMatrix.ts index 4f4dd79c2f..e52c156e9d 100644 --- a/src/client/initMatrix.ts +++ b/src/client/initMatrix.ts @@ -10,10 +10,12 @@ import { IndexedDBStore, IndexedDBCryptoStore, KnownMembership, + OAuth2, SyncState, } from '$types/matrix-sdk'; import { fetch } from '$utils/fetch'; import { clearMediaCache } from '$utils/mediaCache'; +import { getAppOrigin } from '$utils/platform'; import { clearNavToActivePathStore } from '$state/navToActivePath'; import type { Session, SessionStoreName } from '$state/sessions'; @@ -22,7 +24,7 @@ import { createLogger } from '$utils/debug'; import { createDebugLogger } from '$utils/debugLogger'; import * as Sentry from '@sentry/react'; import { pushSessionToSW } from '../sw-session'; -import { SessionOidcTokenRefresher } from './oidcTokenRefresher'; +import { createSessionTokenRefresher } from './oidcTokenRefresher'; import { cryptoCallbacks } from './secretStorageKeys'; import type { SlidingSyncDiagnostics } from './slidingSync'; import { scopeEphemeralExtensions, SlidingSyncManager } from './slidingSync'; @@ -231,7 +233,7 @@ type BuiltClient = { indexedDBStore: IndexedDBStore; }; -const buildClient = (session: Session): BuiltClient => { +const buildClient = async (session: Session): Promise => { const storeName = getSessionStoreName(session); const indexedDBStore = new IndexedDBStore({ @@ -242,8 +244,15 @@ const buildClient = (session: Session): BuiltClient => { const legacyCryptoStore = new IndexedDBCryptoStore(global.indexedDB, storeName.crypto); - const tokenRefresher = - session.oidc && session.refreshToken ? new SessionOidcTokenRefresher(session) : undefined; + const tempClient = createClient({ + baseUrl: session.baseUrl, + fetchFn: fetch, + accessToken: session.accessToken, + refreshToken: session.refreshToken, + userId: session.userId, + deviceId: session.deviceId, + }); + const tokenRefresher = await createSessionTokenRefresher(session, tempClient); const mx = createClient({ baseUrl: session.baseUrl, @@ -257,9 +266,7 @@ const buildClient = (session: Session): BuiltClient => { timelineSupport: true, cryptoCallbacks: cryptoCallbacks as unknown as CryptoCallbacks, verificationMethods: ['m.sas.v1'], - tokenRefreshFunction: tokenRefresher - ? (refreshToken) => tokenRefresher.doRefreshAccessToken(refreshToken) - : undefined, + tokenRefreshFunction: tokenRefresher?.tokenRefreshFunction, }); return { mx, indexedDBStore }; @@ -275,7 +282,7 @@ const initializeClient = async ( ): Promise => { let builtClient: BuiltClient; try { - builtClient = buildClient(session); + builtClient = await buildClient(session); } catch (error) { return { ok: false, error, phase: 'sync_store' }; } @@ -545,37 +552,23 @@ export const getClientSyncDiagnostics = (mx: MatrixClient): ClientSyncDiagnostic }; }; -const revokeOidcToken = async ( - endpoint: string, - token: string, - tokenTypeHint: 'access_token' | 'refresh_token', - clientId: string -): Promise => { - const res = await fetch(endpoint, { - method: 'POST', - headers: { 'Content-Type': 'application/x-www-form-urlencoded' }, - body: new URLSearchParams({ token, token_type_hint: tokenTypeHint, client_id: clientId }), - }); - if (!res.ok) { - throw new Error(`OIDC token revocation failed (${res.status})`); - } -}; - const revokeOidcSession = async (mx: MatrixClient, session: Session): Promise => { const clientId = session.oidc?.clientId; if (!clientId) return; const metadata = await mx.getAuthMetadata(); - const endpoint = metadata.revocation_endpoint; - if (!endpoint) return; + + const oauth2 = new OAuth2(metadata, { + clientId, + redirectUri: getAppOrigin(), + deviceId: session.deviceId, + }); const accessToken = mx.getAccessToken() ?? undefined; const results = await Promise.allSettled([ session.refreshToken - ? revokeOidcToken(endpoint, session.refreshToken, 'refresh_token', clientId) - : Promise.resolve(), - accessToken - ? revokeOidcToken(endpoint, accessToken, 'access_token', clientId) + ? oauth2.revokeToken(session.refreshToken, 'refresh_token') : Promise.resolve(), + accessToken ? oauth2.revokeToken(accessToken, 'access_token') : Promise.resolve(), ]); if (results.some((r) => r.status === 'rejected')) { debugLog.warn('general', 'OIDC token revocation had failures', { userId: session.userId }); diff --git a/src/client/oidcTokenRefresher.ts b/src/client/oidcTokenRefresher.ts index 4e8295f2fe..d8de063160 100644 --- a/src/client/oidcTokenRefresher.ts +++ b/src/client/oidcTokenRefresher.ts @@ -1,5 +1,6 @@ -import type { AccessTokens } from '$types/matrix-sdk'; -import { OidcTokenRefresher } from '$types/matrix-sdk'; +import type { AccessTokens, ValidatedAuthMetadata } from '$types/matrix-sdk'; +import { OAuth2, TokenRefresher } from '$types/matrix-sdk'; +import type { MatrixClient } from '$types/matrix-sdk'; import type { Session } from '$state/sessions'; import { ACTIVE_SESSION_KEY, @@ -10,45 +11,37 @@ import { getLocalStorageItem } from '$state/utils/atomWithLocalStorage'; import { pushSessionToSW } from '../sw-session'; import { getAppOrigin } from '$utils/platform'; -export class SessionOidcTokenRefresher extends OidcTokenRefresher { - private readonly userId: string; +export const createSessionTokenRefresher = async ( + session: Session, + mx: MatrixClient +): Promise => { + if (!session.oidc || !session.refreshToken) return undefined; - private readonly baseUrl: string; + const metadata: ValidatedAuthMetadata = await mx.getAuthMetadata(); - public constructor(session: Session) { - if (!session.oidc) { - throw new Error('SessionOidcTokenRefresher requires an OIDC session'); - } - super( - session.oidc.issuer, - session.oidc.clientId, - getAppOrigin(), - session.deviceId, - session.oidc.idTokenClaims ?? ({} as never) - ); - this.userId = session.userId; - this.baseUrl = session.baseUrl; - } - - // Another tab may have rotated the token; reusing a consumed one revokes the session. - public override doRefreshAccessToken(refreshToken: string): Promise { - return super.doRefreshAccessToken(getStoredSessionRefreshToken(this.userId) ?? refreshToken); - } + const oauth2 = new OAuth2(metadata, { + clientId: session.oidc.clientId, + redirectUri: getAppOrigin(), + deviceId: session.deviceId, + }); - protected async persistTokens(tokens: { - accessToken: string; - refreshToken?: string; - }): Promise { - const { expiry } = tokens as { expiry?: Date }; - updateSessionTokens(this.userId, { + const onRefresh = async (tokens: AccessTokens): Promise => { + updateSessionTokens(session.userId, { accessToken: tokens.accessToken, refreshToken: tokens.refreshToken, - expiresInMs: expiry ? expiry.getTime() - Date.now() : undefined, + expiresInMs: tokens.expiry ? tokens.expiry.getTime() - Date.now() : undefined, }); // Only the active session owns the single service-worker session. const activeSessionId = getLocalStorageItem(ACTIVE_SESSION_KEY, undefined); - if (activeSessionId === this.userId) { - pushSessionToSW(this.baseUrl, tokens.accessToken, this.userId); + if (activeSessionId === session.userId) { + pushSessionToSW(session.baseUrl, tokens.accessToken, session.userId); } - } -} + }; + + const tokenRefresher = new TokenRefresher(oauth2, onRefresh); + const refresh = tokenRefresher.tokenRefreshFunction; + // Another tab may have rotated the token; reusing a consumed one revokes the session. + tokenRefresher.tokenRefreshFunction = (refreshToken) => + refresh(getStoredSessionRefreshToken(session.userId) ?? refreshToken); + return tokenRefresher; +}; diff --git a/src/types/matrix-sdk.ts b/src/types/matrix-sdk.ts index 561941238b..f1bea4e091 100644 --- a/src/types/matrix-sdk.ts +++ b/src/types/matrix-sdk.ts @@ -47,14 +47,7 @@ export * from 'matrix-js-sdk/lib/@types/registration'; export * from 'matrix-js-sdk/lib/feature'; -export * from 'matrix-js-sdk/lib/oidc'; - -// oidc-client-ts is a transitive dep (not resolvable here), so derive its type from the SDK. -import type { completeAuthorizationCodeGrant } from 'matrix-js-sdk/lib/oidc'; - -export type IdTokenClaims = Awaited< - ReturnType ->['idTokenClaims']; +export * from 'matrix-js-sdk/lib/oauth'; export { VerificationMethod } from 'matrix-js-sdk/lib/types'; export * from 'matrix-js-sdk/lib/pushprocessor'; export * from 'matrix-js-sdk/lib/common-crypto/CryptoBackend'; diff --git a/vitest.config.ts b/vitest.config.ts index 063634e42b..2149b5b109 100644 --- a/vitest.config.ts +++ b/vitest.config.ts @@ -9,6 +9,8 @@ export default defineConfig({ plugins: [react(), vanillaExtractPlugin()], resolve: { alias: { + 'matrix-js-sdk/lib': path.resolve(__dirname, 'node_modules/matrix-js-sdk/lib'), + 'matrix-js-sdk': path.resolve(__dirname, 'node_modules/matrix-js-sdk/lib/matrix.js'), $hooks: path.resolve(__dirname, 'src/app/hooks'), $plugins: path.resolve(__dirname, 'src/app/plugins'), $components: path.resolve(__dirname, 'src/app/components'), @@ -38,6 +40,11 @@ export default defineConfig({ globals: true, setupFiles: ['./src/test/setup.ts'], include: ['src/**/*.{test,spec}.{ts,tsx}'], + server: { + deps: { + inline: [/matrix-js-sdk\/lib\//], + }, + }, coverage: { provider: 'v8', reporter: ['text', 'html', 'lcov'], From 59b5e87cf5c67130249165d73ed339b34c266a28 Mon Sep 17 00:00:00 2001 From: 7w1 Date: Tue, 21 Jul 2026 18:23:53 -0500 Subject: [PATCH 2/2] cleanup --- .changeset/update-oauth-implementation.md | 5 ++ .../settings/devices/OtherDevices.tsx | 30 +++---- .../settings/devices/Verification.tsx | 10 +-- src/app/hooks/useAccountManagement.ts | 16 ++++ src/app/pages/App.tsx | 2 + src/app/pages/Router.tsx | 1 + src/app/pages/TauriDeepLinkBridge.tsx | 20 +++-- src/app/pages/auth/SSOTauri.ts | 39 ++++----- src/app/pages/auth/login/Login.tsx | 6 +- src/app/pages/auth/login/OidcLogin.tsx | 2 + src/app/pages/auth/login/oidcLoginUtil.ts | 83 ++++++++++++------- src/app/pages/auth/register/Register.tsx | 2 +- .../pages/client/BackgroundNotifications.tsx | 2 +- src/app/utils/oauthCallback.ts | 34 ++++++++ src/client/initMatrix.ts | 38 ++++----- src/client/oauthTokenRevocation.ts | 29 +++++++ src/client/oidcTokenRefresher.ts | 69 +++++++++++---- 17 files changed, 264 insertions(+), 124 deletions(-) create mode 100644 .changeset/update-oauth-implementation.md create mode 100644 src/app/utils/oauthCallback.ts create mode 100644 src/client/oauthTokenRevocation.ts diff --git a/.changeset/update-oauth-implementation.md b/.changeset/update-oauth-implementation.md new file mode 100644 index 0000000000..84818be386 --- /dev/null +++ b/.changeset/update-oauth-implementation.md @@ -0,0 +1,5 @@ +--- +default: patch +--- + +Updated OIDC implementation with new matrix-js-sdk version. diff --git a/src/app/features/settings/devices/OtherDevices.tsx b/src/app/features/settings/devices/OtherDevices.tsx index 3b248157d0..117bc77a8c 100644 --- a/src/app/features/settings/devices/OtherDevices.tsx +++ b/src/app/features/settings/devices/OtherDevices.tsx @@ -11,8 +11,7 @@ import { useUIAMatrixError } from '$hooks/useUIAFlows'; import { DeviceVerificationStatus } from '$components/DeviceVerificationStatus'; import { VerificationStatus } from '$hooks/useDeviceVerificationStatus'; import { useAuthMetadata } from '$hooks/useAuthMetadata'; -import { withSearchParam } from '$pages/pathUtils'; -import { useAccountManagementActions } from '$hooks/useAccountManagement'; +import { getAccountManagementUrl, useAccountManagementActions } from '$hooks/useAccountManagement'; import { SettingTile } from '$components/setting-tile'; import { SequenceCardStyle } from '$features/settings/styles.css'; import { VerifyOtherDeviceTile } from './Verification'; @@ -63,12 +62,8 @@ export function OtherDevices({ devices, refreshDeviceList, showVerification }: O const [deleted, setDeleted] = useState(new Set()); const handleDashboardOIDC = useCallback(() => { - const authUrl = authMetadata?.account_management_uri ?? authMetadata?.issuer; - if (!authUrl) return; - - const url = withSearchParam(authUrl, { - action: accountManagementActions.sessionsList, - }); + const url = getAccountManagementUrl(authMetadata, accountManagementActions.sessionsList); + if (!url) return; if (isTauri()) { import('@tauri-apps/plugin-opener') .then(({ openUrl }) => openUrl(url)) @@ -80,13 +75,12 @@ export function OtherDevices({ devices, refreshDeviceList, showVerification }: O const handleDeleteOIDC = useCallback( (deviceId: string) => { - const authUrl = authMetadata?.account_management_uri ?? authMetadata?.issuer; - if (!authUrl) return; - - const url = withSearchParam(authUrl, { - action: accountManagementActions.sessionEnd, - device_id: deviceId, - }); + const url = getAccountManagementUrl( + authMetadata, + accountManagementActions.sessionEnd, + deviceId + ); + if (!url) return; if (isTauri()) { import('@tauri-apps/plugin-opener') .then(({ openUrl }) => openUrl(url)) @@ -159,7 +153,7 @@ export function OtherDevices({ devices, refreshDeviceList, showVerification }: O <> Others - {authMetadata && ( + {authMetadata?.account_management_uri && ( - ) : ( + ) : authMetadata ? undefined : ( { setMenuCords(undefined); - if (authMetadata) { - const authUrl = authMetadata.account_management_uri ?? authMetadata.issuer; - const url = withSearchParam(authUrl, { - action: accountManagementActions.crossSigningReset, - }); + const url = getAccountManagementUrl(authMetadata, accountManagementActions.crossSigningReset); + if (url) { if (isTauri()) { import('@tauri-apps/plugin-opener') .then(({ openUrl }) => openUrl(url)) diff --git a/src/app/hooks/useAccountManagement.ts b/src/app/hooks/useAccountManagement.ts index aeccd5814a..4e10c4287b 100644 --- a/src/app/hooks/useAccountManagement.ts +++ b/src/app/hooks/useAccountManagement.ts @@ -1,4 +1,20 @@ import { useMemo } from 'react'; +import type { ValidatedAuthMetadata } from '$types/matrix-sdk'; + +export const getAccountManagementUrl = ( + metadata: ValidatedAuthMetadata | undefined, + action: string, + deviceId?: string +): string | undefined => { + if (!metadata?.account_management_uri) return undefined; + + const url = new URL(metadata.account_management_uri); + if (metadata.account_management_actions_supported?.includes(action)) { + url.searchParams.set('action', action); + if (deviceId) url.searchParams.set('device_id', deviceId); + } + return url.toString(); +}; export const useAccountManagementActions = () => { const actions = useMemo( diff --git a/src/app/pages/App.tsx b/src/app/pages/App.tsx index 4d0682a4b0..ededc64a0e 100644 --- a/src/app/pages/App.tsx +++ b/src/app/pages/App.tsx @@ -18,6 +18,7 @@ import { createRouter } from './Router'; import { isReactQueryDevtoolsEnabled } from './reactQueryDevtoolsGate'; import { bootstrapSettingsStore } from '$state/settings'; import { AppShell } from '$components/app-shell'; +import { normalizeOAuthCallbackUrl } from '$utils/oauthCallback'; const queryClient = new QueryClient(); const ReactQueryDevtools = lazy(async () => { @@ -32,6 +33,7 @@ type BootstrappedAppShellProps = { }; function BootstrappedAppShell({ clientConfig, screenSize }: BootstrappedAppShellProps) { + normalizeOAuthCallbackUrl(clientConfig.hashRouter); const jotaiStoreRef = useRef>(); if (!jotaiStoreRef.current) { jotaiStoreRef.current = createStore(); diff --git a/src/app/pages/Router.tsx b/src/app/pages/Router.tsx index 60bd04f9b1..1976ae2894 100644 --- a/src/app/pages/Router.tsx +++ b/src/app/pages/Router.tsx @@ -176,6 +176,7 @@ export const createRouter = (clientConfig: ClientConfig, screenSize: ScreenSize) if (url.searchParams.get('addAccount') === '1') return null; if (url.searchParams.has('loginToken')) return null; if (url.searchParams.has('code') && url.searchParams.has('state')) return null; + if (url.searchParams.has('error') && url.searchParams.has('state')) return null; if (hasStoredSession()) return redirect(getHomePath()); return null; }} diff --git a/src/app/pages/TauriDeepLinkBridge.tsx b/src/app/pages/TauriDeepLinkBridge.tsx index e73852b13a..9ae148b0c4 100644 --- a/src/app/pages/TauriDeepLinkBridge.tsx +++ b/src/app/pages/TauriDeepLinkBridge.tsx @@ -5,9 +5,9 @@ import { createLogger } from '$utils/debug'; import { parseTauriOidcCallback, parseTauriSsoCallback, - takeTauriOidcServer, takeTauriSsoNonce, } from '$pages/auth/SSOTauri'; +import { getOauthContextServer } from '$pages/auth/login/oidcLoginUtil'; import { getLoginPath, withSearchParam } from './pathUtils'; const log = createLogger('TauriDeepLinkBridge'); @@ -27,10 +27,20 @@ export const mapDeepLinkToLoginPath = (rawUrl: string): string | undefined => { const oidcCallback = parseTauriOidcCallback(rawUrl); if (oidcCallback) { - return withSearchParam(getLoginPath(takeTauriOidcServer()), { - code: oidcCallback.code, - state: oidcCallback.state, - }); + const loginPath = getLoginPath(getOauthContextServer(oidcCallback.state)); + return 'code' in oidcCallback + ? withSearchParam(loginPath, { + code: oidcCallback.code, + state: oidcCallback.state, + }) + : withSearchParam(loginPath, { + error: oidcCallback.error, + ...(oidcCallback.errorDescription + ? { error_description: oidcCallback.errorDescription } + : {}), + ...(oidcCallback.errorUri ? { error_uri: oidcCallback.errorUri } : {}), + state: oidcCallback.state, + }); } return undefined; diff --git a/src/app/pages/auth/SSOTauri.ts b/src/app/pages/auth/SSOTauri.ts index 2a21e4b86e..6a60baa28b 100644 --- a/src/app/pages/auth/SSOTauri.ts +++ b/src/app/pages/auth/SSOTauri.ts @@ -63,32 +63,15 @@ export const TAURI_OIDC_CLIENT_URI = 'https://app.sable.moe'; const TAURI_OIDC_PROTOCOL = 'moe.sable.app:'; const TAURI_OIDC_PATH = '/login'; -const TAURI_OIDC_SERVER_KEY = 'sable:tauri-oidc:server'; export const buildTauriOidcRedirectUrl = (): string => `${TAURI_OIDC_PROTOCOL}${TAURI_OIDC_PATH}`; -export const rememberTauriOidcServer = (server?: string): void => { - try { - if (server) localStorage.setItem(TAURI_OIDC_SERVER_KEY, server); - else localStorage.removeItem(TAURI_OIDC_SERVER_KEY); - } catch { - // ignore storage failures - } -}; - -export const takeTauriOidcServer = (): string | undefined => { - try { - const server = localStorage.getItem(TAURI_OIDC_SERVER_KEY) ?? undefined; - localStorage.removeItem(TAURI_OIDC_SERVER_KEY); - return server; - } catch { - return undefined; - } -}; - export const parseTauriOidcCallback = ( rawUrl: string -): { code: string; state: string } | undefined => { +): + | { code: string; state: string } + | { error: string; errorDescription?: string; errorUri?: string; state: string } + | undefined => { try { const callbackUrl = new URL(rawUrl); if (callbackUrl.protocol !== TAURI_OIDC_PROTOCOL) return undefined; @@ -96,9 +79,19 @@ export const parseTauriOidcCallback = ( const code = callbackUrl.searchParams.get('code'); const state = callbackUrl.searchParams.get('state'); - if (!code || !state) return undefined; + if (!state) return undefined; - return { code, state }; + if (code) return { code, state }; + + const error = callbackUrl.searchParams.get('error'); + if (!error) return undefined; + + return { + error, + errorDescription: callbackUrl.searchParams.get('error_description') ?? undefined, + errorUri: callbackUrl.searchParams.get('error_uri') ?? undefined, + state, + }; } catch { return undefined; } diff --git a/src/app/pages/auth/login/Login.tsx b/src/app/pages/auth/login/Login.tsx index a02ffa279a..cdb3f91f0c 100644 --- a/src/app/pages/auth/login/Login.tsx +++ b/src/app/pages/auth/login/Login.tsx @@ -125,8 +125,10 @@ export function Login() { const oidcCode = searchParams.get('code') ?? undefined; const oidcState = searchParams.get('state') ?? undefined; - const oidcNotice = external.error - ? (external.errorDescription ?? `Sign-in was not completed (${external.error}).`) + const oidcError = searchParams.get('error') ?? external.error; + const oidcErrorDescription = searchParams.get('error_description') ?? external.errorDescription; + const oidcNotice = oidcError + ? (oidcErrorDescription ?? `Sign-in was not completed (${oidcError}).`) : undefined; const showOidc = authMetadata !== undefined; diff --git a/src/app/pages/auth/login/OidcLogin.tsx b/src/app/pages/auth/login/OidcLogin.tsx index d6f71b34c2..7bc48828b7 100644 --- a/src/app/pages/auth/login/OidcLogin.tsx +++ b/src/app/pages/auth/login/OidcLogin.tsx @@ -26,6 +26,8 @@ const oidcErrorMessage = (error: unknown): string => { return 'Sign-in could not be completed. Please try again.'; case OidcLoginError.MissingDeviceId: return 'The homeserver did not grant a device for this session.'; + case OidcLoginError.MissingRefreshToken: + return 'The authorization server did not issue the required refresh token.'; default: return 'Failed to sign in with single sign-on.'; } diff --git a/src/app/pages/auth/login/oidcLoginUtil.ts b/src/app/pages/auth/login/oidcLoginUtil.ts index 8b160aec96..f1ce71c475 100644 --- a/src/app/pages/auth/login/oidcLoginUtil.ts +++ b/src/app/pages/auth/login/oidcLoginUtil.ts @@ -10,11 +10,7 @@ import { openUrl } from '@tauri-apps/plugin-opener'; import { createLogger } from '$utils/debug'; import { fetch } from '$utils/fetch'; import type { Session } from '$state/sessions'; -import { - TAURI_OIDC_CLIENT_URI, - buildTauriOidcRedirectUrl, - rememberTauriOidcServer, -} from '$pages/auth/SSOTauri'; +import { TAURI_OIDC_CLIENT_URI, buildTauriOidcRedirectUrl } from '$pages/auth/SSOTauri'; const log = createLogger('oidcLogin'); @@ -22,9 +18,9 @@ const CLIENT_NAME = 'Sable'; export enum OidcLoginError { RegistrationFailed = 'RegistrationFailed', - UserCancelled = 'UserCancelled', CodeExchangeFailed = 'CodeExchangeFailed', MissingDeviceId = 'MissingDeviceId', + MissingRefreshToken = 'MissingRefreshToken', MissingOauthContext = 'MissingOauthContext', Unknown = 'Unknown', } @@ -36,32 +32,45 @@ export class OidcLoginFailure extends Error { } } -const OAUTH_CONTEXT_KEY = 'oauth_login_context'; +const OAUTH_CONTEXT_KEY_PREFIX = 'oauth_login_context:'; type OauthLoginContext = { - metadata: ValidatedAuthMetadata; + issuer: string; clientId: string; redirectUri: string; deviceId?: string; codeVerifier?: string; homeserverUrl: string; + server?: string; }; -const persistOauthContext = (state: string, ctx: OauthLoginContext): void => { - sessionStorage.setItem(OAUTH_CONTEXT_KEY, JSON.stringify({ state, ...ctx })); +const oauthContextKey = (state: string): string => `${OAUTH_CONTEXT_KEY_PREFIX}${state}`; + +export const persistOauthContext = (state: string, ctx: OauthLoginContext): void => { + // A public OAuth client must retain its one-time PKCE verifier across the full-page redirect. + // This per-tab context contains no access token, refresh token, or user credentials and is + // removed before code exchange. Encrypting it with a key available to the same origin would + // not protect against an attacker capable of reading sessionStorage. + sessionStorage.setItem(oauthContextKey(state), JSON.stringify(ctx)); }; -const consumeOauthContext = (state: string): OauthLoginContext | undefined => { - const raw = sessionStorage.getItem(OAUTH_CONTEXT_KEY); +export const consumeOauthContext = (state: string): OauthLoginContext | undefined => { + const key = oauthContextKey(state); + const raw = sessionStorage.getItem(key); if (!raw) return undefined; + sessionStorage.removeItem(key); try { - const stored = JSON.parse(raw) as OauthLoginContext & { state: string }; - if (stored.state !== state) return undefined; - sessionStorage.removeItem(OAUTH_CONTEXT_KEY); - const { state: _state, ...ctx } = stored; - return ctx; + return JSON.parse(raw) as OauthLoginContext; + } catch { + return undefined; + } +}; + +export const getOauthContextServer = (state: string): string | undefined => { + try { + const raw = sessionStorage.getItem(oauthContextKey(state)); + return raw ? (JSON.parse(raw) as OauthLoginContext).server : undefined; } catch { - sessionStorage.removeItem(OAUTH_CONTEXT_KEY); return undefined; } }; @@ -91,19 +100,23 @@ export const startOidcLogin = async ( const state = crypto.randomUUID(); const oauth2 = new OAuth2(authMetadata, { clientId, redirectUri: effectiveRedirectUri }); + const authUrl = await oauth2.generateAuthorizationCodeGrantUrl( + state, + tauri ? 'query' : 'fragment', + opts?.prompt + ); + persistOauthContext(state, { - metadata: authMetadata, + issuer: authMetadata.issuer, clientId, redirectUri: effectiveRedirectUri, deviceId: oauth2.context.deviceId, codeVerifier: oauth2.context.codeVerifier, homeserverUrl, + server: opts?.server, }); - const authUrl = await oauth2.generateAuthorizationCodeGrantUrl(state, 'query', opts?.prompt); - if (tauri) { - rememberTauriOidcServer(opts?.server); await openUrl(authUrl); return; } @@ -121,6 +134,13 @@ export const expiresInMsFromToken = (token: BearerTokenResponse): number | undef return undefined; }; +export const requireRefreshToken = (token: BearerTokenResponse): string => { + if (!token.refresh_token) { + throw new OidcLoginFailure(OidcLoginError.MissingRefreshToken); + } + return token.refresh_token; +}; + export type OidcLoginResult = { baseUrl: string; session: Session; @@ -133,7 +153,14 @@ export const completeOidcLogin = async (code: string, state: string): Promise => userId: session.userId, deviceId: session.deviceId, }); - const tokenRefresher = await createSessionTokenRefresher(session, tempClient); + const tokenRefresher = createSessionTokenRefresher(session, tempClient); const mx = createClient({ baseUrl: session.baseUrl, diff --git a/src/app/utils/oauthCallback.ts b/src/app/utils/oauthCallback.ts new file mode 100644 index 0000000000..3bce43c947 --- /dev/null +++ b/src/app/utils/oauthCallback.ts @@ -0,0 +1,34 @@ +import type { HashRouterConfig } from '$hooks/useClientConfig'; +import { trimSlash } from './common'; + +const getCallbackParams = (hash: string): URLSearchParams | undefined => { + const params = new URLSearchParams(hash.startsWith('#') ? hash.slice(1) : hash); + if (!params.has('state') || (!params.has('code') && !params.has('error'))) return undefined; + return params; +}; + +export const normalizeOAuthCallbackUrl = (hashRouter?: HashRouterConfig): void => { + const params = getCallbackParams(window.location.hash); + if (!params) return; + + const url = new URL(window.location.href); + url.hash = ''; + + if (hashRouter?.enabled) { + const basePath = `/${trimSlash(import.meta.env.BASE_URL ?? '')}`; + const appPath = url.pathname.startsWith(basePath) + ? url.pathname.slice(basePath.length) + : url.pathname; + const route = [hashRouter.basename, appPath] + .map((part) => trimSlash(part ?? '')) + .filter(Boolean) + .join('/'); + url.pathname = basePath; + url.search = ''; + url.hash = `/${route}?${params}`; + } else { + params.forEach((value, key) => url.searchParams.set(key, value)); + } + + window.history.replaceState(window.history.state, '', url); +}; diff --git a/src/client/initMatrix.ts b/src/client/initMatrix.ts index e52c156e9d..3d556e7946 100644 --- a/src/client/initMatrix.ts +++ b/src/client/initMatrix.ts @@ -10,21 +10,20 @@ import { IndexedDBStore, IndexedDBCryptoStore, KnownMembership, - OAuth2, SyncState, } from '$types/matrix-sdk'; import { fetch } from '$utils/fetch'; import { clearMediaCache } from '$utils/mediaCache'; -import { getAppOrigin } from '$utils/platform'; import { clearNavToActivePathStore } from '$state/navToActivePath'; import type { Session, SessionStoreName } from '$state/sessions'; -import { getSessionStoreName } from '$state/sessions'; +import { getSessionStoreName, getStoredSessionRefreshToken } from '$state/sessions'; import { createLogger } from '$utils/debug'; import { createDebugLogger } from '$utils/debugLogger'; import * as Sentry from '@sentry/react'; import { pushSessionToSW } from '../sw-session'; -import { createSessionTokenRefresher } from './oidcTokenRefresher'; +import { assertAuthMetadataIssuer, createSessionTokenRefresher } from './oidcTokenRefresher'; +import { revokeOAuthToken } from './oauthTokenRevocation'; import { cryptoCallbacks } from './secretStorageKeys'; import type { SlidingSyncDiagnostics } from './slidingSync'; import { scopeEphemeralExtensions, SlidingSyncManager } from './slidingSync'; @@ -252,7 +251,7 @@ const buildClient = async (session: Session): Promise => { userId: session.userId, deviceId: session.deviceId, }); - const tokenRefresher = await createSessionTokenRefresher(session, tempClient); + const tokenRefresher = createSessionTokenRefresher(session, tempClient); const mx = createClient({ baseUrl: session.baseUrl, @@ -553,24 +552,23 @@ export const getClientSyncDiagnostics = (mx: MatrixClient): ClientSyncDiagnostic }; const revokeOidcSession = async (mx: MatrixClient, session: Session): Promise => { - const clientId = session.oidc?.clientId; - if (!clientId) return; + const oidc = session.oidc; + if (!oidc) return; const metadata = await mx.getAuthMetadata(); + assertAuthMetadataIssuer(oidc.issuer, metadata); - const oauth2 = new OAuth2(metadata, { - clientId, - redirectUri: getAppOrigin(), - deviceId: session.deviceId, - }); + const refreshToken = getStoredSessionRefreshToken(session.userId) ?? session.refreshToken; + const token = refreshToken ?? mx.getAccessToken() ?? undefined; + if (!token) return; - const accessToken = mx.getAccessToken() ?? undefined; - const results = await Promise.allSettled([ - session.refreshToken - ? oauth2.revokeToken(session.refreshToken, 'refresh_token') - : Promise.resolve(), - accessToken ? oauth2.revokeToken(accessToken, 'access_token') : Promise.resolve(), - ]); - if (results.some((r) => r.status === 'rejected')) { + try { + await revokeOAuthToken( + metadata, + oidc.clientId, + token, + refreshToken ? 'refresh_token' : 'access_token' + ); + } catch { debugLog.warn('general', 'OIDC token revocation had failures', { userId: session.userId }); } }; diff --git a/src/client/oauthTokenRevocation.ts b/src/client/oauthTokenRevocation.ts new file mode 100644 index 0000000000..61b8cc0549 --- /dev/null +++ b/src/client/oauthTokenRevocation.ts @@ -0,0 +1,29 @@ +import type { ValidatedAuthMetadata } from '$types/matrix-sdk'; +import { fetch } from '$utils/fetch'; + +type Fetch = typeof globalThis.fetch; + +export const revokeOAuthToken = async ( + metadata: ValidatedAuthMetadata, + clientId: string, + token: string, + tokenTypeHint: 'access_token' | 'refresh_token', + fetchFn: Fetch = fetch +): Promise => { + const response = await fetchFn(metadata.revocation_endpoint, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/x-www-form-urlencoded', + }, + body: new URLSearchParams({ + token, + token_type_hint: tokenTypeHint, + client_id: clientId, + }), + }); + + if (!response.ok) { + throw new Error(`OAuth token revocation failed (${response.status})`); + } +}; diff --git a/src/client/oidcTokenRefresher.ts b/src/client/oidcTokenRefresher.ts index d8de063160..8e2db84934 100644 --- a/src/client/oidcTokenRefresher.ts +++ b/src/client/oidcTokenRefresher.ts @@ -11,19 +11,25 @@ import { getLocalStorageItem } from '$state/utils/atomWithLocalStorage'; import { pushSessionToSW } from '../sw-session'; import { getAppOrigin } from '$utils/platform'; -export const createSessionTokenRefresher = async ( - session: Session, - mx: MatrixClient -): Promise => { - if (!session.oidc || !session.refreshToken) return undefined; +export const assertAuthMetadataIssuer = ( + expectedIssuer: string, + metadata: ValidatedAuthMetadata +): void => { + if (metadata.issuer !== expectedIssuer) { + throw new Error( + `OAuth issuer changed for the stored session: expected ${expectedIssuer}, received ${metadata.issuer}` + ); + } +}; - const metadata: ValidatedAuthMetadata = await mx.getAuthMetadata(); +export type SessionTokenRefresher = Pick; - const oauth2 = new OAuth2(metadata, { - clientId: session.oidc.clientId, - redirectUri: getAppOrigin(), - deviceId: session.deviceId, - }); +export const createSessionTokenRefresher = ( + session: Session, + mx: MatrixClient +): SessionTokenRefresher | undefined => { + const { oidc } = session; + if (!oidc || !session.refreshToken) return undefined; const onRefresh = async (tokens: AccessTokens): Promise => { updateSessionTokens(session.userId, { @@ -38,10 +44,39 @@ export const createSessionTokenRefresher = async ( } }; - const tokenRefresher = new TokenRefresher(oauth2, onRefresh); - const refresh = tokenRefresher.tokenRefreshFunction; - // Another tab may have rotated the token; reusing a consumed one revokes the session. - tokenRefresher.tokenRefreshFunction = (refreshToken) => - refresh(getStoredSessionRefreshToken(session.userId) ?? refreshToken); - return tokenRefresher; + let tokenRefresherPromise: Promise | undefined; + const getTokenRefresher = (): Promise => { + if (!tokenRefresherPromise) { + tokenRefresherPromise = mx + .getAuthMetadata() + .then((metadata: ValidatedAuthMetadata) => { + assertAuthMetadataIssuer(oidc.issuer, metadata); + const oauth2 = new OAuth2(metadata, { + clientId: oidc.clientId, + redirectUri: getAppOrigin(), + deviceId: session.deviceId, + }); + return new TokenRefresher(oauth2, onRefresh); + }) + .catch((error: unknown) => { + tokenRefresherPromise = undefined; + throw error; + }); + } + return tokenRefresherPromise; + }; + + return { + tokenRefreshFunction: async (refreshToken) => { + const tokenRefresher = await getTokenRefresher(); + // Another tab may have rotated the token; reusing a consumed one revokes the session. + const latestRefreshToken = getStoredSessionRefreshToken(session.userId) ?? refreshToken; + const tokens = await tokenRefresher.tokenRefreshFunction(latestRefreshToken); + return { + ...tokens, + // OAuth servers may omit a replacement refresh token, in which case the old one remains valid. + refreshToken: tokens.refreshToken ?? latestRefreshToken, + }; + }, + }; };