diff --git a/e2e/specs/mock/auth.spec.ts b/e2e/specs/mock/auth.spec.ts index d965d4270f..7443a64fe9 100644 --- a/e2e/specs/mock/auth.spec.ts +++ b/e2e/specs/mock/auth.spec.ts @@ -5,6 +5,26 @@ import { getSecondaryE2EUser } from '../../setup/users.mock'; import cleanupUser from '../../setup/cleanupUser'; import { NEW_CHAT_PATH } from './helpers'; +type AuthRecoveryTestEvent = { + type: string; + detail: unknown; +}; + +type AuthRecoveryTestWindow = Window & { + __authRecoveryTestEvents: AuthRecoveryTestEvent[]; +}; + +type RefreshTokenBody = { + token?: string; +}; + +function createJwt(expiresAtMs: number) { + const payload = Buffer.from(JSON.stringify({ exp: Math.floor(expiresAtMs / 1000) })).toString( + 'base64url', + ); + return `header.${payload}.signature`; +} + async function getIsolatedStorageState(request: APIRequestContext, user: User) { await cleanupUser(user); @@ -79,4 +99,83 @@ test.describe('auth session', () => { await cleanupUser(user); } }); + + test('recovers from an expired bearer during app bootstrap without redirect looping', async ({ + page, + }) => { + test.setTimeout(30000); + + const expiredToken = createJwt(Date.now() - 60_000); + const expiredBearerPaths: string[] = []; + let refreshCalls = 0; + + await page.addInitScript(() => { + const testWindow = window as AuthRecoveryTestWindow; + testWindow.__authRecoveryTestEvents = []; + window.addEventListener('authRecovery', (event) => { + testWindow.__authRecoveryTestEvents.push({ + type: 'authRecovery', + detail: (event as CustomEvent).detail, + }); + }); + window.addEventListener('authRedirectStarted', (event) => { + testWindow.__authRecoveryTestEvents.push({ + type: 'authRedirectStarted', + detail: (event as CustomEvent).detail, + }); + }); + }); + + await page.route('**/api/**', async (route) => { + const request = route.request(); + const pathname = new URL(request.url()).pathname; + + if (pathname === '/api/auth/refresh') { + refreshCalls += 1; + const response = await route.fetch(); + if (refreshCalls === 1) { + const body = (await response.json()) as RefreshTokenBody; + await route.fulfill({ + response, + json: { + ...body, + token: expiredToken, + }, + }); + return; + } + + await route.fulfill({ response }); + return; + } + + if (request.headers().authorization === `Bearer ${expiredToken}`) { + expiredBearerPaths.push(pathname); + await route.fulfill({ + status: 401, + contentType: 'application/json', + json: { message: 'jwt expired' }, + }); + return; + } + + await route.continue(); + }); + + await page.goto(NEW_CHAT_PATH, { timeout: 10000 }); + + await expect(page).not.toHaveURL(/\/login/); + await expect(page.getByRole('textbox', { name: 'Message input' })).toBeVisible(); + await expect.poll(() => refreshCalls).toBe(2); + expect(expiredBearerPaths.length).toBeGreaterThan(0); + + const events = await page.evaluate( + () => (window as AuthRecoveryTestWindow).__authRecoveryTestEvents, + ); + + expect(events.filter((event) => event.type === 'authRedirectStarted')).toHaveLength(0); + expect( + events.filter((event) => event.type === 'authRecovery').map((event) => event.detail), + ).toEqual([{ state: 'started' }, { state: 'finished' }]); + }); }); diff --git a/packages/data-provider/specs/request-interceptor.spec.ts b/packages/data-provider/specs/request-interceptor.spec.ts index b5e43736cc..8de8d541a1 100644 --- a/packages/data-provider/specs/request-interceptor.spec.ts +++ b/packages/data-provider/specs/request-interceptor.spec.ts @@ -2,6 +2,7 @@ * @jest-environment @happy-dom/jest-environment */ import axios from 'axios'; +import type { InternalAxiosRequestConfig } from 'axios'; import { setTokenHeader } from '../src/headers-helpers'; /** @@ -20,6 +21,55 @@ const mockAdapter = jest.fn(); let originalAdapter: typeof axios.defaults.adapter; let savedLocation: Location; +type RetryableAdapterConfig = InternalAxiosRequestConfig & { _retry?: boolean }; + +function createAdapterResponse(config: InternalAxiosRequestConfig, data: unknown = {}) { + return Promise.resolve({ + data, + status: 200, + headers: {}, + config, + }); +} + +function create401Error(config: InternalAxiosRequestConfig) { + return Promise.reject({ + response: { status: 401 }, + config, + }); +} + +function getCallsForUrl(urlPart: string) { + return mockAdapter.mock.calls.filter(([config]) => config.url?.includes(urlPart) === true); +} + +function createDeferred() { + let resolve!: (value: T) => void; + let reject!: (reason?: unknown) => void; + const promise = new Promise((res, rej) => { + resolve = res; + reject = rej; + }); + return { promise, resolve, reject }; +} + +async function waitForAdapterCall(urlPart: string) { + for (let i = 0; i < 10; i++) { + if (getCallsForUrl(urlPart).length > 0) { + return; + } + await new Promise((resolve) => setTimeout(resolve, 0)); + } + throw new Error(`Adapter was not called for ${urlPart}`); +} + +function createJwt(expiresAtMs: number) { + const payload = Buffer.from(JSON.stringify({ exp: Math.floor(expiresAtMs / 1000) })).toString( + 'base64url', + ); + return `header.${payload}.signature`; +} + beforeAll(async () => { originalAdapter = axios.defaults.adapter; axios.defaults.adapter = mockAdapter; @@ -38,6 +88,8 @@ afterAll(() => { afterEach(() => { delete axios.defaults.headers.common['Authorization']; + window.localStorage.clear(); + delete (window as Window & { __librechatAuthRecovery?: unknown }).__librechatAuthRecovery; Object.defineProperty(window, 'location', { value: savedLocation, writable: true, @@ -53,6 +105,27 @@ function setWindowLocation(overrides: Partial) { }); } +function setTrackedWindowLocation(overrides: Partial) { + let href = overrides.href ?? window.location.href; + const hrefWrites: string[] = []; + Object.defineProperty(window, 'location', { + value: { + ...window.location, + ...overrides, + get href() { + return href; + }, + set href(value: string) { + hrefWrites.push(value); + href = value; + }, + }, + writable: true, + configurable: true, + }); + return hrefWrites; +} + describe('axios 401 interceptor — Authorization header guard', () => { it('skips refresh and rejects when Authorization header is cleared', async () => { expect.assertions(1); @@ -301,4 +374,160 @@ describe('axios 401 interceptor — Authorization header guard', () => { const refreshCall = mockAdapter.mock.calls[1]; expect(refreshCall[0].url).toContain('api/auth/refresh'); }); + + it('coalesces concurrent 401 responses into one refresh and retries with the new token', async () => { + expect.assertions(3); + setTokenHeader('expired-token'); + + mockAdapter.mockImplementation((config: RetryableAdapterConfig) => { + if (config.url?.includes('/api/auth/refresh') === true) { + return createAdapterResponse(config, { token: 'new-token' }); + } + if (config._retry === true) { + return createAdapterResponse(config, { ok: true }); + } + return create401Error(config); + }); + + const responses = await Promise.all([ + axios.get('/api/messages'), + axios.get('/api/convos'), + axios.get('/api/files'), + ]); + + expect(responses.map((response) => response.data)).toEqual([ + { ok: true }, + { ok: true }, + { ok: true }, + ]); + + expect(getCallsForUrl('/api/auth/refresh')).toHaveLength(1); + expect( + mockAdapter.mock.calls + .filter(([config]) => (config as RetryableAdapterConfig)._retry === true) + .every(([config]) => config.headers?.Authorization === 'Bearer new-token'), + ).toBe(true); + }); + + it('holds new requests behind an in-flight auth recovery', async () => { + expect.assertions(4); + setTokenHeader('expired-token'); + const refresh = createDeferred(); + + mockAdapter.mockImplementation((config: RetryableAdapterConfig) => { + if (config.url?.includes('/api/auth/refresh') === true) { + return refresh.promise.then((token) => createAdapterResponse(config, { token })); + } + if (config.url === '/api/messages' && config._retry !== true) { + return create401Error(config); + } + return createAdapterResponse(config, { ok: true }); + }); + + const firstRequest = axios.get('/api/messages'); + await waitForAdapterCall('/api/auth/refresh'); + + const secondRequest = axios.get('/api/projects'); + await Promise.resolve(); + + expect(getCallsForUrl('/api/projects')).toHaveLength(0); + + refresh.resolve('new-token'); + const responses = await Promise.all([firstRequest, secondRequest]); + expect(responses.map((response) => response.data)).toEqual([{ ok: true }, { ok: true }]); + + expect(getCallsForUrl('/api/auth/refresh')).toHaveLength(1); + expect(getCallsForUrl('/api/projects')[0][0].headers?.Authorization).toBe('Bearer new-token'); + }); + + it('redirects once when a burst of 401s cannot refresh a token', async () => { + expect.assertions(3); + setTokenHeader('expired-token'); + const hrefWrites = setTrackedWindowLocation({ + href: 'http://localhost/c/race', + pathname: '/c/race', + search: '', + hash: '', + } as Partial); + + mockAdapter.mockImplementation((config: InternalAxiosRequestConfig) => { + if (config.url?.includes('/api/auth/refresh') === true) { + return createAdapterResponse(config, { token: '' }); + } + return create401Error(config); + }); + + await Promise.allSettled([ + axios.get('/api/messages'), + axios.get('/api/convos'), + axios.get('/api/files'), + ]); + + expect(getCallsForUrl('/api/auth/refresh')).toHaveLength(1); + expect(hrefWrites).toHaveLength(1); + expect(hrefWrites[0]).toBe('/login?redirect_to=%2Fc%2Frace'); + }); + + it('keeps redirect deduping when the storage timestamp is corrupted', async () => { + expect.assertions(2); + setTokenHeader('expired-token'); + const hrefWrites = setTrackedWindowLocation({ + href: 'http://localhost/c/race', + pathname: '/c/race', + search: '', + hash: '', + } as Partial); + + mockAdapter.mockImplementation((config: InternalAxiosRequestConfig) => { + if (config.url?.includes('/api/auth/refresh') === true) { + return createAdapterResponse(config, { token: '' }); + } + return create401Error(config); + }); + + await axios.get('/api/messages').catch(() => undefined); + window.localStorage.setItem('librechat.auth.redirect.startedAt', 'not-a-number'); + await axios.get('/api/convos').catch(() => undefined); + + expect(getCallsForUrl('/api/auth/refresh')).toHaveLength(1); + expect(hrefWrites).toEqual(['/login?redirect_to=%2Fc%2Frace']); + }); + + it('refreshes a near-expiry bearer token before sending a request', async () => { + expect.assertions(4); + setTokenHeader(createJwt(Date.now() + 60_000)); + + mockAdapter.mockImplementation((config: InternalAxiosRequestConfig) => { + if (config.url?.includes('/api/auth/refresh') === true) { + return createAdapterResponse(config, { token: 'fresh-token' }); + } + return createAdapterResponse(config, { ok: true }); + }); + + const response = await axios.get('/api/messages'); + + expect(response.data).toEqual({ ok: true }); + + expect(mockAdapter.mock.calls[0][0].url).toContain('/api/auth/refresh'); + expect(mockAdapter.mock.calls[1][0].url).toBe('/api/messages'); + expect(mockAdapter.mock.calls[1][0].headers?.Authorization).toBe('Bearer fresh-token'); + }); + + it('does not wait on the in-flight recovery when the refresh request itself fails', async () => { + expect.assertions(3); + setTokenHeader(createJwt(Date.now() + 60_000)); + + mockAdapter.mockImplementation((config: InternalAxiosRequestConfig) => { + if (config.url?.includes('/api/auth/refresh') === true) { + return create401Error(config); + } + return createAdapterResponse(config, { ok: true }); + }); + + const response = await axios.get('/api/messages'); + + expect(response.data).toEqual({ ok: true }); + expect(getCallsForUrl('/api/auth/refresh')).toHaveLength(1); + expect(getCallsForUrl('/api/messages')).toHaveLength(1); + }); }); diff --git a/packages/data-provider/src/request.ts b/packages/data-provider/src/request.ts index 5021b150e3..d6a1a85dea 100644 --- a/packages/data-provider/src/request.ts +++ b/packages/data-provider/src/request.ts @@ -1,8 +1,8 @@ /* eslint-disable @typescript-eslint/no-explicit-any */ -import axios, { AxiosError, AxiosRequestConfig } from 'axios'; +import axios, { AxiosRequestConfig } from 'axios'; +import type * as t from './types'; import { setTokenHeader } from './headers-helpers'; import * as endpoints from './api-endpoints'; -import type * as t from './types'; async function _get(url: string, options?: AxiosRequestConfig): Promise { const response = await axios.get(url, { ...options }); @@ -61,41 +61,215 @@ async function _patch(url: string, data?: any) { return response.data; } -let isRefreshing = false; -let failedQueue: { resolve: (value?: any) => void; reject: (reason?: any) => void }[] = []; +const AUTH_RECOVERY_EVENT = 'authRecovery'; +const AUTH_REDIRECT_EVENT = 'authRedirectStarted'; +const AUTH_REDIRECT_STORAGE_KEY = 'librechat.auth.redirect.startedAt'; +const AUTH_REDIRECT_DEDUPE_MS = 15_000; +const TOKEN_REFRESH_BUFFER_MS = 2 * 60 * 1000; + +type RetryableAxiosRequestConfig = AxiosRequestConfig & { _retry?: boolean }; + +type AuthRecoveryState = { + lastRedirectStartedAt: number; + refreshPromise: Promise | null; +}; + +type AuthRecoveryWindow = Window & { + __librechatAuthRecovery?: AuthRecoveryState; +}; const refreshToken = (retry?: boolean): Promise => _post(endpoints.refreshToken(retry)); const dispatchTokenUpdatedEvent = (token: string) => { setTokenHeader(token); + clearAuthRedirectStartedAt(); window.dispatchEvent(new CustomEvent('tokenUpdated', { detail: token })); }; -const processQueue = (error: AxiosError | null, token: string | null = null) => { - failedQueue.forEach((prom) => { - if (error) { - prom.reject(error); - } else { - prom.resolve(token); - } - }); - failedQueue = []; +const getAuthRecoveryState = (): AuthRecoveryState => { + const browserWindow = window as AuthRecoveryWindow; + browserWindow.__librechatAuthRecovery ??= { + lastRedirectStartedAt: 0, + refreshPromise: null, + }; + return browserWindow.__librechatAuthRecovery; +}; + +const getAuthRedirectStartedAt = () => { + const state = getAuthRecoveryState(); + try { + const startedAt = window.localStorage.getItem(AUTH_REDIRECT_STORAGE_KEY); + const storedStartedAt = startedAt != null ? Number(startedAt) : 0; + const finiteStartedAt = Number.isFinite(storedStartedAt) ? storedStartedAt : 0; + return Math.max(finiteStartedAt, state.lastRedirectStartedAt); + } catch { + return state.lastRedirectStartedAt; + } +}; + +const setAuthRedirectStartedAt = () => { + const state = getAuthRecoveryState(); + state.lastRedirectStartedAt = Date.now(); + try { + window.localStorage.setItem(AUTH_REDIRECT_STORAGE_KEY, String(state.lastRedirectStartedAt)); + } catch { + // localStorage can be blocked in embedded/private contexts. + } +}; + +const clearAuthRedirectStartedAt = () => { + const state = getAuthRecoveryState(); + state.lastRedirectStartedAt = 0; + try { + window.localStorage.removeItem(AUTH_REDIRECT_STORAGE_KEY); + } catch { + // Ignore unavailable storage. + } +}; + +const isAuthRedirectInProgress = () => { + const startedAt = getAuthRedirectStartedAt(); + return ( + Number.isFinite(startedAt) && startedAt > 0 && Date.now() - startedAt < AUTH_REDIRECT_DEDUPE_MS + ); +}; + +const dispatchAuthRecoveryEvent = (state: 'started' | 'finished') => { + window.dispatchEvent(new CustomEvent(AUTH_RECOVERY_EVENT, { detail: { state } })); +}; + +const setRequestAuthorizationHeader = (config: AxiosRequestConfig, token: string) => { + const headers = (config.headers ?? {}) as Record; + headers['Authorization'] = `Bearer ${token}`; + config.headers = headers; +}; + +const isAuthRecoveryEndpoint = (url?: string) => + url?.includes('/api/auth/2fa') === true || + url?.includes('/api/auth/logout') === true || + url?.includes('/api/auth/refresh') === true; + +const startAuthRecovery = (retryRefresh?: boolean) => { + const state = getAuthRecoveryState(); + if (state.refreshPromise) { + return state.refreshPromise; + } + + dispatchAuthRecoveryEvent('started'); + state.refreshPromise = refreshToken(retryRefresh) + .then((response) => { + const token = response?.token ?? ''; + if (!token) { + return null; + } + dispatchTokenUpdatedEvent(token); + return token; + }) + .finally(() => { + state.refreshPromise = null; + dispatchAuthRecoveryEvent('finished'); + }); + + return state.refreshPromise; +}; + +const redirectToLoginOnce = () => { + if (isAuthRedirectInProgress()) { + return; + } + + const href = endpoints.apiBaseUrl() + endpoints.buildLoginRedirectUrl(); + setAuthRedirectStartedAt(); + window.dispatchEvent(new CustomEvent(AUTH_REDIRECT_EVENT, { detail: { href } })); + window.location.href = href; +}; + +const getBearerToken = () => { + const authorization = axios.defaults.headers.common['Authorization']; + if (typeof authorization !== 'string' || !authorization.startsWith('Bearer ')) { + return null; + } + return authorization.slice('Bearer '.length); +}; + +const getJwtExpiryMs = (token: string) => { + const payload = token.split('.')[1]; + if (!payload) { + return null; + } + + try { + const normalizedPayload = payload.replace(/-/g, '+').replace(/_/g, '/'); + const paddedPayload = normalizedPayload.padEnd( + normalizedPayload.length + ((4 - (normalizedPayload.length % 4)) % 4), + '=', + ); + const decodedPayload = JSON.parse(window.atob(paddedPayload)) as { exp?: number }; + return typeof decodedPayload.exp === 'number' ? decodedPayload.exp * 1000 : null; + } catch { + return null; + } +}; + +const shouldRefreshBeforeRequest = (url?: string) => { + if (isAuthRecoveryEndpoint(url) || isAuthRedirectInProgress()) { + return false; + } + + const token = getBearerToken(); + if (!token) { + return false; + } + + const expiresAt = getJwtExpiryMs(token); + if (expiresAt == null) { + return false; + } + + const timeUntilExpiry = expiresAt - Date.now(); + return timeUntilExpiry > 0 && timeUntilExpiry <= TOKEN_REFRESH_BUFFER_MS; }; if (typeof window !== 'undefined') { + axios.interceptors.request.use(async (config) => { + const state = getAuthRecoveryState(); + if (state.refreshPromise && !isAuthRecoveryEndpoint(config.url)) { + const token = await state.refreshPromise.catch(() => null); + if (token) { + setRequestAuthorizationHeader(config, token); + } + return config; + } + + if (!shouldRefreshBeforeRequest(config.url)) { + return config; + } + + const token = await startAuthRecovery(false).catch(() => null); + if (token) { + setRequestAuthorizationHeader(config, token); + } + return config; + }); + axios.interceptors.response.use( (response) => response, async (error) => { - const originalRequest = error.config; + const originalRequest = error.config as RetryableAxiosRequestConfig | undefined; if (!error.response) { return Promise.reject(error); } - - if (originalRequest.url?.includes('/api/auth/2fa') === true) { + if (!originalRequest) { return Promise.reject(error); } - if (originalRequest.url?.includes('/api/auth/logout') === true) { + + const isRefreshRequest = originalRequest.url?.includes('/api/auth/refresh') === true; + if (isAuthRecoveryEndpoint(originalRequest.url) && !isRefreshRequest) { + return Promise.reject(error); + } + + if (isRefreshRequest && getAuthRecoveryState().refreshPromise) { return Promise.reject(error); } @@ -108,46 +282,32 @@ if (typeof window !== 'undefined') { return Promise.reject(error); } + if (isAuthRedirectInProgress()) { + return Promise.reject(error); + } + if (error.response.status === 401 && !originalRequest._retry) { - console.warn('401 error, refreshing token'); + const hasActiveRecovery = getAuthRecoveryState().refreshPromise != null; + if (!hasActiveRecovery) { + console.warn('401 error, refreshing token'); + } originalRequest._retry = true; - if (isRefreshing) { - try { - const token = await new Promise((resolve, reject) => { - failedQueue.push({ resolve, reject }); - }); - originalRequest.headers['Authorization'] = 'Bearer ' + token; - return await axios(originalRequest); - } catch (err) { - return Promise.reject(err); - } - } - - isRefreshing = true; - try { - const response = await refreshToken( + const token = await startAuthRecovery( // Handle edge case where we get a blank screen if the initial 401 error is from a refresh token request - originalRequest.url?.includes('api/auth/refresh') === true ? true : false, + isRefreshRequest, ); - const token = response?.token ?? ''; - if (token) { - originalRequest.headers['Authorization'] = 'Bearer ' + token; - dispatchTokenUpdatedEvent(token); - processQueue(null, token); + setRequestAuthorizationHeader(originalRequest, token); return await axios(originalRequest); - } else { - processQueue(error, null); - window.location.href = endpoints.apiBaseUrl() + endpoints.buildLoginRedirectUrl(); } + + redirectToLoginOnce(); + return Promise.reject(error); } catch (err) { - processQueue(err as AxiosError, null); return Promise.reject(err); - } finally { - isRefreshing = false; } }