diff --git a/api/server/controllers/assistants/v1.js b/api/server/controllers/assistants/v1.js index 7116edda75..2f53c07632 100644 --- a/api/server/controllers/assistants/v1.js +++ b/api/server/controllers/assistants/v1.js @@ -44,6 +44,7 @@ const createAssistant = async (req, res) => { const { toolDefinitions, accessibleServerNames } = await getAssistantToolDefinitions({ req, + res, tools, }); const healedTools = await healMcpToolNames({ @@ -172,6 +173,7 @@ const patchAssistant = async (req, res) => { const { toolDefinitions, accessibleServerNames } = await getAssistantToolDefinitions({ req, + res, tools: updateData.tools, }); const healedTools = await healMcpToolNames({ diff --git a/api/server/controllers/assistants/v2.js b/api/server/controllers/assistants/v2.js index afaf04f5d0..e6e9a5d0ef 100644 --- a/api/server/controllers/assistants/v2.js +++ b/api/server/controllers/assistants/v2.js @@ -34,6 +34,7 @@ const createAssistant = async (req, res) => { const { toolDefinitions, accessibleServerNames } = await getAssistantToolDefinitions({ req, + res, tools, }); const healedTools = await healMcpToolNames({ @@ -153,6 +154,7 @@ const updateAssistant = async ({ req, openai, assistant_id, updateData }) => { let hasFileSearch = false; const { toolDefinitions, accessibleServerNames } = await getAssistantToolDefinitions({ req, + res: req.res, tools: updateData.tools, }); const healedTools = await healMcpToolNames({ diff --git a/api/server/services/MCP.js b/api/server/services/MCP.js index dc2699f9a3..725554b523 100644 --- a/api/server/services/MCP.js +++ b/api/server/services/MCP.js @@ -369,12 +369,24 @@ async function healMcpToolNames({ req, tools, toolDefinitions, accessibleServerN * server slices instead of relying on the static aggregate cache. * @param {object} params * @param {ServerRequest} params.req + * @param {ServerResponse} [params.res] * @param {Array} [params.tools] * @returns {Promise} */ -async function getAssistantToolDefinitions({ req, tools }) { +async function getAssistantToolDefinitions({ req, res, tools }) { const registry = getMCPServersRegistry(); const appConfig = await getAppConfigForRequest(req); + const oboIdentityContext = createAuthIdentityContext({ + user: req.user, + tenantId: getTenantId(), + }); + const upstreamTokenProvider = createOpenIDSessionTokenProvider({ + req, + res: res ?? req.res, + user: req.user, + identityContext: oboIdentityContext, + tokenPreference: 'access_token', + }); return await loadAssistantToolDefinitions( { user: req.user, @@ -401,26 +413,12 @@ async function getAssistantToolDefinitions({ req, tools }) { servers: [serverName], findPluginAuthsByKeys, }); - /** - * An OBO server refuses to connect without a live upstream-token closure, so this - * recovery has to carry one like every other reinit path. There is no `res` on an - * assistant write; a refresh that rotates the token falls back to the recovery bridge. - */ - const oboIdentityContext = createAuthIdentityContext({ - user: req.user, - tenantId: getTenantId(), - }); const result = await reinitMCPServer({ user: req.user, serverName, serverConfig, userMCPAuthMap, - upstreamTokenProvider: createOpenIDSessionTokenProvider({ - req, - user: req.user, - identityContext: oboIdentityContext, - tokenPreference: 'access_token', - }), + upstreamTokenProvider, oboIdentityContext, }); return result?.availableTools ?? null; diff --git a/api/server/services/OpenIDRefreshFlight.js b/api/server/services/OpenIDRefreshFlight.js index fbb2555e4d..c9d2e18070 100644 --- a/api/server/services/OpenIDRefreshFlight.js +++ b/api/server/services/OpenIDRefreshFlight.js @@ -8,6 +8,7 @@ const DEFAULT_FLIGHT_TTL_MS = 2 * 60 * 1000; const DEFAULT_LOCK_TTL_MS = 30 * 1000; const DEFAULT_WAIT_TIMEOUT_MS = DEFAULT_LOCK_TTL_MS + 1000; const DEFAULT_WAIT_INTERVAL_MS = 100; +const DEFAULT_HEARTBEAT_INTERVAL_MS = 10 * 1000; const INTERNAL_BROWSER_REFRESH_TOKEN_FIELD = '__browserRefreshToken'; function sha256(value) { @@ -83,6 +84,74 @@ async function completeOpenIDRefreshFlight({ key, ownerId, tokens, ttl = DEFAULT }); } +async function renewOpenIDRefreshFlight({ + key, + ownerId, + lockTtl = DEFAULT_LOCK_TTL_MS, + ttl = DEFAULT_FLIGHT_TTL_MS, +}) { + if (!key || !ownerId) { + return null; + } + + return db.renewOpenIDRefreshFlight({ + key, + ownerId, + lockExpiresAt: new Date(Date.now() + lockTtl), + expiresAt: new Date(Date.now() + ttl), + }); +} + +/** + * Keeps the Mongo lease alive while the IdP grant is in progress. Without + * renewal, a slow-but-live owner can outlast the lease and a follower may + * reclaim the flight, admitting a second rotating refresh-token grant. + */ +async function withOpenIDRefreshFlightLease({ + key, + ownerId, + operation, + heartbeatInterval = DEFAULT_HEARTBEAT_INTERVAL_MS, + lockTtl = DEFAULT_LOCK_TTL_MS, + ttl = DEFAULT_FLIGHT_TTL_MS, +}) { + if (!key || !ownerId) { + return operation(); + } + + let renewalPromise = null; + const heartbeat = setInterval(() => { + if (renewalPromise) { + return; + } + renewalPromise = renewOpenIDRefreshFlight({ key, ownerId, lockTtl, ttl }) + .then((flight) => { + if (!flight) { + logger.warn('[OpenIDRefreshFlight] Refresh flight lease ownership was lost', { key }); + } + }) + .catch((error) => { + logger.warn('[OpenIDRefreshFlight] Failed to renew refresh flight lease', { + key, + error: error?.message, + }); + }) + .finally(() => { + renewalPromise = null; + }); + }, heartbeatInterval); + heartbeat.unref?.(); + + try { + return await operation(); + } finally { + clearInterval(heartbeat); + if (renewalPromise) { + await renewalPromise; + } + } +} + async function failOpenIDRefreshFlight({ key, ownerId, error, ttl = DEFAULT_FLIGHT_TTL_MS }) { if (!key || !ownerId) { return null; @@ -157,7 +226,9 @@ module.exports = { completeOpenIDRefreshFlight, createOpenIDRefreshFlightKey, failOpenIDRefreshFlight, + renewOpenIDRefreshFlight, waitForOpenIDRefreshFlight, + withOpenIDRefreshFlightLease, __internals: { sha256, readCompletedFlight, @@ -165,5 +236,6 @@ module.exports = { DEFAULT_LOCK_TTL_MS, DEFAULT_WAIT_TIMEOUT_MS, DEFAULT_WAIT_INTERVAL_MS, + DEFAULT_HEARTBEAT_INTERVAL_MS, }, }; diff --git a/api/server/services/OpenIDRefreshFlight.spec.js b/api/server/services/OpenIDRefreshFlight.spec.js index 8206cbcbf1..ae1e45ecd3 100644 --- a/api/server/services/OpenIDRefreshFlight.spec.js +++ b/api/server/services/OpenIDRefreshFlight.spec.js @@ -6,11 +6,28 @@ jest.mock('@librechat/data-schemas', () => ({ decryptV2: jest.fn(async (value) => value.replace(/^encrypted:/, '')), })); +jest.mock('@librechat/api', () => ({ + createOpenIDRefreshIdentityTuple: ({ user, requestUser }) => { + const subject = user?.openidId || user?.id || requestUser?.openidId || requestUser?.id; + if (!subject) { + return null; + } + return { + tenantId: user?.tenantId || requestUser?.tenantId || 'no-tenant', + openidIssuer: user?.openidIssuer || requestUser?.openidIssuer || 'no-issuer', + subject, + }; + }, + serializeAuthIdentityTuple: ({ tenantId, openidIssuer, subject }) => + [tenantId, openidIssuer, subject].join('\x1f'), +})); + jest.mock('~/models', () => ({ acquireOpenIDRefreshFlight: jest.fn(), completeOpenIDRefreshFlight: jest.fn(), failOpenIDRefreshFlight: jest.fn(), findOpenIDRefreshFlight: jest.fn(), + renewOpenIDRefreshFlight: jest.fn(), })); const { encryptV2, decryptV2 } = require('@librechat/data-schemas'); @@ -20,7 +37,9 @@ const { completeOpenIDRefreshFlight, createOpenIDRefreshFlightKey, failOpenIDRefreshFlight, + renewOpenIDRefreshFlight, waitForOpenIDRefreshFlight, + withOpenIDRefreshFlightLease, __internals, } = require('./OpenIDRefreshFlight'); @@ -31,6 +50,7 @@ describe('OpenIDRefreshFlight', () => { db.completeOpenIDRefreshFlight.mockResolvedValue({}); db.failOpenIDRefreshFlight.mockResolvedValue({}); db.findOpenIDRefreshFlight.mockResolvedValue(null); + db.renewOpenIDRefreshFlight.mockResolvedValue({ ownerId: 'owner-1', status: 'pending' }); }); it('creates a stable hash key from session, user, issuer, tenant, and refresh token', () => { @@ -110,6 +130,55 @@ describe('OpenIDRefreshFlight', () => { }); }); + it('renews only the owning pending flight lease', async () => { + await renewOpenIDRefreshFlight({ + key: 'flight-key', + ownerId: 'owner-1', + lockTtl: 30000, + ttl: 60000, + }); + + expect(db.renewOpenIDRefreshFlight).toHaveBeenCalledWith({ + key: 'flight-key', + ownerId: 'owner-1', + lockExpiresAt: expect.any(Date), + expiresAt: expect.any(Date), + }); + }); + + it('keeps renewing a leader lease until its refresh operation settles', async () => { + jest.useFakeTimers(); + let resolveOperation; + const operation = jest.fn( + () => + new Promise((resolve) => { + resolveOperation = resolve; + }), + ); + + try { + const resultPromise = withOpenIDRefreshFlightLease({ + key: 'flight-key', + ownerId: 'owner-1', + heartbeatInterval: 1000, + lockTtl: 30000, + ttl: 60000, + operation, + }); + + await jest.advanceTimersByTimeAsync(1000); + expect(db.renewOpenIDRefreshFlight).toHaveBeenCalledTimes(1); + + resolveOperation('tokens'); + await expect(resultPromise).resolves.toBe('tokens'); + + await jest.advanceTimersByTimeAsync(2000); + expect(db.renewOpenIDRefreshFlight).toHaveBeenCalledTimes(1); + } finally { + jest.useRealTimers(); + } + }); + it('encrypts completed token results before storing them', async () => { const tokens = { access_token: 'access', diff --git a/api/server/services/OpenIDSessionRefresh.js b/api/server/services/OpenIDSessionRefresh.js index cb5aa30d69..64d37cef48 100644 --- a/api/server/services/OpenIDSessionRefresh.js +++ b/api/server/services/OpenIDSessionRefresh.js @@ -1,4 +1,5 @@ const jwt = require('jsonwebtoken'); +const cookies = require('cookie'); const crypto = require('node:crypto'); const openIdClient = require('openid-client'); const { logger, DEFAULT_REFRESH_TOKEN_EXPIRY } = require('@librechat/data-schemas'); @@ -25,6 +26,7 @@ const { createOpenIDRefreshFlightKey, failOpenIDRefreshFlight, waitForOpenIDRefreshFlight, + withOpenIDRefreshFlightLease, } = require('./OpenIDRefreshFlight'); /** @@ -79,9 +81,9 @@ const IDENTITY_PART_SEPARATOR = '\x1f'; * session coalesces into one IdP refresh-token grant. Mirrors the * single-flight pattern in `OboTokenService.js`. * - * Process-local: multi-worker deployments may double-refresh on the very - * first concurrent miss across workers — acceptable because the IdP accepts - * both and the session store uses last-write-wins. + * Process-local coalescing is backed by a renewable Mongo lease in + * `performIdpRefresh`, so distinct workers do not admit parallel rotating-token + * grants for the same key. */ const inFlightRefreshes = new Map(); @@ -167,6 +169,55 @@ function resolveExpectedOpenIDSessionIdentity(req, user, identityContext) { }); } +function hasAnyOpenIDSessionIdentity(sessionTokens) { + return ['appUserId', 'openidSubject', 'tenantId', 'openidIssuer'].some( + (field) => sessionTokens?.[field] != null, + ); +} + +function canBindLegacyOpenIDSession(req, sessionTokens, expectedIdentity) { + if ( + hasAnyOpenIDSessionIdentity(sessionTokens) || + !expectedIdentity.appUserId || + !expectedIdentity.openidSubject || + !process.env.JWT_REFRESH_SECRET + ) { + return false; + } + + const parsedCookies = req?.headers?.cookie ? cookies.parse(req.headers.cookie) : {}; + const browserRefreshToken = parsedCookies.refreshToken; + const expectedBrowserRefreshToken = + sessionTokens.browserRefreshToken || sessionTokens.refreshToken; + if ( + !browserRefreshToken || + !expectedBrowserRefreshToken || + browserRefreshToken !== expectedBrowserRefreshToken || + !parsedCookies.openid_user_id + ) { + return false; + } + + try { + const marker = jwt.verify(parsedCookies.openid_user_id, process.env.JWT_REFRESH_SECRET); + if ( + typeof marker !== 'object' || + marker == null || + marker.id !== expectedIdentity.appUserId || + typeof marker.refreshTokenHash !== 'string' + ) { + return false; + } + const refreshTokenHash = crypto + .createHash('sha256') + .update(browserRefreshToken) + .digest('base64url'); + return marker.refreshTokenHash === refreshTokenHash; + } catch { + return false; + } +} + function assertOpenIDSessionIdentityMatch(req, user, identityContext) { const sessionTokens = req?.session?.openidTokens; if (!sessionTokens) { @@ -178,6 +229,22 @@ function assertOpenIDSessionIdentityMatch(req, user, identityContext) { return; } + /** + * Sessions minted before identity stamping was deployed have none of these + * fields. During a rolling upgrade, bind that legacy record only when the + * signed browser marker proves the current app user and refresh-token cookie + * are the ones that created it. Partial or unverifiable metadata still fails + * closed, preventing cross-user token adoption. + */ + if (canBindLegacyOpenIDSession(req, sessionTokens, expectedIdentity)) { + Object.assign(sessionTokens, expectedIdentity); + return persistSession(req).then(() => { + logger.info('[OpenIDSessionRefresh] Bound verified legacy OpenID session identity', { + userId: expectedIdentity.appUserId, + }); + }); + } + logger.warn('[OpenIDSessionRefresh] OpenID session token identity mismatch; refusing reuse', { userId: expectedIdentity.appUserId, has_session_user_id: Boolean(sessionTokens.appUserId), @@ -601,11 +668,8 @@ async function performIdpRefresh(req, res, user, tokenPreference, identityContex try { flight = await acquireOpenIDRefreshFlight({ key }); } catch (error) { - logger.warn( - '[OpenIDSessionRefresh] Failed to acquire shared refresh flight; refreshing directly', - error, - ); - return performIdpRefreshGrant(req, res, user, tokenPreference, identityContext); + logger.warn('[OpenIDSessionRefresh] Failed to acquire shared refresh flight', error); + throw new Error('OpenID refresh coordination is temporarily unavailable', { cause: error }); } if (!flight.acquired) { @@ -618,44 +682,50 @@ async function performIdpRefresh(req, res, user, tokenPreference, identityContex return resolvedTokens; } - logger.warn('[OpenIDSessionRefresh] Shared refresh flight unavailable; refreshing directly', { + logger.warn('[OpenIDSessionRefresh] Shared refresh flight remained unresolved', { key: hashKeyForLogs(key), }); - return performIdpRefreshGrant(req, res, user, tokenPreference, identityContext); + throw new Error('OpenID refresh coordination is temporarily unavailable'); } - try { - const resolvedTokens = await performIdpRefreshGrant( - req, - res, - user, - tokenPreference, - identityContext, - ); - try { - await completeOpenIDRefreshFlight({ - key, - ownerId: flight.ownerId, - tokens: resolvedTokens, - }); - } catch (flightError) { - logger.warn('[OpenIDSessionRefresh] Failed to complete shared refresh flight', { - key: hashKeyForLogs(key), - error: flightError?.message, - }); - } - return resolvedTokens; - } catch (error) { - try { - await failOpenIDRefreshFlight({ key, ownerId: flight.ownerId, error }); - } catch (flightError) { - logger.warn('[OpenIDSessionRefresh] Failed to mark shared refresh flight failed', { - key: hashKeyForLogs(key), - error: flightError?.message, - }); - } - throw error; - } + return withOpenIDRefreshFlightLease({ + key, + ownerId: flight.ownerId, + operation: async () => { + try { + const resolvedTokens = await performIdpRefreshGrant( + req, + res, + user, + tokenPreference, + identityContext, + ); + try { + await completeOpenIDRefreshFlight({ + key, + ownerId: flight.ownerId, + tokens: resolvedTokens, + }); + } catch (flightError) { + logger.warn('[OpenIDSessionRefresh] Failed to complete shared refresh flight', { + key: hashKeyForLogs(key), + error: flightError?.message, + }); + } + return resolvedTokens; + } catch (error) { + try { + await failOpenIDRefreshFlight({ key, ownerId: flight.ownerId, error }); + } catch (flightError) { + logger.warn('[OpenIDSessionRefresh] Failed to mark shared refresh flight failed', { + key: hashKeyForLogs(key), + error: flightError?.message, + }); + } + throw error; + } + }, + }); } /** @@ -743,7 +813,10 @@ async function refreshOrReuseSession(req, res, user, tokenPreference, identityCo * returned `expires_at`. OBO callers pass 'access_token'. */ async function refreshOpenIDSession(req, res, user, tokenPreference, identityContext) { - assertOpenIDSessionIdentityMatch(req, user, identityContext); + const identityBinding = assertOpenIDSessionIdentityMatch(req, user, identityContext); + if (identityBinding) { + await identityBinding; + } const key = getSingleFlightKey(req, user, identityContext); if (!key) { return refreshOrReuseSession(req, res, user, tokenPreference, identityContext); diff --git a/api/server/services/OpenIDSessionRefresh.spec.js b/api/server/services/OpenIDSessionRefresh.spec.js index 14195e8f6f..8ee5d80669 100644 --- a/api/server/services/OpenIDSessionRefresh.spec.js +++ b/api/server/services/OpenIDSessionRefresh.spec.js @@ -128,9 +128,11 @@ jest.mock('./OpenIDRefreshFlight', () => ({ createOpenIDRefreshFlightKey: jest.fn(), failOpenIDRefreshFlight: jest.fn(), waitForOpenIDRefreshFlight: jest.fn(), + withOpenIDRefreshFlightLease: jest.fn(({ operation }) => operation()), })); const jwt = require('jsonwebtoken'); +const crypto = require('node:crypto'); const openIdClient = require('openid-client'); const { isEnabled, @@ -148,6 +150,7 @@ const { createOpenIDRefreshFlightKey, failOpenIDRefreshFlight, waitForOpenIDRefreshFlight, + withOpenIDRefreshFlightLease, } = require('./OpenIDRefreshFlight'); const { createOpenIDSessionTokenProvider, @@ -210,6 +213,7 @@ describe('OpenIDSessionRefresh', () => { completeOpenIDRefreshFlight.mockResolvedValue({}); failOpenIDRefreshFlight.mockResolvedValue({}); waitForOpenIDRefreshFlight.mockResolvedValue(null); + withOpenIDRefreshFlightLease.mockImplementation(({ operation }) => operation()); }); describe('createOpenIDSessionTokenProvider closure no-op cases', () => { @@ -303,7 +307,7 @@ describe('OpenIDSessionRefresh', () => { }); }); - it('rejects session tokens that are missing identity metadata', async () => { + it('rejects legacy session tokens without a verifiable signed marker', async () => { const farFutureExp = Math.floor(Date.now() / 1000) + 600; const sessionTokens = { accessToken: makeJwt(farFutureExp), @@ -318,6 +322,45 @@ describe('OpenIDSessionRefresh', () => { expect(openIdClient.refreshTokenGrant).not.toHaveBeenCalled(); }); + it('binds a verified legacy session during rolling upgrades before token reuse', async () => { + const farFutureExp = Math.floor(Date.now() / 1000) + 600; + const refreshToken = 'rt-legacy'; + const sessionTokens = { + accessToken: makeJwt(farFutureExp), + idToken: makeJwt(farFutureExp), + refreshToken, + }; + const req = buildReq(sessionTokens, 'session-legacy', { bindIdentity: false }); + const previousSecret = process.env.JWT_REFRESH_SECRET; + process.env.JWT_REFRESH_SECRET = SECRET; + const refreshTokenHash = crypto.createHash('sha256').update(refreshToken).digest('base64url'); + const marker = jwt.sign({ id: 'local-id-1', refreshTokenHash }, SECRET); + req.headers = { + cookie: `refreshToken=${refreshToken}; openid_user_id=${marker}`, + }; + + try { + await expect( + refreshOpenIDSession(req, undefined, makeOpenIdUser(), 'access_token'), + ).resolves.toEqual( + expect.objectContaining({ + access_token: sessionTokens.accessToken, + refresh_token: refreshToken, + }), + ); + } finally { + if (previousSecret == null) { + delete process.env.JWT_REFRESH_SECRET; + } else { + process.env.JWT_REFRESH_SECRET = previousSecret; + } + } + + expect(req.session.openidTokens).toEqual(expect.objectContaining(DEFAULT_SESSION_IDENTITY)); + expect(req.session.save).toHaveBeenCalledTimes(1); + expect(openIdClient.refreshTokenGrant).not.toHaveBeenCalled(); + }); + it('rejects session tokens bound to a different OpenID identity', async () => { const farFutureExp = Math.floor(Date.now() / 1000) + 600; const sessionTokens = { @@ -1290,6 +1333,54 @@ describe('OpenIDSessionRefresh', () => { expires_at: expect.any(Number), }), }); + expect(withOpenIDRefreshFlightLease).toHaveBeenCalledWith({ + key: 'flight:session-cross-worker:rt-cross-worker', + ownerId: 'owner-leader', + operation: expect.any(Function), + }); + }); + + it('fails closed when a shared flight times out instead of issuing a duplicate grant', async () => { + const expiredExp = Math.floor(Date.now() / 1000) - 60; + const req = buildReq( + { + accessToken: makeJwt(expiredExp), + idToken: makeJwt(expiredExp), + refreshToken: 'rt-cross-worker', + }, + 'session-cross-worker', + ); + acquireOpenIDRefreshFlight.mockResolvedValueOnce({ + acquired: false, + ownerId: 'owner-joiner', + }); + waitForOpenIDRefreshFlight.mockResolvedValueOnce(null); + + await expect( + refreshOpenIDSession(req, undefined, makeOpenIdUser(), 'access_token'), + ).rejects.toThrow('OpenID refresh coordination is temporarily unavailable'); + + expect(openIdClient.refreshTokenGrant).not.toHaveBeenCalled(); + expect(withOpenIDRefreshFlightLease).not.toHaveBeenCalled(); + }); + + it('fails closed when Mongo flight acquisition is unavailable', async () => { + const expiredExp = Math.floor(Date.now() / 1000) - 60; + const req = buildReq( + { + accessToken: makeJwt(expiredExp), + idToken: makeJwt(expiredExp), + refreshToken: 'rt-coordination-down', + }, + 'session-coordination-down', + ); + acquireOpenIDRefreshFlight.mockRejectedValueOnce(new Error('mongo unavailable')); + + await expect( + refreshOpenIDSession(req, undefined, makeOpenIdUser(), 'access_token'), + ).rejects.toThrow('OpenID refresh coordination is temporarily unavailable'); + + expect(openIdClient.refreshTokenGrant).not.toHaveBeenCalled(); }); }); diff --git a/api/server/services/__tests__/MCP.spec.js b/api/server/services/__tests__/MCP.spec.js index 0fae7a3058..e2b8ebca36 100644 --- a/api/server/services/__tests__/MCP.spec.js +++ b/api/server/services/__tests__/MCP.spec.js @@ -2,6 +2,8 @@ const mockRegistry = { ensureConfigServers: jest.fn(), getAllServerConfigs: jest.fn(), }; +const mockUpstreamTokenProvider = jest.fn().mockResolvedValue(null); +const mockCreateOpenIDSessionTokenProvider = jest.fn(() => mockUpstreamTokenProvider); jest.mock('~/config', () => ({ getMCPServersRegistry: jest.fn(() => mockRegistry), @@ -34,6 +36,12 @@ jest.mock('@librechat/api', () => ({ GenerationJobManager: jest.fn(), buildOAuthToolCallName: jest.fn((name) => name), getUserMCPAuthMap: jest.fn(), + createAuthIdentityContext: ({ user, tenantId }) => ({ + appUserId: user?._id?.toString?.() ?? user?.id, + openidSubject: user?.openidId, + tenantId: tenantId ?? user?.tenantId, + openidIssuer: user?.openidIssuer, + }), /** Mirrors the real resolver so these tests still exercise the wrapper's own * plumbing - loading the request config and degrading on failure - rather than * the resolution logic, which is unit-tested in packages/api. Like the real @@ -66,6 +74,9 @@ jest.mock('~/server/services/OboTokenService', () => ({ jest.mock('~/server/services/OboPolicyService', () => ({ createOboTrustChecker: jest.fn(() => async () => true), })); +jest.mock('~/server/services/OpenIDSessionRefresh', () => ({ + createOpenIDSessionTokenProvider: (...args) => mockCreateOpenIDSessionTokenProvider(...args), +})); jest.mock('~/server/services/Tools/mcp', () => ({ reinitMCPServer: jest.fn(), })); @@ -159,10 +170,11 @@ describe('getAssistantToolDefinitions', () => { const getServerToolFunctionsSnapshot = jest.fn().mockResolvedValue({ tools: null }); require('~/config').getMCPManager.mockReturnValue({ getServerToolFunctionsSnapshot }); const userMCPAuthMap = { 'mcp_app-server': { API_KEY: 'saved' } }; + const res = { cookie: jest.fn() }; getUserMCPAuthMap.mockResolvedValue(userMCPAuthMap); reinitMCPServer.mockResolvedValue({ availableTools: { [toolKey]: mcpDefinition } }); - await expect(getAssistantToolDefinitions({ req, tools: [toolKey] })).resolves.toEqual({ + await expect(getAssistantToolDefinitions({ req, res, tools: [toolKey] })).resolves.toEqual({ toolDefinitions: { [toolKey]: mcpDefinition }, accessibleServerNames: ['app-server'], }); @@ -171,6 +183,25 @@ describe('getAssistantToolDefinitions', () => { serverName: 'app-server', serverConfig, userMCPAuthMap, + upstreamTokenProvider: mockUpstreamTokenProvider, + oboIdentityContext: { + appUserId: 'u1', + openidSubject: undefined, + tenantId: 'tenant-1', + openidIssuer: undefined, + }, + }); + expect(mockCreateOpenIDSessionTokenProvider).toHaveBeenCalledWith({ + req, + res, + user: req.user, + identityContext: { + appUserId: 'u1', + openidSubject: undefined, + tenantId: 'tenant-1', + openidIssuer: undefined, + }, + tokenPreference: 'access_token', }); expect(getUserMCPAuthMap).toHaveBeenCalledWith({ userId: 'u1', diff --git a/packages/data-schemas/src/methods/openidRefreshFlight.spec.ts b/packages/data-schemas/src/methods/openidRefreshFlight.spec.ts index a9842796dd..d3b4060eea 100644 --- a/packages/data-schemas/src/methods/openidRefreshFlight.spec.ts +++ b/packages/data-schemas/src/methods/openidRefreshFlight.spec.ts @@ -118,6 +118,38 @@ describe('OpenIDRefreshFlight Methods', () => { expect(reclaimed.flight?.errorMessage).toBeUndefined(); }); + it('renews a lease only for the owning pending worker', async () => { + await methods.acquireOpenIDRefreshFlight({ + key: 'flight-key', + ownerId: 'owner-1', + lockExpiresAt: new Date(Date.now() + 30000), + expiresAt: new Date(Date.now() + 60000), + }); + + const nextLockExpiry = new Date(Date.now() + 45000); + const nextFlightExpiry = new Date(Date.now() + 90000); + await expect( + methods.renewOpenIDRefreshFlight({ + key: 'flight-key', + ownerId: 'owner-2', + lockExpiresAt: nextLockExpiry, + expiresAt: nextFlightExpiry, + }), + ).resolves.toBeNull(); + + const renewed = await methods.renewOpenIDRefreshFlight({ + key: 'flight-key', + ownerId: 'owner-1', + lockExpiresAt: nextLockExpiry, + expiresAt: nextFlightExpiry, + }); + + expect(renewed?.ownerId).toBe('owner-1'); + expect(renewed?.status).toBe('pending'); + expect(renewed?.lockExpiresAt.getTime()).toBe(nextLockExpiry.getTime()); + expect(renewed?.expiresAt.getTime()).toBe(nextFlightExpiry.getTime()); + }); + it('completes a flight only for the owning pending worker', async () => { await methods.acquireOpenIDRefreshFlight({ key: 'flight-key', diff --git a/packages/data-schemas/src/methods/openidRefreshFlight.ts b/packages/data-schemas/src/methods/openidRefreshFlight.ts index 7ddfff6ae9..5781cae534 100644 --- a/packages/data-schemas/src/methods/openidRefreshFlight.ts +++ b/packages/data-schemas/src/methods/openidRefreshFlight.ts @@ -3,6 +3,7 @@ import type { IOpenIDRefreshFlight, OpenIDRefreshFlightCreateData, OpenIDRefreshFlightCompleteData, + OpenIDRefreshFlightRenewData, OpenIDRefreshFlightFailData, OpenIDRefreshFlightQuery, OpenIDRefreshFlightAcquireResult, @@ -30,6 +31,9 @@ export function createOpenIDRefreshFlightMethods(mongoose: typeof import('mongoo completeOpenIDRefreshFlight: ( data: OpenIDRefreshFlightCompleteData, ) => Promise; + renewOpenIDRefreshFlight: ( + data: OpenIDRefreshFlightRenewData, + ) => Promise; failOpenIDRefreshFlight: ( data: OpenIDRefreshFlightFailData, ) => Promise; @@ -147,6 +151,33 @@ export function createOpenIDRefreshFlightMethods(mongoose: typeof import('mongoo } } + async function renewOpenIDRefreshFlight( + data: OpenIDRefreshFlightRenewData, + ): Promise { + try { + const OpenIDRefreshFlight = mongoose.models + .OpenIDRefreshFlight as Model; + return await OpenIDRefreshFlight.findOneAndUpdate( + { + key: data.key, + ownerId: data.ownerId, + status: 'pending', + }, + { + $set: { + lockExpiresAt: data.lockExpiresAt, + expiresAt: data.expiresAt, + updatedAt: new Date(), + }, + }, + { new: true }, + ).lean(); + } catch (error) { + logger.debug('[renewOpenIDRefreshFlight] Error renewing flight:', error); + throw error; + } + } + async function failOpenIDRefreshFlight( data: OpenIDRefreshFlightFailData, ): Promise { @@ -196,6 +227,7 @@ export function createOpenIDRefreshFlightMethods(mongoose: typeof import('mongoo return { acquireOpenIDRefreshFlight, + renewOpenIDRefreshFlight, completeOpenIDRefreshFlight, failOpenIDRefreshFlight, findOpenIDRefreshFlight, diff --git a/packages/data-schemas/src/methods/refreshTokenBridge.spec.ts b/packages/data-schemas/src/methods/refreshTokenBridge.spec.ts index a65a8c4b97..3ba280168c 100644 --- a/packages/data-schemas/src/methods/refreshTokenBridge.spec.ts +++ b/packages/data-schemas/src/methods/refreshTokenBridge.spec.ts @@ -32,25 +32,12 @@ beforeEach(async () => { }); describe('RefreshTokenBridge Methods', () => { - it('keeps the lookup indexes aligned with the data-layer query shape', () => { - const indexKeys = mongoose.models.RefreshTokenBridge.schema.indexes().map(([key]) => key); - - expect(indexKeys).toContainEqual({ oldRefreshTokenHash: 1, userId: 1, tenantId: 1 }); - expect(indexKeys).not.toContainEqual({ - oldRefreshTokenHash: 1, - userId: 1, - tenantId: 1, - openidIssuer: 1, - }); - }); - - /** The TTL index is the only thing that deletes a bridge, and a bridge holds an encrypted - * refresh token — so a deployment with `MONGO_AUTO_INDEX=false` must not silently keep them. */ - it('creates the lookup and TTL indexes before the first write', async () => { + it('creates uniqueness and TTL indexes before the first write', async () => { await methods.upsertRefreshTokenBridge({ oldRefreshTokenHash: 'old-hash', encryptedNewRefreshToken: 'encrypted-new', userId: 'user-1', + tenantId: 'tenant-1', expiresAt: new Date(Date.now() + 60000), }); @@ -66,6 +53,18 @@ describe('RefreshTokenBridge Methods', () => { ); }); + it('keeps the lookup indexes aligned with the data-layer query shape', () => { + const indexKeys = mongoose.models.RefreshTokenBridge.schema.indexes().map(([key]) => key); + + expect(indexKeys).toContainEqual({ oldRefreshTokenHash: 1, userId: 1, tenantId: 1 }); + expect(indexKeys).not.toContainEqual({ + oldRefreshTokenHash: 1, + userId: 1, + tenantId: 1, + openidIssuer: 1, + }); + }); + it('upserts and finds a bridge by old token hash, user, and tenant', async () => { await methods.upsertRefreshTokenBridge({ oldRefreshTokenHash: 'old-hash', diff --git a/packages/data-schemas/src/methods/refreshTokenBridge.ts b/packages/data-schemas/src/methods/refreshTokenBridge.ts index 6ae8cc7c45..29cd08be03 100644 --- a/packages/data-schemas/src/methods/refreshTokenBridge.ts +++ b/packages/data-schemas/src/methods/refreshTokenBridge.ts @@ -25,11 +25,11 @@ export function createRefreshTokenBridgeMethods(mongoose: typeof import('mongoos ) => Promise; findRefreshTokenBridge: (query: RefreshTokenBridgeQuery) => Promise; } { + let indexesPromise: Promise | null = null; + const getRefreshTokenBridgeModel = () => mongoose.models.RefreshTokenBridge as Model; - let indexesPromise: Promise | null = null; - /** * A bridge holds an encrypted refresh token, and the TTL index is the only thing that ever * deletes one. `MONGO_AUTO_INDEX=false` is a supported deployment setting, and under it Mongoose @@ -51,8 +51,8 @@ export function createRefreshTokenBridgeMethods(mongoose: typeof import('mongoos async function upsertRefreshTokenBridge( bridgeData: RefreshTokenBridgeCreateData, ): Promise { - await ensureIndexes(); try { + await ensureIndexes(); const RefreshTokenBridge = getRefreshTokenBridgeModel(); const filter = bridgeFilter(bridgeData); const update: UpdateQuery = { diff --git a/packages/data-schemas/src/types/openidRefreshFlight.ts b/packages/data-schemas/src/types/openidRefreshFlight.ts index 1b80a5b130..b6b396a221 100644 --- a/packages/data-schemas/src/types/openidRefreshFlight.ts +++ b/packages/data-schemas/src/types/openidRefreshFlight.ts @@ -28,6 +28,13 @@ export interface OpenIDRefreshFlightCompleteData { expiresAt: Date; } +export interface OpenIDRefreshFlightRenewData { + key: string; + ownerId: string; + lockExpiresAt: Date; + expiresAt: Date; +} + export interface OpenIDRefreshFlightFailData { key: string; ownerId: string;