diff --git a/.env.example b/.env.example index 42b88f61bcf..b310ffe2cce 100644 --- a/.env.example +++ b/.env.example @@ -987,6 +987,12 @@ OPENID_REUSE_TOKENS= #is not rotated/revoked out from under downstream consumers (e.g. MCP servers that introspect the bearer). #When OPENID_REUSE_TOKENS=true, the OpenID session cookie maxAge is extended to at least this value. OPENID_REUSE_MAX_SESSION_AGE_MS= +# Discovery attempts during startup (0-100, default 1). Set to 0 to use background retries only. +# librechat.yaml `registration.openidDiscovery.startupAttempts` takes precedence. +OPENID_DISCOVERY_RETRY_ATTEMPTS= +# Delay in milliseconds between startup and background discovery retries (100-3600000, default 5000). +# librechat.yaml `registration.openidDiscovery.retryDelayMs` takes precedence. +OPENID_DISCOVERY_RETRY_DELAY_MS= #Short recovery window for a rotated OpenID refresh token while LibreChat publishes the refreshed session. Default 60000 ms (1 min). #Accepts arithmetic expressions. Increase only when slow session persistence or cross-replica publication needs more time. OPENID_REFRESH_BRIDGE_GRACE_MS= diff --git a/.github/workflows/static-checks.yml b/.github/workflows/static-checks.yml index 5782dab05f7..fe3c49e9ab6 100644 --- a/.github/workflows/static-checks.yml +++ b/.github/workflows/static-checks.yml @@ -131,7 +131,11 @@ jobs: # Run ESLint # --no-warn-ignored: changed files under config-ignored paths # (e.g. packages/data-schemas/misc/**) must not fail --max-warnings=0 - npx eslint --no-error-on-unmatched-pattern \ + # Invoke the installed binary, not `npx`: npm exec joins the whole + # command into one shell string, and Linux rejects a single argv + # string over 128 KiB (MAX_ARG_STRLEN) — past ~2,200 changed files + # npx dies with exit 249 and no output. + node_modules/.bin/eslint --no-error-on-unmatched-pattern \ --config eslint.config.mjs \ --no-warn-ignored \ --max-warnings=0 \ @@ -162,7 +166,7 @@ jobs: # `prettier --check` exits non-zero if any file would be reformatted. # Suggest the local fix in the failure message so contributors aren't # left guessing how to resolve. - if ! npx prettier --check --no-error-on-unmatched-pattern -- "${CHANGED_FILES[@]}"; then + if ! node_modules/.bin/prettier --check --no-error-on-unmatched-pattern -- "${CHANGED_FILES[@]}"; then echo "" echo "::error::Prettier formatting drift detected. Fix locally with:" echo "::error:: npx prettier --write " @@ -235,7 +239,7 @@ jobs: if: always() && steps.paths.outputs.eslint_config == 'true' continue-on-error: true run: | - npx eslint --config eslint.config.mjs \ + node_modules/.bin/eslint --config eslint.config.mjs \ api/server/index.js client/src/main.jsx packages/api/src/index.ts - name: Restore data-provider build cache @@ -808,8 +812,16 @@ jobs: run: | run_sweep() { set +e + # Two rules are switched off for the sweep only: `prettier/prettier` + # reformats every file (formatting drift is caught per changed file + # by the Static checks job and is not a config regression), and + # `import/no-cycle` walks the whole import graph from every file + # while config/circular-deps.mjs already owns cycle detection. + # Together they were ~85% of a full-tree lint. timeout -k 15 "$ESLINT_SWEEP_BUDGET_SECONDS" \ - npx eslint --config "$1" api client packages -f json -o "$2" + node_modules/.bin/eslint --config "$1" api client packages -f json -o "$2" \ + --rule 'prettier/prettier: off' \ + --rule 'import/no-cycle: off' local status=$? set -e # 124 = timeout sent TERM; 137 = it escalated to KILL. @@ -837,6 +849,13 @@ jobs: # flat-config files/ignores patterns and plugin imports resolve # relative to the config's own directory, so a temp-dir copy would # scope to nothing and the comparison would pass vacuously. + # An unchanged config cannot regress: the head sweep above already + # proved this workflow still runs it, so a second identical sweep + # would only double the job's runtime. + if git diff --quiet "$BASE_SHA" HEAD -- eslint.config.mjs; then + echo "::notice title=ESLint sweep::eslint.config.mjs is unchanged from base; skipping the regression comparison." + exit 0 + fi trap 'rm -f eslint.config.base.mjs' EXIT if ! git show "$BASE_SHA:eslint.config.mjs" > eslint.config.base.mjs 2>/dev/null; then echo "::notice title=ESLint sweep::No eslint.config.mjs at base ref; skipping regression comparison." diff --git a/api/app/clients/BaseClient.js b/api/app/clients/BaseClient.js index 950efff0026..248f336a223 100644 --- a/api/app/clients/BaseClient.js +++ b/api/app/clients/BaseClient.js @@ -20,6 +20,7 @@ const { collectModelBoundHistoricalFileIdState, projectModelBoundSourceFiles, isModelBoundAttachmentFile, + withBalanceReservations, } = require('@librechat/api'); const { Constants, @@ -732,6 +733,18 @@ class BaseClient { } async sendMessage(message, opts = {}) { + return withBalanceReservations((balanceReservations) => + this.sendReservedMessage(message, opts, balanceReservations), + ); + } + + /** + * @param {string} message + * @param {Record} opts + * @param {BalanceReservations} balanceReservations - Holds the balance reservation admitting + * this message; released once its usage is recorded, and by `sendMessage` on any other exit. + */ + async sendReservedMessage(message, opts, balanceReservations) { const appConfig = this.options.req?.config; /** @type {Promise} */ let userMessagePromise; @@ -974,7 +987,7 @@ class BaseClient { balanceConfig?.enabled && supportsBalanceCheck[this.options.endpointType ?? this.options.endpoint] ) { - await checkBalance( + const balanceAdmission = checkBalance( { req: this.options.req, res: this.options.res, @@ -990,12 +1003,13 @@ class BaseClient { { logViolation, getMultiplier: db.getMultiplier, - findBalanceByUser: db.findBalanceByUser, - createAutoRefillTransaction: db.createAutoRefillTransaction, + reserveBalance: db.reserveBalance, + renewBalanceReservation: db.renewBalanceReservation, + releaseBalanceReservation: db.releaseBalanceReservation, balanceConfig, - upsertBalanceFields: db.upsertBalanceFields, }, ); + await balanceReservations.track(balanceAdmission); } completionResult = await this.sendCompletion(payload, opts); @@ -1155,6 +1169,7 @@ class BaseClient { completionTokens, }); } + await balanceReservations.release(); if (userMessagePromise) { await userMessagePromise; diff --git a/api/app/clients/specs/BaseClient.test.js b/api/app/clients/specs/BaseClient.test.js index 8f2bbbafeef..3aa013cc8e8 100644 --- a/api/app/clients/specs/BaseClient.test.js +++ b/api/app/clients/specs/BaseClient.test.js @@ -48,9 +48,22 @@ jest.mock('~/models', () => ({ deleteFiles: jest.fn(), getFiles: jest.fn(), updateFileUsage: jest.fn(), + getMultiplier: jest.fn(), + reserveBalance: jest.fn(), + renewBalanceReservation: jest.fn(), + releaseBalanceReservation: jest.fn(), })); -const { getConvo, getFiles, getMessages, saveConvo, saveMessage } = require('~/models'); +const { + releaseBalanceReservation, + reserveBalance, + getMultiplier, + saveMessage, + getMessages, + saveConvo, + getFiles, + getConvo, +} = require('~/models'); jest.mock('@librechat/agents', () => { const actual = jest.requireActual('@librechat/agents'); @@ -2225,6 +2238,91 @@ describe('BaseClient', () => { }); }); + describe('balance reservation lifecycle', () => { + let priorEndpoint; + let priorEndpointType; + let events; + + beforeEach(() => { + priorEndpoint = TestClient.options.endpoint; + priorEndpointType = TestClient.options.endpointType; + TestClient.options.endpoint = EModelEndpoint.openAI; + delete TestClient.options.endpointType; + TestClient.options.req = { config: { balance: { enabled: true } } }; + + events = []; + getMultiplier.mockReturnValue(1); + reserveBalance.mockImplementation(async () => { + events.push('reserve'); + return { reserved: true, balance: 1000 }; + }); + releaseBalanceReservation.mockImplementation(async () => { + events.push('release'); + }); + TestClient.sendCompletion.mockImplementation(async () => { + events.push('completion'); + return { completion: 'Mock response text', metadata: undefined }; + }); + TestClient.getTokenCountForResponse = jest.fn().mockReturnValue(50); + TestClient.recordTokenUsage = jest.fn(async () => { + events.push('usage'); + }); + TestClient.buildMessages.mockReturnValue({ + prompt: [], + tokenCountMap: { res: 50 }, + }); + }); + + afterEach(() => { + delete TestClient.options.req; + TestClient.options.endpoint = priorEndpoint; + TestClient.options.endpointType = priorEndpointType; + }); + + test('releases the reservation once the response usage is recorded, before persistence', async () => { + const beforeResponsePersistence = jest.fn(async () => { + events.push('persist'); + return true; + }); + + await TestClient.sendMessage('Hello', { beforeResponsePersistence }); + + expect(events).toEqual(['reserve', 'completion', 'usage', 'release', 'persist']); + const [{ reservationId, amount }] = reserveBalance.mock.calls[0]; + expect(releaseBalanceReservation).toHaveBeenCalledTimes(1); + expect(releaseBalanceReservation).toHaveBeenCalledWith({ + user: TestClient.user, + reservationId, + amount, + }); + }); + + test('releases the reservation when the completion fails', async () => { + TestClient.sendCompletion.mockRejectedValue(new Error('provider unavailable')); + + await expect(TestClient.sendMessage('Hello', {})).rejects.toThrow('provider unavailable'); + + expect(events).toEqual(['reserve', 'release']); + }); + + test('releases the reservation when work after the completion fails', async () => { + TestClient.recordTokenUsage.mockRejectedValue(new Error('usage write failed')); + + await expect(TestClient.sendMessage('Hello', {})).rejects.toThrow('usage write failed'); + + expect(events).toEqual(['reserve', 'completion', 'release']); + }); + + test('takes no reservation when the balance check refuses the request', async () => { + reserveBalance.mockResolvedValue({ reserved: false, balance: 0 }); + + await expect(TestClient.sendMessage('Hello', {})).rejects.toThrow(); + + expect(TestClient.sendCompletion).not.toHaveBeenCalled(); + expect(releaseBalanceReservation).not.toHaveBeenCalled(); + }); + }); + describe('getMessagesWithinTokenLimit with instructions', () => { test('should always include instructions when present', async () => { TestClient.maxContextTokens = 50; diff --git a/api/jest.config.js b/api/jest.config.js index 4ba0853094f..d0cabf39c5a 100644 --- a/api/jest.config.js +++ b/api/jest.config.js @@ -1,3 +1,5 @@ +const { maxWorkers } = require('../config/jest.workers.cjs'); + const esModules = [ 'openid-client', 'oauth4webapi', @@ -21,7 +23,7 @@ module.exports = { clearMocks: true, roots: [''], coverageDirectory: 'coverage', - maxWorkers: '50%', + maxWorkers, testTimeout: 30000, // 30 seconds timeout for all tests setupFiles: ['./test/jestSetup.js', './test/__mocks__/logger.js'], moduleNameMapper: { diff --git a/api/server/controllers/PermissionsController.js b/api/server/controllers/PermissionsController.js index 6bfe9028f97..fabc5610936 100644 --- a/api/server/controllers/PermissionsController.js +++ b/api/server/controllers/PermissionsController.js @@ -7,6 +7,7 @@ const { logger, getTenantId, SYSTEM_TENANT_ID } = require('@librechat/data-schem const { ResourceType, PrincipalType, PermissionBits } = require('librechat-data-provider'); const { enrichRemoteAgentPrincipals, + createPrincipalSearch, backfillRemoteAgentPermissions, auditInsightsPermissionChanges, getInsightsPrincipalState, @@ -471,120 +472,13 @@ const getUserEffectivePermissions = async (req, res) => { * Supports hybrid local database + Entra ID search when configured * @route GET /api/permissions/search-principals */ -const searchPrincipals = async (req, res) => { - try { - const { q: rawQuery, limit = 20, types } = req.query; - - if (typeof rawQuery !== 'string' || rawQuery.trim().length === 0) { - return res.status(400).json({ - error: 'Query parameter "q" is required and must not be empty', - }); - } - - const query = rawQuery.trim(); - - if (query.length < 2) { - return res.status(400).json({ - error: 'Query must be at least 2 characters long', - }); - } - - const searchLimit = Math.min(Math.max(1, parseInt(limit) || 10), 50); - - let typeFilters = null; - if (types) { - const typesArray = Array.isArray(types) ? types : types.split(','); - const validTypes = typesArray.filter((t) => - [PrincipalType.USER, PrincipalType.GROUP, PrincipalType.ROLE].includes(t), - ); - typeFilters = validTypes.length > 0 ? validTypes : null; - } - - const localResults = await db.searchPrincipals(query, searchLimit, typeFilters); - let allPrincipals = [...localResults]; - - const useEntraId = entraIdPrincipalFeatureEnabled(req.user); - - if (useEntraId && localResults.length < searchLimit) { - try { - let graphType = 'all'; - if (typeFilters && typeFilters.length === 1) { - const graphTypeMap = { - [PrincipalType.USER]: 'users', - [PrincipalType.GROUP]: 'groups', - }; - const mappedType = graphTypeMap[typeFilters[0]]; - if (mappedType) { - graphType = mappedType; - } - } - - const authHeader = req.headers.authorization; - const accessToken = - authHeader && authHeader.startsWith('Bearer ') ? authHeader.substring(7) : null; - - if (accessToken) { - const graphResults = await searchEntraIdPrincipals( - accessToken, - req.user.openidId, - query, - graphType, - searchLimit - localResults.length, - ); - - const localEmails = new Set( - localResults.map((p) => p.email?.toLowerCase()).filter(Boolean), - ); - const localGroupSourceIds = new Set( - localResults.map((p) => p.idOnTheSource).filter(Boolean), - ); - - for (const principal of graphResults) { - const isDuplicateByEmail = - principal.email && localEmails.has(principal.email.toLowerCase()); - const isDuplicateBySourceId = - principal.idOnTheSource && localGroupSourceIds.has(principal.idOnTheSource); - - if (!isDuplicateByEmail && !isDuplicateBySourceId) { - allPrincipals.push(principal); - } - } - } - } catch (graphError) { - logger.warn('Graph API search failed, falling back to local results:', graphError.message); - } - } - const scoredResults = allPrincipals.map((item) => ({ - ...item, - _searchScore: db.calculateRelevanceScore(item, query), - })); - - const finalResults = db - .sortPrincipalsByRelevance(scoredResults) - .slice(0, searchLimit) - .map((result) => { - const { _searchScore, ...resultWithoutScore } = result; - return resultWithoutScore; - }); - - res.status(200).json({ - query, - limit: searchLimit, - types: typeFilters, - results: finalResults, - count: finalResults.length, - sources: { - local: finalResults.filter((r) => r.source === 'local').length, - entra: finalResults.filter((r) => r.source === 'entra').length, - }, - }); - } catch (error) { - logger.error('Error searching principals:', error); - res.status(500).json({ - error: 'Failed to search principals', - }); - } -}; +const searchPrincipals = createPrincipalSearch({ + searchPrincipals: db.searchPrincipals, + calculateRelevanceScore: db.calculateRelevanceScore, + sortPrincipalsByRelevance: db.sortPrincipalsByRelevance, + entraIdPrincipalFeatureEnabled, + searchEntraIdPrincipals, +}); /** * Get user's effective permissions for all accessible resources of a type diff --git a/api/server/controllers/__tests__/PermissionsController.spec.js b/api/server/controllers/__tests__/PermissionsController.spec.js index 2a578b75e7e..0b83041a56e 100644 --- a/api/server/controllers/__tests__/PermissionsController.spec.js +++ b/api/server/controllers/__tests__/PermissionsController.spec.js @@ -9,8 +9,9 @@ jest.mock('@librechat/data-schemas', () => ({ SYSTEM_TENANT_ID: '__SYSTEM__', })); -const { AccessRoleIds, ResourceType, PrincipalType, SystemRoles } = +const { AccessRoleIds, ResourceType, PrincipalType, SystemRoles, PermissionTypes, Permissions } = jest.requireActual('librechat-data-provider'); +const { createPeoplePickerAccess } = jest.requireActual('@librechat/api'); jest.mock('librechat-data-provider', () => ({ ...jest.requireActual('librechat-data-provider'), @@ -100,68 +101,32 @@ describe('PermissionsController', () => { db.sortPrincipalsByRelevance.mockImplementation((results) => results); }); - it('rejects non-string query parameters', async () => { - const req = createMockReq({ - query: { q: ['alice'] }, - }); - const res = createMockRes(); - - await searchPrincipals(req, res); - - expect(res.status).toHaveBeenCalledWith(400); - expect(res.json).toHaveBeenCalledWith({ - error: 'Query parameter "q" is required and must not be empty', - }); - expect(db.searchPrincipals).not.toHaveBeenCalled(); - }); - - it('searches with the trimmed literal query', async () => { - db.searchPrincipals.mockResolvedValue([ - { - id: 'user-1', - type: PrincipalType.USER, - name: 'Regex [invalid User', - source: 'local', - }, - ]); - - const req = createMockReq({ - query: { q: ' [invalid ', limit: '5', types: PrincipalType.USER }, - }); - const res = createMockRes(); - - await searchPrincipals(req, res); - - expect(db.searchPrincipals).toHaveBeenCalledWith('[invalid', 5, [PrincipalType.USER]); - expect(db.calculateRelevanceScore).toHaveBeenCalledWith( - expect.objectContaining({ name: 'Regex [invalid User' }), - '[invalid', - ); - expect(res.status).toHaveBeenCalledWith(200); - expect(res.json).toHaveBeenCalledWith( - expect.objectContaining({ - query: '[invalid', - limit: 5, - count: 1, - }), - ); - }); - - it('does not expose internal error details on search failures', async () => { - db.searchPrincipals.mockRejectedValue(new Error('database failure with internal detail')); - - const req = createMockReq({ - query: { q: 'alice' }, - }); - const res = createMockRes(); - - await searchPrincipals(req, res); - - expect(res.status).toHaveBeenCalledWith(500); - expect(res.json).toHaveBeenCalledWith({ - error: 'Failed to search principals', - }); - }); + it.each([{ q: 'al', type: PrincipalType.GROUP }, { q: 'al', types: 'foobar' }, { q: 'al' }])( + 'searches only the types the people picker check resolved for %j', + async (query) => { + const checkAccess = createPeoplePickerAccess({ + getRoleByName: async () => ({ + permissions: { + [PermissionTypes.PEOPLE_PICKER]: { + [Permissions.VIEW_USERS]: false, + [Permissions.VIEW_GROUPS]: true, + [Permissions.VIEW_ROLES]: false, + }, + }, + }), + }); + const req = createMockReq({ query }); + const res = createMockRes(); + + await checkAccess(req, res, () => searchPrincipals(req, res)); + + expect(db.searchPrincipals).toHaveBeenCalledWith('al', 20, [PrincipalType.GROUP]); + expect(res.status).toHaveBeenCalledWith(200); + expect(res.json).toHaveBeenCalledWith( + expect.objectContaining({ types: [PrincipalType.GROUP] }), + ); + }, + ); }); describe('getResourcePermissions — principal details', () => { diff --git a/api/server/controllers/agents/__tests__/callbacks.spec.js b/api/server/controllers/agents/__tests__/callbacks.spec.js index cd9c1cd2ae8..318adaf8bc8 100644 --- a/api/server/controllers/agents/__tests__/callbacks.spec.js +++ b/api/server/controllers/agents/__tests__/callbacks.spec.js @@ -24,6 +24,7 @@ jest.mock('@librechat/api', () => ({ ), isCodeArtifactToolOutput: jest.requireActual('@librechat/api').isCodeArtifactToolOutput, isCodeSessionToolName: jest.requireActual('@librechat/api').isCodeSessionToolName, + collectToolCallIds: jest.requireActual('@librechat/api').collectToolCallIds, })); jest.mock('@librechat/data-schemas', () => ({ @@ -210,10 +211,17 @@ describe('resumable event generation fencing', () => { ); const contextUsageSink = { latest: null, count: 0, onSnapshot }; const usageEmitSink = [{ input_tokens: 10 }]; + /** The call this snapshot already accounts for: a later tool-limit stop counts + * the results of the calls missing from this set. */ + const contentParts = [ + { type: 'text', text: 'answering' }, + { type: 'tool_call', tool_call: { id: 'call_1', name: 'grep' } }, + ]; const data = { contextBudget: 1000, remainingContextTokens: 400 }; const handlers = getDefaultHandlers({ res: { write: jest.fn() }, aggregateContent: jest.fn(), + contentParts, toolEndCallback: jest.fn(), collectedUsage: [], streamId: 'conversation-1', @@ -230,7 +238,12 @@ describe('resumable event generation fencing', () => { }); await Promise.resolve(); - expect(contextUsageSink).toMatchObject({ latest: data, count: 1, latestUsageIndex: 1 }); + expect(contextUsageSink).toMatchObject({ + latest: data, + count: 1, + latestUsageIndex: 1, + latestToolCallIds: new Set(['call_1']), + }); expect(onSnapshot).toHaveBeenCalledTimes(1); expect(settled).toBe(false); releaseSnapshot(); diff --git a/api/server/controllers/agents/__tests__/client.contextMetadata.spec.js b/api/server/controllers/agents/__tests__/client.contextMetadata.spec.js index 6df38efb6c5..6fe18af8638 100644 --- a/api/server/controllers/agents/__tests__/client.contextMetadata.spec.js +++ b/api/server/controllers/agents/__tests__/client.contextMetadata.spec.js @@ -1,3 +1,12 @@ +/** The counting itself is exercised in packages/api (it needs a loaded encoding, + * which this environment cannot load); here the resolver stands in so the WIRING + * is pinned: which content the save path hands it, and where its figure lands. */ +const mockResolveRetainedToolTokens = jest.fn(); +jest.mock('@librechat/api', () => ({ + ...jest.requireActual('@librechat/api'), + resolveRetainedToolTokens: (...args) => mockResolveRetainedToolTokens(...args), +})); + const AgentClient = require('../client'); /** Minimal post-(maybe-)summary snapshot. baseUsed = maxContextTokens(1000) - @@ -34,18 +43,41 @@ const primaryFor = (runId, output_tokens) => ({ runId, }); -function buildMeta({ snap, latestUsageIndex, usageEvents }) { +const toolPart = (id, name, output) => ({ + type: 'tool_call', + tool_call: { id, name, args: '{"path":"a"}', output }, +}); + +function buildMeta({ + snap, + latestUsageIndex, + usageEvents, + stepLimitReached = false, + latestToolCallIds, + contentParts, + maxRetainedToolCountChars, +}) { const self = { collectedThoughtSignatures: null, usageEmitSink: usageEvents, + stepLimitReached, + contentParts, + getEncoding: () => 'o200k_base', + options: { + req: { config: { endpoints: { agents: { maxRetainedToolCountChars } } } }, + }, contextUsageSink: snap - ? { latest: snap, count: 1, latestUsageIndex } + ? { latest: snap, count: 1, latestUsageIndex, latestToolCallIds } : { latest: null, count: 0 }, }; return AgentClient.prototype.buildResponseMetadata.call(self); } describe('AgentClient.buildResponseMetadata — snapshot persistence + summary marker', () => { + beforeEach(() => { + mockResolveRetainedToolTokens.mockReset(); + }); + it('persists the snapshot when a primary usage follows it (normal turn)', () => { const meta = buildMeta({ snap: snapshot(0), latestUsageIndex: 0, usageEvents: [primary] }); expect(meta.contextUsage).toBeDefined(); @@ -136,4 +168,54 @@ describe('AgentClient.buildResponseMetadata — snapshot persistence + summary m /** run-1's own primary follows the snapshot → snapshot persisted with output 5. */ expect(meta.contextUsage.completedOutputTokens).toBe(5); }); + + /** A turn that stops at the tool-call limit keeps the results of the tools its + * final call ran. The snapshot describing that call precedes them and no further + * call is made, so the counted figure has to ride along or the client's gauge + * misses the retained result until the next turn. */ + it('hands the resolver this snapshot’s call boundary and persists its figure', () => { + mockResolveRetainedToolTokens.mockReturnValue(180); + const contentParts = [ + toolPart('call_1', 'grep', 'the result the snapshot already counts'), + toolPart('call_2', 'read_file', 'the retained result'), + ]; + const latestToolCallIds = new Set(['call_1']); + const meta = buildMeta({ + snap: snapshot(0), + latestUsageIndex: 0, + usageEvents: [primary], + stepLimitReached: true, + latestToolCallIds, + contentParts, + maxRetainedToolCountChars: 1_048_576, + }); + expect(mockResolveRetainedToolTokens).toHaveBeenCalledWith({ + stoppedAtToolLimit: true, + contentParts, + /** The calls the snapshot already saw; only the rest are retained. */ + priorToolCallIds: latestToolCallIds, + encoding: 'o200k_base', + /** The deployment's ceiling on the tokenization this costs. */ + maxCountChars: 1_048_576, + }); + expect(meta.contextUsage.retainedToolTokens).toBe(180); + }); + + it('reports a turn that did not stop at the tool-call limit as such', () => { + mockResolveRetainedToolTokens.mockReturnValue(undefined); + const meta = buildMeta({ + snap: snapshot(0), + latestUsageIndex: 0, + usageEvents: [primary], + latestToolCallIds: new Set(), + contentParts: [toolPart('call_1', 'read_file', 'a result the next call re-counted')], + }); + expect(mockResolveRetainedToolTokens).toHaveBeenCalledWith( + expect.objectContaining({ stoppedAtToolLimit: false }), + ); + /** Nothing to add, so the blob stays exactly as it was before this change. */ + expect(Object.prototype.hasOwnProperty.call(meta.contextUsage, 'retainedToolTokens')).toBe( + false, + ); + }); }); diff --git a/api/server/controllers/agents/__tests__/request.resumeMetadata.spec.js b/api/server/controllers/agents/__tests__/request.resumeMetadata.spec.js index a7594950bb7..d778f6b2e58 100644 --- a/api/server/controllers/agents/__tests__/request.resumeMetadata.spec.js +++ b/api/server/controllers/agents/__tests__/request.resumeMetadata.spec.js @@ -283,6 +283,7 @@ jest.mock('@librechat/api', () => ({ jest.requireActual('@librechat/api').shouldPersistCodeWorkspaceInitializationError, getSafeErrorMetadata: jest.requireActual('@librechat/api').getSafeErrorMetadata, getSafeErrorText: jest.requireActual('@librechat/api').getSafeErrorText, + resolveFailedTurnContent: jest.requireActual('@librechat/api').resolveFailedTurnContent, getFailedTurnTraceFields: (...args) => mockGetFailedTurnTraceFields(...args), GenerationJobManager: mockGenerationJobManager, getReferencedQuotes: jest.fn((quotes) => { diff --git a/api/server/controllers/agents/callbacks.js b/api/server/controllers/agents/callbacks.js index 6c295b6568a..4b27bde7fed 100644 --- a/api/server/controllers/agents/callbacks.js +++ b/api/server/controllers/agents/callbacks.js @@ -31,6 +31,7 @@ const { shouldSignalSandboxStart, getToolInputValidationDetails, captureSubagentIdentity, + collectToolCallIds, } = require('@librechat/api'); const { processFileCitations } = require('~/server/services/Files/Citations'); const { processCodeOutput, runPreviewFinalize } = require('~/server/services/Files/Code/process'); @@ -385,7 +386,9 @@ function feedSubagentAggregator(aggregator, event) { * @param {UsageCostDeps} [options.usageCost] - Pricing context for authoritative per-event cost. * @param {{ latest: TContextUsageEvent | null, count: number }} [options.contextUsageSink] - Mutable * holder for the latest visible context snapshot + a count of visible snapshots (model calls), - * used to persist the breakdown only when the final call emitted usage. + * used to persist the breakdown only when the final call emitted usage. Also records that + * snapshot's position in the usage stream and in `contentParts`, so the save path can tell + * which usage events and which content parts came after it. * @param {Array} [options.usageEmitSink] - Array collecting each emitted * `on_token_usage` payload (incl. cost) so the response's usage rollup can be persisted. * @param {(toolName: string, agentId?: string) => string | undefined} [options.resolveMcpServerName] @@ -851,6 +854,13 @@ function getDefaultHandlers({ contextUsageSink.latest = data; contextUsageSink.count = (contextUsageSink.count ?? 0) + 1; contextUsageSink.latestUsageIndex = usageEmitSink?.length ?? 0; + /** Which tool calls this snapshot already accounts for. A turn that + * stops at the tool-call limit counts the results of the calls missing + * from this set — the ones its own call produced, which no later + * snapshot describes. Ids, not a content index: completion reshapes the + * array (skill cards unshifted, hidden sequential output filtered), so + * an index recorded here would mean something else by save time. */ + contextUsageSink.latestToolCallIds = collectToolCallIds(contentParts); } /** Every agent's snapshot publishes the run's context meta, hidden * sequential agents included: their model calls latch tiers too, and a diff --git a/api/server/controllers/agents/client.js b/api/server/controllers/agents/client.js index e37693d986c..5b9b9552575 100644 --- a/api/server/controllers/agents/client.js +++ b/api/server/controllers/agents/client.js @@ -26,12 +26,15 @@ const { applyContextToAgent, isMemoryAgentEnabled, recordCollectedUsage, + resolveRunUsageContext, + recordFallbackTokenUsage, createDetachedSubagentUsageRecorder, sendEvent, computeUsageCostUSD, aggregateEmittedUsage, resolveAgentTokenConfig, buildPersistedContextUsage, + resolveRetainedToolTokens, computeSummaryUsedTokens, priorRunOutputTokens, createSubagentUsageSink, @@ -171,6 +174,8 @@ const { resolveToolRoleGrants, createTerminalRunErrorObserver, isAgentRunCancellation, + getSummaryPartText, + markCompactionOutcome, } = require('@librechat/api'); const { Run, @@ -331,41 +336,6 @@ function captureRunContextMeta(client) { }); } -/** Text of a summary content part; empty for anything else. */ -function getSummaryPartText(part) { - if (part?.type !== ContentTypes.SUMMARY || !Array.isArray(part.content)) { - return ''; - } - return part.content - .map((block) => (typeof block?.text === 'string' ? block.text : '')) - .join('') - .trim(); -} - -/** - * A compaction turn's response is its summary. The run emits no text, so a - * completion without a usable summary part means the summarizer produced - * nothing. A run that already recorded why (an error part, e.g. a skipped - * compaction) persists with that explanation; one that ended with neither - * fails as a typed error instead of persisting an empty assistant message. - * @param {Array} contentParts - */ -function markCompactionSummary(contentParts) { - const summary = contentParts.find( - (part) => part?.failed !== true && getSummaryPartText(part).length > 0, - ); - if (summary != null) { - summary.initiatedBy = 'user'; - return; - } - if (contentParts.some((part) => part?.type === ContentTypes.ERROR)) { - return; - } - throw Object.assign(new Error(JSON.stringify({ type: ErrorTypes.COMPACTION_FAILED })), { - code: 'COMPACTION_FAILED', - }); -} - function getLatestEventActorSummary(contentParts) { if (!Array.isArray(contentParts)) { return undefined; @@ -3568,7 +3538,9 @@ class AgentClient extends BaseClient { const completion = filterMalformedContentParts(this.contentParts); if (this.isCompactionTurn()) { - markCompactionSummary(completion); + markCompactionOutcome(completion, { + aborted: this.abortController?.signal?.aborted === true, + }); } const metadata = this.buildResponseMetadata(); return metadata ? { completion, metadata } : { completion }; @@ -3625,7 +3597,19 @@ class AgentClient extends BaseClient { event.runId === latestSnapshotRunId), ); if (latestSnapshot && hasPrimaryAfterSnapshot) { - metadata.contextUsage = buildPersistedContextUsage(latestSnapshot, usageEvents); + /** The counted tool results this turn keeps past that snapshot — only a + * tool-call-limit stop has any; see `resolveRetainedToolTokens`. */ + metadata.contextUsage = buildPersistedContextUsage(latestSnapshot, usageEvents, { + retainedToolTokens: resolveRetainedToolTokens({ + stoppedAtToolLimit: this.stepLimitReached === true, + contentParts: this.contentParts, + priorToolCallIds: this.contextUsageSink?.latestToolCallIds, + encoding: this.getEncoding(), + maxCountChars: + this.options?.req?.config?.endpoints?.[EModelEndpoint.agents] + ?.maxRetainedToolCountChars, + }), + }); } /** Lightweight summarization marker — persisted whenever this turn compacted * the context, INDEPENDENT of the snapshot guard above. When the client has @@ -5170,20 +5154,14 @@ class AgentClient extends BaseClient { this.artifactPromises.push(...attachments); } - /** Skip token spending if aborted - the abort handler (abortMiddleware.js) handles it - This prevents double-spending when user aborts via `/api/agents/chat/abort` */ - const wasAborted = abortController?.signal?.aborted; - if (!wasAborted) { - await this.recordCollectedUsage({ - context: 'message', - balance: balanceConfig, - transactions: transactionsConfig, - }); - } else { - logger.debug( - '[api/server/controllers/agents/client.js #chatCompletion] Skipping token spending - handled by abort middleware', - ); - } + /** The run owns its usage even when stopped: `/api/agents/chat/abort` + * only signals the abort, so nothing else records what was consumed. + * A stopped turn is labelled as such on its transactions. */ + await this.recordCollectedUsage({ + context: resolveRunUsageContext(abortController?.signal?.aborted === true), + balance: balanceConfig, + transactions: transactionsConfig, + }); } catch (err) { logger.error( '[api/server/controllers/agents/client.js #chatCompletion] Error in cleanup phase', @@ -5840,14 +5818,11 @@ class AgentClient extends BaseClient { } try { - const wasAborted = abortController?.signal?.aborted; - if (!wasAborted) { - await this.recordCollectedUsage({ - context: 'message', - balance: balanceConfig, - transactions: transactionsConfig, - }); - } + await this.recordCollectedUsage({ + context: resolveRunUsageContext(abortController?.signal?.aborted === true), + balance: balanceConfig, + transactions: transactionsConfig, + }); } catch (err) { logger.error( '[api/server/controllers/agents/client.js #resumeCompletion] Error in cleanup phase', @@ -6155,13 +6130,19 @@ class AgentClient extends BaseClient { transactions, promptTokens, completionTokens, - context = 'message', + context, }) { - try { - await db.spendTokens( - { + await recordFallbackTokenUsage( + { spendTokens: db.spendTokens }, + { + usage, + context, + collectedUsage: this.collectedUsage, + aborted: this.abortController?.signal?.aborted === true, + promptTokens, + completionTokens, + txMetadata: { model, - context, balance, transactions, messageId: this.responseMessageId, @@ -6169,35 +6150,8 @@ class AgentClient extends BaseClient { user: this.user ?? this.options.req.user?.id, endpointTokenConfig: this.options.endpointTokenConfig, }, - { promptTokens, completionTokens }, - ); - - if ( - usage && - typeof usage === 'object' && - 'reasoning_tokens' in usage && - typeof usage.reasoning_tokens === 'number' - ) { - await db.spendTokens( - { - model, - balance, - transactions, - context: 'reasoning', - messageId: this.responseMessageId, - conversationId: this.conversationId, - user: this.user ?? this.options.req.user?.id, - endpointTokenConfig: this.options.endpointTokenConfig, - }, - { completionTokens: usage.reasoning_tokens }, - ); - } - } catch (error) { - logger.error( - '[api/server/controllers/agents/client.js #recordTokenUsage] Error recording token usage', - getSafeErrorMetadata(error), - ); - } + }, + ); } /** Anthropic Claude models use a distinct BPE tokenizer; all others default to o200k_base. */ diff --git a/api/server/controllers/agents/client.test.js b/api/server/controllers/agents/client.test.js index 54b869ca885..8e63b061b1d 100644 --- a/api/server/controllers/agents/client.test.js +++ b/api/server/controllers/agents/client.test.js @@ -21,7 +21,7 @@ const mockStripActivityLabelParts = jest.fn((payload) => ); const { Providers } = require('@librechat/agents'); -const { Constants, ContentTypes, EModelEndpoint } = require('librechat-data-provider'); +const { Constants, ContentTypes, EModelEndpoint, ErrorTypes } = require('librechat-data-provider'); const { GenerationJobManager, createStreamServices, @@ -1512,6 +1512,56 @@ describe('AgentClient - interrupt discovery persistence', () => { expect(metaWhenResumed).toEqual(seed); }); + it('records collected usage as an abort when a resumed run is stopped', async () => { + jest.clearAllMocks(); + const streamId = 'conversation-resume-stopped'; + const job = await GenerationJobManager.createJob(streamId, 'user-123', streamId); + const abortController = new AbortController(); + mockCreateRun.mockImplementationOnce(async () => ({ + Graph: null, + resume: jest.fn(async () => { + abortController.abort(); + }), + processStream: jest.fn().mockResolvedValue(), + getCalibrationRatio: jest.fn(() => 0), + getInterrupt: jest.fn(() => undefined), + })); + const client = new AgentClient({ + req: { + user: { id: 'user-123' }, + body: { endpoint: EModelEndpoint.agents, agent_id: 'agent-123', isTemporary: true }, + config: { endpoints: { [EModelEndpoint.agents]: {} } }, + _resumableStreamId: streamId, + }, + res: {}, + agent: { + id: 'agent-123', + endpoint: EModelEndpoint.openAI, + provider: EModelEndpoint.openAI, + model_parameters: { model: 'gpt-4' }, + }, + contentParts: [], + collectedUsage: [{ input_tokens: 10, output_tokens: 5 }], + artifactPromises: [], + jobCreatedAt: job.createdAt, + }); + client.conversationId = streamId; + client.responseMessageId = 'response-resume-stopped'; + client.recordCollectedUsage = jest.fn().mockResolvedValue(); + + await client.resumeCompletion({ + resumeValue: { decisions: [] }, + streamId, + checkpointNamespace: 'resume-stopped', + abortController, + }); + + expect(client.recordCollectedUsage).toHaveBeenCalledTimes(1); + expect(client.recordCollectedUsage).toHaveBeenCalledWith( + expect.objectContaining({ context: 'abort' }), + ); + }); + it('publishes the inherited context meta before a fresh run streams', async () => { const streamId = 'conversation-context-meta-stream-seed'; const job = await GenerationJobManager.createJob(streamId, 'user-123', streamId); @@ -2886,6 +2936,93 @@ describe('AgentClient - startup telemetry', () => { ); }); + it('records collected usage as an abort when the run is stopped', async () => { + jest.clearAllMocks(); + const abortController = new AbortController(); + mockCreateRun.mockResolvedValue({ + Graph: null, + processStream: jest.fn(async () => { + abortController.abort(); + }), + getCalibrationRatio: jest.fn(() => 0), + }); + mockIsHITLEnabled.mockReturnValue(false); + const client = new AgentClient({ + req: { + user: { id: 'user-123' }, + body: {}, + config: { endpoints: { [EModelEndpoint.agents]: {} } }, + _resumableStreamId: 'conversation-stopped', + }, + res: {}, + agent: { + id: 'agent-123', + endpoint: EModelEndpoint.openAI, + provider: EModelEndpoint.openAI, + model_parameters: { model: 'gpt-4' }, + hide_sequential_outputs: false, + }, + endpointTokenConfig: {}, + eventHandlers: {}, + contentParts: [], + collectedUsage: [{ input_tokens: 10, output_tokens: 5 }], + artifactPromises: [], + }); + client.conversationId = 'conversation-stopped'; + client.responseMessageId = 'response-conversation-stopped'; + client.parentMessageId = 'parent-conversation-stopped'; + client.recordCollectedUsage = jest.fn().mockResolvedValue(); + + await client.chatCompletion({ payload: [], abortController }); + + expect(client.recordCollectedUsage).toHaveBeenCalledTimes(1); + expect(client.recordCollectedUsage).toHaveBeenCalledWith( + expect.objectContaining({ context: 'abort' }), + ); + }); + + it('records collected usage as a message when the run completes', async () => { + jest.clearAllMocks(); + mockCreateRun.mockResolvedValue({ + Graph: null, + processStream: jest.fn().mockResolvedValue(), + getCalibrationRatio: jest.fn(() => 0), + }); + mockIsHITLEnabled.mockReturnValue(false); + const client = new AgentClient({ + req: { + user: { id: 'user-123' }, + body: {}, + config: { endpoints: { [EModelEndpoint.agents]: {} } }, + _resumableStreamId: 'conversation-completed', + }, + res: {}, + agent: { + id: 'agent-123', + endpoint: EModelEndpoint.openAI, + provider: EModelEndpoint.openAI, + model_parameters: { model: 'gpt-4' }, + hide_sequential_outputs: false, + }, + endpointTokenConfig: {}, + eventHandlers: {}, + contentParts: [], + collectedUsage: [{ input_tokens: 10, output_tokens: 5 }], + artifactPromises: [], + }); + client.conversationId = 'conversation-completed'; + client.responseMessageId = 'response-conversation-completed'; + client.parentMessageId = 'parent-conversation-completed'; + client.recordCollectedUsage = jest.fn().mockResolvedValue(); + + await client.chatCompletion({ payload: [], abortController: new AbortController() }); + + expect(client.recordCollectedUsage).toHaveBeenCalledTimes(1); + expect(client.recordCollectedUsage).toHaveBeenCalledWith( + expect.objectContaining({ context: 'message' }), + ); + }); + it('classifies a terminal chat-model failure without logging provider content', async () => { jest.clearAllMocks(); const { logger } = require('@librechat/data-schemas'); @@ -2985,6 +3122,102 @@ describe('AgentClient - startup telemetry', () => { errorSpy.mockRestore(); }); + /** A compaction's only record of having been one is the marker on the part it + * produced, and Compact runs on whatever leaf the branch ends with. Without + * the marker on the failure, a compaction that failed on a user leaf keeps a + * Regenerate that answers that user message instead of redoing the run. */ + it('marks the failure a compaction turn persists instead of a summary', async () => { + jest.clearAllMocks(); + mockCreateRun.mockImplementation(async () => ({ + Graph: null, + processStream: jest.fn(async () => { + throw new Error('summarizer unavailable'); + }), + getCalibrationRatio: jest.fn(() => 0), + })); + mockIsHITLEnabled.mockReturnValue(false); + const client = new AgentClient({ + req: { + user: { id: 'user-123' }, + body: { compact: true }, + config: { endpoints: { [EModelEndpoint.agents]: {} } }, + _resumableStreamId: 'conversation-compaction-failure', + }, + res: {}, + agent: { + id: 'agent-123', + endpoint: EModelEndpoint.openAI, + provider: EModelEndpoint.openAI, + model_parameters: { model: 'gpt-4' }, + hide_sequential_outputs: false, + }, + endpointTokenConfig: {}, + eventHandlers: {}, + contentParts: [], + collectedUsage: [], + artifactPromises: [], + }); + client.conversationId = 'conversation-compaction-failure'; + client.responseMessageId = 'response-compaction-failure'; + client.parentMessageId = 'parent-compaction-failure'; + client.recordCollectedUsage = jest.fn().mockResolvedValue(); + + const { completion } = await client.sendCompletion([]); + + expect(completion).toEqual([ + expect.objectContaining({ type: ContentTypes.ERROR, initiatedBy: 'user' }), + ]); + }); + + /** A summarizer that returns nothing emits no content at all, so the run has + * neither a summary nor an explanation. The turn records the typed failure + * itself instead of being saved as a bare error row the client cannot tell + * apart from an answer to the message it hangs off. */ + it('records a marked typed failure when a compaction run produces nothing', async () => { + jest.clearAllMocks(); + mockCreateRun.mockImplementation(async () => ({ + Graph: null, + processStream: jest.fn(async () => {}), + getCalibrationRatio: jest.fn(() => 0), + })); + mockIsHITLEnabled.mockReturnValue(false); + const client = new AgentClient({ + req: { + user: { id: 'user-123' }, + body: { compact: true }, + config: { endpoints: { [EModelEndpoint.agents]: {} } }, + _resumableStreamId: 'conversation-compaction-empty', + }, + res: {}, + agent: { + id: 'agent-123', + endpoint: EModelEndpoint.openAI, + provider: EModelEndpoint.openAI, + model_parameters: { model: 'gpt-4' }, + hide_sequential_outputs: false, + }, + endpointTokenConfig: {}, + eventHandlers: {}, + contentParts: [], + collectedUsage: [], + artifactPromises: [], + }); + client.conversationId = 'conversation-compaction-empty'; + client.responseMessageId = 'response-compaction-empty'; + client.parentMessageId = 'parent-compaction-empty'; + client.recordCollectedUsage = jest.fn().mockResolvedValue(); + + const { completion } = await client.sendCompletion([]); + + expect(completion).toEqual([ + { + type: ContentTypes.ERROR, + error: JSON.stringify({ type: ErrorTypes.COMPACTION_FAILED }), + initiatedBy: 'user', + }, + ]); + }); + it('keeps a later non-provider run failure on the generic error path', async () => { jest.clearAllMocks(); const { logger } = require('@librechat/data-schemas'); diff --git a/api/server/controllers/agents/recordCollectedUsage.spec.js b/api/server/controllers/agents/recordCollectedUsage.spec.js index 009c5b262ca..6e009041071 100644 --- a/api/server/controllers/agents/recordCollectedUsage.spec.js +++ b/api/server/controllers/agents/recordCollectedUsage.spec.js @@ -89,6 +89,77 @@ describe('AgentClient - recordCollectedUsage', () => { client.user = 'user-123'; }); + describe('recordTokenUsage fallback', () => { + const estimate = { promptTokens: 40, completionTokens: 7 }; + + it('does not bill the estimate when provider usage was already recorded', async () => { + await client.recordTokenUsage({ + ...estimate, + usage: { input_tokens: 40, output_tokens: 0 }, + model: 'gpt-4', + }); + + expect(mockSpendTokens).not.toHaveBeenCalled(); + }); + + it('bills the estimate when no provider usage was recorded', async () => { + await client.recordTokenUsage({ ...estimate, usage: undefined, model: 'gpt-4' }); + + expect(mockSpendTokens).toHaveBeenCalledTimes(1); + expect(mockSpendTokens).toHaveBeenCalledWith( + expect.objectContaining({ model: 'gpt-4', context: 'message' }), + estimate, + ); + }); + + it('labels the estimate as an abort when the run was stopped and no context is given', async () => { + client.abortController = { signal: { aborted: true } }; + + await client.recordTokenUsage({ ...estimate, usage: undefined, model: 'gpt-4' }); + + expect(mockSpendTokens).toHaveBeenCalledWith( + expect.objectContaining({ context: 'abort' }), + estimate, + ); + }); + + it('labels the estimate as a message when the run completed and no context is given', async () => { + client.abortController = { signal: { aborted: false } }; + + await client.recordTokenUsage({ ...estimate, usage: undefined, model: 'gpt-4' }); + + expect(mockSpendTokens).toHaveBeenCalledWith( + expect.objectContaining({ context: 'message' }), + estimate, + ); + }); + + it('does not bill the estimate when a later primary call was billed but the aggregate hides it', async () => { + client.collectedUsage = [ + { input_tokens: 0, output_tokens: 0 }, + { input_tokens: 5, output_tokens: 0 }, + ]; + + await client.recordTokenUsage({ + ...estimate, + usage: { input_tokens: 0, output_tokens: 0 }, + model: 'gpt-4', + }); + + expect(mockSpendTokens).not.toHaveBeenCalled(); + }); + + it('still bills the estimate when the recorded report is all zero', async () => { + await client.recordTokenUsage({ + ...estimate, + usage: { input_tokens: 0, output_tokens: 0 }, + model: 'gpt-4', + }); + + expect(mockSpendTokens).toHaveBeenCalledTimes(1); + }); + }); + describe('basic functionality', () => { it('should delegate to recordCollectedUsage with full deps', async () => { const collectedUsage = [{ input_tokens: 100, output_tokens: 50, model: 'gpt-4' }]; diff --git a/api/server/controllers/agents/request.js b/api/server/controllers/agents/request.js index 65b92bfec57..9f9b5935d74 100644 --- a/api/server/controllers/agents/request.js +++ b/api/server/controllers/agents/request.js @@ -46,6 +46,7 @@ const { getCodeWorkspaceSelectionErrorDetails, shouldPersistCodeWorkspaceInitializationError, getFailedTurnTraceFields, + resolveFailedTurnContent, } = require('@librechat/api'); const { disposeClient } = require('~/server/cleanup'); const { @@ -490,6 +491,7 @@ async function saveErrorTurn( error: true, unfinished: false, isCreatedByUser: false, + ...resolveFailedTurnContent(req.body, errorText), }, { context }, ); diff --git a/api/server/controllers/assistants/chat.contentFilter.spec.js b/api/server/controllers/assistants/chat.contentFilter.spec.js index 9b3e6301024..f5f9ccc7e80 100644 --- a/api/server/controllers/assistants/chat.contentFilter.spec.js +++ b/api/server/controllers/assistants/chat.contentFilter.spec.js @@ -112,9 +112,9 @@ jest.mock('~/server/middleware/error', () => ({ })); jest.mock('~/models', () => ({ - createAutoRefillTransaction: jest.fn(), - findBalanceByUser: jest.fn(), - upsertBalanceFields: jest.fn(), + releaseBalanceReservation: jest.fn(), + renewBalanceReservation: jest.fn(), + reserveBalance: jest.fn(), getTransactions: jest.fn(), getMultiplier: jest.fn(), getConvo: (...args) => mockGetConvo(...args), @@ -137,6 +137,7 @@ jest.mock('./helpers', () => ({ const chatV1 = require('./chatV1'); const chatV2 = require('./chatV2'); const { logger } = require('@librechat/data-schemas'); +const { checkBalance, getBalanceConfig } = require('@librechat/api'); const { ImageVisionTool } = require('librechat-data-provider'); describe.each([ @@ -237,6 +238,22 @@ describe.each([ expect(mockSendResponse).not.toHaveBeenCalled(); } + it('releases a balance reservation that settles after thread initialization fails', async () => { + req.config.filters = {}; + const release = jest.fn().mockResolvedValue(undefined); + getBalanceConfig.mockReturnValue({ enabled: true }); + checkBalance.mockImplementation( + () => new Promise((resolve) => setTimeout(() => resolve({ release }), 50)), + ); + mockInitThread.mockRejectedValueOnce(new Error('stop after initThread')); + + await chatController(req, res); + + expect(mockInitThread).toHaveBeenCalledTimes(1); + expect(checkBalance).toHaveBeenCalledTimes(1); + expect(release).toHaveBeenCalledTimes(1); + }); + it('blocks persisted instructions before thread, message, run, or stream side effects', async () => { req.config.filters = { agentInstructions: { @@ -571,6 +588,46 @@ describe.each([ } if (_version === 'v2') { + it('keeps the balance reservation until a run that continued in the background settles', async () => { + req.config.filters = {}; + const release = jest.fn().mockResolvedValue(undefined); + getBalanceConfig.mockReturnValue({ enabled: true }); + checkBalance.mockResolvedValue({ release }); + mockInitThread.mockResolvedValueOnce({ thread_id: 'thread-existing' }); + let finishBackgroundRun = () => undefined; + const usage = { prompt_tokens: 1, completion_tokens: 1 }; + mockStreamRunManager + .mockImplementationOnce(() => ({ + runAssistant: jest.fn().mockResolvedValue(undefined), + run: { id: 'run-1', status: 'in_progress', usage }, + intermediateText: '', + messages: [], + })) + .mockImplementationOnce(() => ({ + runAssistant: jest.fn(() => new Promise((resolve) => (finishBackgroundRun = resolve))), + run: { id: 'run-1', status: 'completed', usage }, + intermediateText: '', + messages: [], + })); + + const handled = chatController(req, res); + for (let i = 0; i < 50 && !res.end.mock.calls.length; i++) { + await new Promise((resolve) => setImmediate(resolve)); + } + for (let i = 0; i < 20; i++) { + await new Promise((resolve) => setImmediate(resolve)); + } + + expect(mockStreamRunManager).toHaveBeenCalledTimes(2); + expect(mockHandleError.mock.calls.map(([error]) => error?.message)).toEqual([]); + expect(res.end).toHaveBeenCalled(); + expect(release).not.toHaveBeenCalled(); + + finishBackgroundRun(); + await handled; + expect(release).toHaveBeenCalledTimes(1); + }); + describe('V2 final conversation-file preflight', () => { beforeEach(() => { req.config.filters = { diff --git a/api/server/controllers/assistants/chatV1.js b/api/server/controllers/assistants/chatV1.js index 5db12b87007..a185e544c91 100644 --- a/api/server/controllers/assistants/chatV1.js +++ b/api/server/controllers/assistants/chatV1.js @@ -5,6 +5,7 @@ const { sendEvent, countTokens, checkBalance, + createBalanceReservations, getBalanceConfig, getSafeErrorText, getModelMaxTokens, @@ -47,10 +48,10 @@ const { createRunBody } = require('~/server/services/createRunBody'); const { sendResponse } = require('~/server/middleware/error'); const setHeaders = require('~/server/middleware/setHeaders'); const { - createAutoRefillTransaction, - findBalanceByUser, - upsertBalanceFields, + releaseBalanceReservation, + renewBalanceReservation, getTransactions, + reserveBalance, getMultiplier, getConvo, getFiles, @@ -126,6 +127,7 @@ const chatV1 = async (req, res) => { /** @type {Run | undefined} - The completed run, undefined if incomplete */ let completedRun; let contentRejected = false; + const balanceReservations = createBalanceReservations(); const handleError = async (error) => { const defaultErrorMessage = @@ -304,7 +306,7 @@ const chatV1 = async (req, res) => { // Count tokens up to the current context window promptTokens = Math.min(promptTokens, getModelMaxTokens(model)); - await checkBalance( + return await checkBalance( { req, res, @@ -316,12 +318,12 @@ const chatV1 = async (req, res) => { }, }, { - findBalanceByUser, getMultiplier, - createAutoRefillTransaction, + reserveBalance, + renewBalanceReservation, + releaseBalanceReservation, logViolation, balanceConfig, - upsertBalanceFields, }, ); }; @@ -564,7 +566,7 @@ const chatV1 = async (req, res) => { } } - const promises = [initializeThread(), checkBalanceBeforeRun()]; + const promises = [initializeThread(), balanceReservations.track(checkBalanceBeforeRun())]; await Promise.all(promises); const sendInitialResponse = () => { @@ -683,7 +685,7 @@ const chatV1 = async (req, res) => { } if (response.run.status === RunStatus.IN_PROGRESS) { - processRun(true); + balanceReservations.holdUntil(processRun(true)); } completedRun = response.run; @@ -755,6 +757,8 @@ const chatV1 = async (req, res) => { } } catch (error) { await handleError(error); + } finally { + await balanceReservations.release(); } }; diff --git a/api/server/controllers/assistants/chatV2.js b/api/server/controllers/assistants/chatV2.js index b79053ba2ad..24f3b97691b 100644 --- a/api/server/controllers/assistants/chatV2.js +++ b/api/server/controllers/assistants/chatV2.js @@ -5,6 +5,7 @@ const { sendEvent, countTokens, checkBalance, + createBalanceReservations, getBalanceConfig, getTransactionsConfig, getModelMaxTokens, @@ -44,9 +45,9 @@ const { getConvo, getMultiplier, getTransactions, - findBalanceByUser, - upsertBalanceFields, - createAutoRefillTransaction, + reserveBalance, + renewBalanceReservation, + releaseBalanceReservation, getFiles, } = require('~/models'); const { logViolation, getLogStores } = require('~/cache'); @@ -118,6 +119,7 @@ const chatV2 = async (req, res) => { /** @type {Run | undefined} - The completed run, undefined if incomplete */ let completedRun; let contentRejected = false; + const balanceReservations = createBalanceReservations(); const getContext = () => ({ openai, @@ -175,7 +177,7 @@ const chatV2 = async (req, res) => { // Count tokens up to the current context window promptTokens = Math.min(promptTokens, getModelMaxTokens(model)); - await checkBalance( + return await checkBalance( { req, res, @@ -187,12 +189,12 @@ const chatV2 = async (req, res) => { }, }, { - findBalanceByUser, getMultiplier, - createAutoRefillTransaction, + reserveBalance, + renewBalanceReservation, + releaseBalanceReservation, logViolation, balanceConfig, - upsertBalanceFields, }, ); }; @@ -395,7 +397,7 @@ const chatV2 = async (req, res) => { } } - const promises = [initializeThread(), checkBalanceBeforeRun()]; + const promises = [initializeThread(), balanceReservations.track(checkBalanceBeforeRun())]; await Promise.all(promises); const sendInitialResponse = () => { @@ -520,7 +522,7 @@ const chatV2 = async (req, res) => { } if (response.run.status === RunStatus.IN_PROGRESS) { - processRun(true); + balanceReservations.holdUntil(processRun(true)); } completedRun = response.run; @@ -593,6 +595,8 @@ const chatV2 = async (req, res) => { } } catch (error) { await handleError(error); + } finally { + await balanceReservations.release(); } }; diff --git a/api/server/experimental.js b/api/server/experimental.js index d0173bd728d..b77d5c33fd8 100644 --- a/api/server/experimental.js +++ b/api/server/experimental.js @@ -37,6 +37,7 @@ const { initializeFileStorage, loadToolApprovalHooks, maybeInjectQueryDevtoolsBootstrap, + injectConfiguredFooterBootstrap, preAuthTenantMiddleware, requestContextMiddleware, configureServerTimeouts, @@ -551,6 +552,15 @@ if (cluster.isMaster) { } } + /* The composer lays out against whether a footer bar sits beneath it, and + `/api/config` answers that only after it has painted. One shell serves + every request, before there is a caller whose overrides could be resolved, + so the answer is the deployment's base configuration, like index.js. */ + indexHTML = injectConfiguredFooterBootstrap(indexHTML, { + customFooter: process.env.CUSTOM_FOOTER, + interfaceConfig: baseAppConfig?.interfaceConfig, + }); + const cspPolicy = createCspPolicy(); const shellCache = shellCacheHeaders(cspPolicy != null); @@ -628,7 +638,7 @@ if (cluster.isMaster) { } if (isEnabled(ALLOW_SOCIAL_LOGIN)) { - await configureSocialLogins(app); + await configureSocialLogins(app, appConfig); } app.use(capabilityContextMiddleware); diff --git a/api/server/index.js b/api/server/index.js index 73ce8e45d79..b8c0ec5b517 100644 --- a/api/server/index.js +++ b/api/server/index.js @@ -43,6 +43,7 @@ const { setPluginHookSource, loadToolApprovalHooks, maybeInjectQueryDevtoolsBootstrap, + injectConfiguredFooterBootstrap, preAuthTenantMiddleware, requestContextMiddleware, registerShutdownTask, @@ -284,6 +285,16 @@ const startServer = async () => { } } + /* The composer lays out against whether a footer bar sits beneath it, and + `/api/config` answers that only after it has painted. One shell serves every + request, before there is a caller whose overrides could be resolved, so the + answer is the deployment's base configuration; `/api/config` resolves the + caller's and the client prefers it. */ + indexHTML = injectConfiguredFooterBootstrap(indexHTML, { + customFooter: process.env.CUSTOM_FOOTER, + interfaceConfig: appConfig?.interfaceConfig, + }); + const cspPolicy = createCspPolicy(); const shellCache = shellCacheHeaders(cspPolicy != null); @@ -373,7 +384,7 @@ const startServer = async () => { } if (isEnabled(ALLOW_SOCIAL_LOGIN)) { - await configureSocialLogins(app); + await configureSocialLogins(app, appConfig); } /* Per-request capability cache — must be registered before any route that calls hasCapability */ diff --git a/api/server/index.spec.js b/api/server/index.spec.js index 9949b5e0bc3..e1b565dcc36 100644 --- a/api/server/index.spec.js +++ b/api/server/index.spec.js @@ -133,6 +133,21 @@ describe('Startup readiness wiring', () => { ).toHaveLength(1); }); + it('configures social logins with the app config loaded at startup in both server entries', () => { + const experimental = fs.readFileSync(path.join(__dirname, 'experimental.js'), 'utf8'); + + for (const [name, contents] of [ + ['index.js', source], + ['experimental.js', experimental], + ]) { + const appConfigIndex = contents.indexOf('const appConfig = await getAppConfig('); + const socialLoginsIndex = contents.indexOf('await configureSocialLogins(app, appConfig);'); + + expect([name, appConfigIndex > -1]).toEqual([name, true]); + expect([name, socialLoginsIndex > appConfigIndex]).toEqual([name, true]); + } + }); + it('awaits the shared Redis client before startup cache access', () => { const redisReadyIndex = source.indexOf('await waitForKeyvRedisClient();'); const connectDbIndex = source.indexOf('await connectDb();'); @@ -265,6 +280,9 @@ describe('Server Configuration', () => { mongoServer = await MongoMemoryServer.create(); process.env.MONGO_URI = mongoServer.getUri(); process.env.PORT = '0'; // Use a random available port + /* This deployment configures a footer, so the shell it serves has to say so + before any `/api/config` request: the composer lays out against it. */ + process.env.CUSTOM_FOOTER = 'Operator policy footer'; /* index.js listens at module scope and exports only the app, so capture the server to close it. */ const listenSpy = jest.spyOn(express.application, 'listen'); app = require('~/server'); @@ -279,6 +297,7 @@ describe('Server Configuration', () => { await promisify(server.close).call(server); await mongoServer.stop(); await mongoose.disconnect(); + delete process.env.CUSTOM_FOOTER; }); it('should return OK for /health', async () => { @@ -376,6 +395,19 @@ describe('Server Configuration', () => { expect(directIndexResponse.text).toContain('"enableQueryDevtools":true'); }); + it('serves the configured-footer answer with the shell', async () => { + const [fallbackResponse, indexResponse] = await Promise.all([ + request(app).get('/this/does/not/exist'), + request(app).get('/index.html'), + ]); + + for (const response of [fallbackResponse, indexResponse]) { + expect(response.status).toBe(200); + expect(response.text).toContain('window.__LIBRECHAT_CONFIG__'); + expect(response.text).toContain('"hasConfiguredFooter":true'); + } + }); + it('should return 500 for unknown errors via ErrorController', async () => { // Testing the error handling here on top of unit tests to ensure the middleware is correctly integrated diff --git a/api/server/middleware/abortMiddleware.js b/api/server/middleware/abortMiddleware.js index 76c9cb1f07a..532c748a7aa 100644 --- a/api/server/middleware/abortMiddleware.js +++ b/api/server/middleware/abortMiddleware.js @@ -6,8 +6,6 @@ const { countTokens, isAbortError, GenerationJobManager, - recordCollectedUsage, - getTransactionsConfig, sanitizeMessageForTransmit, buildAbortedResponseMetadata, } = require('@librechat/api'); @@ -17,58 +15,6 @@ const { sendError } = require('~/server/middleware/error'); const { abortRun } = require('./abortRun'); const db = require('~/models'); -/** - * Spend tokens for all models from collected usage. - * This handles both sequential and parallel agent execution. - * - * IMPORTANT: After spending, this function clears the collectedUsage array - * to prevent double-spending. The array is shared with AgentClient.collectedUsage, - * so clearing it here prevents the finally block from also spending tokens. - * - * @param {Object} params - * @param {string} params.userId - User ID - * @param {string} params.conversationId - Conversation ID - * @param {Array} params.collectedUsage - Usage metadata from all models - * @param {string} [params.fallbackModel] - Fallback model name if not in usage - * @param {string} [params.messageId] - The response message ID for transaction correlation - * @param {AppConfig['transactions']} [params.transactions] - Resolved transactions config - */ -async function spendCollectedUsage({ - userId, - conversationId, - collectedUsage, - fallbackModel, - messageId, - transactions, -}) { - if (!collectedUsage || collectedUsage.length === 0) { - return; - } - - await recordCollectedUsage( - { - spendTokens: db.spendTokens, - spendStructuredTokens: db.spendStructuredTokens, - pricing: { getMultiplier: db.getMultiplier, getCacheMultiplier: db.getCacheMultiplier }, - bulkWriteOps: { insertMany: db.bulkInsertTransactions, updateBalance: db.updateBalance }, - }, - { - user: userId, - conversationId, - collectedUsage, - context: 'abort', - messageId, - model: fallbackModel, - transactions, - }, - ); - - // Clear the array to prevent double-spending from the AgentClient finally block. - // The collectedUsage array is shared by reference with AgentClient.collectedUsage, - // so clearing it here ensures recordCollectedUsage() sees an empty array and returns early. - collectedUsage.length = 0; -} - /** * Abort an active message generation. * Uses GenerationJobManager for all agent requests. @@ -94,10 +40,9 @@ async function abortMessage(req, res) { return; } - const { jobData, content, text, collectedUsage } = abortResult; + const { jobData, content, text } = abortResult; const completionTokens = await countTokens(text); - const promptTokens = jobData?.promptTokens ?? 0; const responseMessage = { messageId: jobData?.responseMessageId, @@ -129,25 +74,9 @@ async function abortMessage(req, res) { responseMessage.metadata = abortMetadata; } - const transactions = getTransactionsConfig(req.config); - - // Spend tokens for ALL models from collectedUsage (handles parallel agents/addedConvo) - if (collectedUsage && collectedUsage.length > 0) { - await spendCollectedUsage({ - userId, - conversationId: jobData?.conversationId, - collectedUsage, - fallbackModel: jobData?.model, - messageId: jobData?.responseMessageId, - transactions, - }); - } else { - // Fallback: no collected usage, use text-based token counting for primary model only - await db.spendTokens( - { ...responseMessage, context: 'incomplete', user: userId, transactions }, - { promptTokens, completionTokens }, - ); - } + /** The run that produced this response records its own usage on exit + * (`AgentClient` labels a stopped turn `'abort'`), so billing here would + * charge it a second time. This route only stops and persists. */ await db.saveMessage( { @@ -302,5 +231,4 @@ const handleAbortError = async (res, req, error, data) => { module.exports = { handleAbort, handleAbortError, - spendCollectedUsage, }; diff --git a/api/server/middleware/abortMiddleware.spec.js b/api/server/middleware/abortMiddleware.spec.js index 8a64feae896..ee7ec98c208 100644 --- a/api/server/middleware/abortMiddleware.spec.js +++ b/api/server/middleware/abortMiddleware.spec.js @@ -1,14 +1,9 @@ /** - * Tests for abortMiddleware - spendCollectedUsage function + * Tests for abortMiddleware. * - * This tests the token spending logic for abort scenarios, - * particularly for parallel agents (addedConvo) where multiple - * models need their tokens spent. - * - * spendCollectedUsage delegates to recordCollectedUsage from @librechat/api, - * passing pricing + bulkWriteOps deps, with context: 'abort'. - * After spending, it clears the collectedUsage array to prevent double-spending - * from the AgentClient finally block (which shares the same array reference). + * The run that produced a stopped response records its own usage on exit + * (AgentClient labels it 'abort'), so this route must not bill: it only stops + * the job and persists the partial response. */ const mockSpendTokens = jest.fn().mockResolvedValue(); @@ -86,7 +81,7 @@ const { logger } = require('@librechat/data-schemas'); const { sendError } = require('~/server/middleware/error'); const { GenerationJobManager } = require('@librechat/api'); const db = require('~/models'); -const { handleAbort, handleAbortError, spendCollectedUsage } = require('./abortMiddleware'); +const { handleAbort, handleAbortError } = require('./abortMiddleware'); const buildAbortRequest = () => ({ body: { @@ -97,169 +92,6 @@ const buildAbortRequest = () => ({ }, }); -describe('abortMiddleware - spendCollectedUsage', () => { - beforeEach(() => { - jest.clearAllMocks(); - }); - - describe('spendCollectedUsage delegation', () => { - it('should return early if collectedUsage is empty', async () => { - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage: [], - fallbackModel: 'gpt-4', - }); - - expect(mockRecordCollectedUsage).not.toHaveBeenCalled(); - }); - - it('should return early if collectedUsage is null', async () => { - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage: null, - fallbackModel: 'gpt-4', - }); - - expect(mockRecordCollectedUsage).not.toHaveBeenCalled(); - }); - - it('should call recordCollectedUsage with abort context and full deps', async () => { - const collectedUsage = [{ input_tokens: 100, output_tokens: 50, model: 'gpt-4' }]; - - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage, - fallbackModel: 'gpt-4', - messageId: 'msg-123', - }); - - expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1); - expect(mockRecordCollectedUsage).toHaveBeenCalledWith( - { - spendTokens: expect.any(Function), - spendStructuredTokens: expect.any(Function), - pricing: { - getMultiplier: mockGetMultiplier, - getCacheMultiplier: mockGetCacheMultiplier, - }, - bulkWriteOps: { - insertMany: mockBulkInsertTransactions, - updateBalance: mockUpdateBalance, - }, - }, - { - user: 'user-123', - conversationId: 'convo-123', - collectedUsage, - context: 'abort', - messageId: 'msg-123', - model: 'gpt-4', - }, - ); - }); - - it('should pass context abort for multiple models (parallel agents)', async () => { - const collectedUsage = [ - { input_tokens: 100, output_tokens: 50, model: 'gpt-4' }, - { input_tokens: 80, output_tokens: 40, model: 'claude-3' }, - { input_tokens: 120, output_tokens: 60, model: 'gemini-pro' }, - ]; - - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage, - fallbackModel: 'gpt-4', - }); - - expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1); - expect(mockRecordCollectedUsage).toHaveBeenCalledWith( - expect.any(Object), - expect.objectContaining({ - context: 'abort', - collectedUsage, - }), - ); - }); - - it('should handle real-world parallel agent abort scenario', async () => { - const collectedUsage = [ - { input_tokens: 31596, output_tokens: 151, model: 'gemini-3-flash-preview' }, - { input_tokens: 28000, output_tokens: 120, model: 'gpt-5.2' }, - ]; - - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage, - fallbackModel: 'gemini-3-flash-preview', - }); - - expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1); - expect(mockRecordCollectedUsage).toHaveBeenCalledWith( - expect.any(Object), - expect.objectContaining({ - user: 'user-123', - conversationId: 'convo-123', - context: 'abort', - model: 'gemini-3-flash-preview', - }), - ); - }); - - /** - * Race condition prevention: after abort middleware spends tokens, - * the collectedUsage array is cleared so AgentClient.recordCollectedUsage() - * (which shares the same array reference) sees an empty array and returns early. - */ - it('should clear collectedUsage array after spending to prevent double-spending', async () => { - const collectedUsage = [ - { input_tokens: 100, output_tokens: 50, model: 'gpt-4' }, - { input_tokens: 80, output_tokens: 40, model: 'claude-3' }, - ]; - - expect(collectedUsage.length).toBe(2); - - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage, - fallbackModel: 'gpt-4', - }); - - expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1); - expect(collectedUsage.length).toBe(0); - }); - - it('should await recordCollectedUsage before clearing array', async () => { - let resolved = false; - mockRecordCollectedUsage.mockImplementation(async () => { - await new Promise((resolve) => setTimeout(resolve, 10)); - resolved = true; - return { input_tokens: 100, output_tokens: 50 }; - }); - - const collectedUsage = [ - { input_tokens: 100, output_tokens: 50, model: 'gpt-4' }, - { input_tokens: 80, output_tokens: 40, model: 'claude-3' }, - ]; - - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage, - fallbackModel: 'gpt-4', - }); - - expect(resolved).toBe(true); - expect(collectedUsage.length).toBe(0); - }); - }); -}); - describe('abortMiddleware - handleAbortError', () => { beforeEach(() => { jest.clearAllMocks(); @@ -328,7 +160,7 @@ describe('abortMiddleware - handleAbortError', () => { * caller-supplied data, so an omitted value is indistinguishable from enabled and * the write proceeds even when `transactions.enabled` is false. */ -describe('abortMiddleware - transactions config', () => { +describe('abortMiddleware - handleAbort billing', () => { const buildJobData = () => ({ model: 'gpt-4', responseMessageId: 'msg-123', @@ -363,25 +195,7 @@ describe('abortMiddleware - transactions config', () => { db.getConvo.mockResolvedValue({ title: 'Test Chat' }); }); - it('forwards transactions through spendCollectedUsage to recordCollectedUsage', async () => { - const collectedUsage = [{ input_tokens: 100, output_tokens: 50, model: 'gpt-4' }]; - - await spendCollectedUsage({ - userId: 'user-123', - conversationId: 'convo-123', - collectedUsage, - fallbackModel: 'gpt-4', - transactions: { enabled: false }, - }); - - expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1); - expect(mockRecordCollectedUsage).toHaveBeenCalledWith( - expect.any(Object), - expect.objectContaining({ context: 'abort', transactions: { enabled: false } }), - ); - }); - - it('resolves the config from req and forwards it on the collected-usage path', async () => { + it('leaves billing to the run even when the stopped job collected usage', async () => { const collectedUsage = [{ input_tokens: 100, output_tokens: 50, model: 'gpt-4' }]; GenerationJobManager.abortJob.mockResolvedValue({ success: true, @@ -391,16 +205,14 @@ describe('abortMiddleware - transactions config', () => { collectedUsage, }); - const req = buildReq(); - await handleAbort()(req, buildRes()); + await handleAbort()(buildReq(), buildRes()); expect(logger.error).not.toHaveBeenCalled(); - expect(mockGetTransactionsConfig).toHaveBeenCalledWith(req.config); - expect(mockRecordCollectedUsage).toHaveBeenCalledTimes(1); - expect(mockRecordCollectedUsage).toHaveBeenCalledWith( - expect.any(Object), - expect.objectContaining({ context: 'abort', transactions: { enabled: false } }), - ); + expect(mockRecordCollectedUsage).not.toHaveBeenCalled(); + expect(mockSpendTokens).not.toHaveBeenCalled(); + expect(mockSpendStructuredTokens).not.toHaveBeenCalled(); + expect(collectedUsage).toHaveLength(1); + expect(db.saveMessage).toHaveBeenCalledTimes(1); }); it('carries the context meta the run published onto the job into the stopped response', async () => { @@ -444,7 +256,7 @@ describe('abortMiddleware - transactions config', () => { expect(savedMessage.contextMeta).toBeNull(); }); - it('resolves the config from req and forwards it on the token-count fallback path', async () => { + it('does not bill a stopped response by token count either', async () => { GenerationJobManager.abortJob.mockResolvedValue({ success: true, jobData: buildJobData(), @@ -453,16 +265,11 @@ describe('abortMiddleware - transactions config', () => { collectedUsage: [], }); - const req = buildReq(); - await handleAbort()(req, buildRes()); + await handleAbort()(buildReq(), buildRes()); expect(logger.error).not.toHaveBeenCalled(); - expect(mockGetTransactionsConfig).toHaveBeenCalledWith(req.config); expect(mockRecordCollectedUsage).not.toHaveBeenCalled(); - expect(mockSpendTokens).toHaveBeenCalledTimes(1); - expect(mockSpendTokens).toHaveBeenCalledWith( - expect.objectContaining({ context: 'incomplete', transactions: { enabled: false } }), - expect.any(Object), - ); + expect(mockSpendTokens).not.toHaveBeenCalled(); + expect(db.saveMessage).toHaveBeenCalledTimes(1); }); }); diff --git a/api/server/middleware/checkPeoplePickerAccess.js b/api/server/middleware/checkPeoplePickerAccess.js index 83329f56a31..016c35f082d 100644 --- a/api/server/middleware/checkPeoplePickerAccess.js +++ b/api/server/middleware/checkPeoplePickerAccess.js @@ -1,114 +1,7 @@ -const { logger } = require('@librechat/data-schemas'); -const { - PrincipalType, - PermissionTypes, - Permissions, - SystemRoles, -} = require('librechat-data-provider'); +const { createPeoplePickerAccess } = require('@librechat/api'); const { getRoleByName } = require('~/models'); -const VALID_PRINCIPAL_TYPES = new Set([ - PrincipalType.USER, - PrincipalType.GROUP, - PrincipalType.ROLE, -]); - -/** - * Middleware to check if user has permission to access people picker functionality. - * Validates requested principal types via `type` (singular) and `types` (comma-separated or array) - * query parameters against the caller's role permissions: - * - user: requires VIEW_USERS permission - * - group: requires VIEW_GROUPS permission - * - role: requires VIEW_ROLES permission - * - no type filter (mixed search): requires at least one of the above - */ -const checkPeoplePickerAccess = async (req, res, next) => { - try { - const user = req.user; - if (!user || !user.role) { - return res.status(401).json({ - error: 'Unauthorized', - message: 'Authentication required', - }); - } - - if (user.role === SystemRoles.ADMIN) { - return next(); - } - - const role = await getRoleByName(user.role); - if (!role || !role.permissions) { - return res.status(403).json({ - error: 'Forbidden', - message: 'No permissions configured for user role', - }); - } - - const { type, types } = req.query; - const peoplePickerPerms = role.permissions[PermissionTypes.PEOPLE_PICKER] || {}; - const canViewUsers = peoplePickerPerms[Permissions.VIEW_USERS] === true; - const canViewGroups = peoplePickerPerms[Permissions.VIEW_GROUPS] === true; - const canViewRoles = peoplePickerPerms[Permissions.VIEW_ROLES] === true; - - const permissionChecks = { - [PrincipalType.USER]: { - hasPermission: canViewUsers, - message: 'Insufficient permissions to search for users', - }, - [PrincipalType.GROUP]: { - hasPermission: canViewGroups, - message: 'Insufficient permissions to search for groups', - }, - [PrincipalType.ROLE]: { - hasPermission: canViewRoles, - message: 'Insufficient permissions to search for roles', - }, - }; - - const requestedTypes = new Set(); - - if (type && VALID_PRINCIPAL_TYPES.has(type)) { - requestedTypes.add(type); - } - - if (types) { - const typesArray = Array.isArray(types) ? types : types.split(','); - for (const t of typesArray) { - if (VALID_PRINCIPAL_TYPES.has(t)) { - requestedTypes.add(t); - } - } - } - - for (const requested of requestedTypes) { - const check = permissionChecks[requested]; - if (!check.hasPermission) { - return res.status(403).json({ - error: 'Forbidden', - message: check.message, - }); - } - } - - if (requestedTypes.size === 0 && !canViewUsers && !canViewGroups && !canViewRoles) { - return res.status(403).json({ - error: 'Forbidden', - message: 'Insufficient permissions to search for users, groups, or roles', - }); - } - - next(); - } catch (error) { - logger.error( - `[checkPeoplePickerAccess][${req.user?.id}] error for type=${req.query.type}, types=${req.query.types}`, - error, - ); - return res.status(500).json({ - error: 'Internal Server Error', - message: 'Failed to check permissions', - }); - } -}; +const checkPeoplePickerAccess = createPeoplePickerAccess({ getRoleByName }); module.exports = { checkPeoplePickerAccess, diff --git a/api/server/middleware/checkPeoplePickerAccess.spec.js b/api/server/middleware/checkPeoplePickerAccess.spec.js deleted file mode 100644 index 45097ef2905..00000000000 --- a/api/server/middleware/checkPeoplePickerAccess.spec.js +++ /dev/null @@ -1,431 +0,0 @@ -const { logger } = require('@librechat/data-schemas'); -const { - PrincipalType, - PermissionTypes, - Permissions, - SystemRoles, -} = require('librechat-data-provider'); -const { checkPeoplePickerAccess } = require('./checkPeoplePickerAccess'); -const { getRoleByName } = require('~/models'); - -jest.mock('~/models'); -jest.mock('@librechat/data-schemas', () => ({ - ...jest.requireActual('@librechat/data-schemas'), - logger: { - error: jest.fn(), - }, -})); - -describe('checkPeoplePickerAccess', () => { - let req, res, next; - - beforeEach(() => { - req = { - user: { id: 'user123', role: 'USER' }, - query: {}, - }; - res = { - status: jest.fn().mockReturnThis(), - json: jest.fn(), - }; - next = jest.fn(); - jest.clearAllMocks(); - }); - - it('should return 401 if user is not authenticated', async () => { - req.user = null; - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(401); - expect(res.json).toHaveBeenCalledWith({ - error: 'Unauthorized', - message: 'Authentication required', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should return 403 if role has no permissions', async () => { - getRoleByName.mockResolvedValue(null); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'No permissions configured for user role', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('allows a literal admin regardless of configured people picker permissions', async () => { - req.user.role = SystemRoles.ADMIN; - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(getRoleByName).not.toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should allow access when searching for users with VIEW_USERS permission', async () => { - req.query.type = PrincipalType.USER; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should deny access when searching for users without VIEW_USERS permission', async () => { - req.query.type = PrincipalType.USER; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: false, - [Permissions.VIEW_GROUPS]: true, - [Permissions.VIEW_ROLES]: true, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for users', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should allow access when searching for groups with VIEW_GROUPS permission', async () => { - req.query.type = PrincipalType.GROUP; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: false, - [Permissions.VIEW_GROUPS]: true, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should deny access when searching for groups without VIEW_GROUPS permission', async () => { - req.query.type = PrincipalType.GROUP; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: true, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for groups', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should allow access when searching for roles with VIEW_ROLES permission', async () => { - req.query.type = PrincipalType.ROLE; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: false, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: true, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should deny access when searching for roles without VIEW_ROLES permission', async () => { - req.query.type = PrincipalType.ROLE; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: true, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for roles', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should deny access when using types param to bypass type-specific check', async () => { - req.query.types = PrincipalType.GROUP; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for groups', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should deny access when types contains any unpermitted type', async () => { - req.query.types = `${PrincipalType.USER},${PrincipalType.ROLE}`; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for roles', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should allow access when all requested types are permitted', async () => { - req.query.types = `${PrincipalType.USER},${PrincipalType.GROUP}`; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: true, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should validate types when provided as array (Express qs parsing)', async () => { - req.query.types = [PrincipalType.GROUP, PrincipalType.ROLE]; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: true, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for groups', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should enforce permissions for combined type and types params', async () => { - req.query.type = PrincipalType.USER; - req.query.types = PrincipalType.GROUP; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for groups', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should treat all-invalid types values as mixed search', async () => { - req.query.types = 'foobar'; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should deny when types is empty string and user has no permissions', async () => { - req.query.types = ''; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: false, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for users, groups, or roles', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should treat types=public as mixed search since PUBLIC is not a searchable principal type', async () => { - req.query.types = PrincipalType.PUBLIC; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: true, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should allow mixed search when user has at least one permission', async () => { - // No type specified = mixed search - req.query.type = undefined; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: false, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: true, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(next).toHaveBeenCalled(); - expect(res.status).not.toHaveBeenCalled(); - }); - - it('should deny mixed search when user has no permissions', async () => { - // No type specified = mixed search - req.query.type = undefined; - getRoleByName.mockResolvedValue({ - permissions: { - [PermissionTypes.PEOPLE_PICKER]: { - [Permissions.VIEW_USERS]: false, - [Permissions.VIEW_GROUPS]: false, - [Permissions.VIEW_ROLES]: false, - }, - }, - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for users, groups, or roles', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should handle errors gracefully', async () => { - const error = new Error('Database error'); - getRoleByName.mockRejectedValue(error); - - await checkPeoplePickerAccess(req, res, next); - - expect(logger.error).toHaveBeenCalledWith( - '[checkPeoplePickerAccess][user123] error for type=undefined, types=undefined', - error, - ); - expect(res.status).toHaveBeenCalledWith(500); - expect(res.json).toHaveBeenCalledWith({ - error: 'Internal Server Error', - message: 'Failed to check permissions', - }); - expect(next).not.toHaveBeenCalled(); - }); - - it('should handle missing permissions object gracefully', async () => { - req.query.type = PrincipalType.USER; - getRoleByName.mockResolvedValue({ - permissions: {}, // No PEOPLE_PICKER permissions - }); - - await checkPeoplePickerAccess(req, res, next); - - expect(res.status).toHaveBeenCalledWith(403); - expect(res.json).toHaveBeenCalledWith({ - error: 'Forbidden', - message: 'Insufficient permissions to search for users', - }); - expect(next).not.toHaveBeenCalled(); - }); -}); diff --git a/api/server/middleware/index.js b/api/server/middleware/index.js index 8f39ea1ee89..814baf075b1 100644 --- a/api/server/middleware/index.js +++ b/api/server/middleware/index.js @@ -11,6 +11,7 @@ const { } = require('./messageValidation'); const checkDomainAllowed = require('./checkDomainAllowed'); const { markOAuthNavigation } = require('./oauthNavigation'); +const requireSameOrigin = require('./requireSameOrigin'); const requireLocalAuth = require('./requireLocalAuth'); const canDeleteAccount = require('./canDeleteAccount'); const accessResources = require('./accessResources'); @@ -51,6 +52,7 @@ module.exports = { checkInviteUser, requireLdapAuth, requireLocalAuth, + requireSameOrigin, canDeleteAccount, configMiddleware, checkDomainAllowed, diff --git a/api/server/middleware/requireSameOrigin.js b/api/server/middleware/requireSameOrigin.js new file mode 100644 index 00000000000..815063a9b58 --- /dev/null +++ b/api/server/middleware/requireSameOrigin.js @@ -0,0 +1,9 @@ +const { createSameOriginGuard } = require('@librechat/api'); + +module.exports = createSameOriginGuard({ + trustedOrigins: [ + process.env.DOMAIN_CLIENT, + process.env.DOMAIN_SERVER, + process.env.ADMIN_PANEL_URL, + ], +}); diff --git a/api/server/routes/admin/auth.js b/api/server/routes/admin/auth.js index fbc53c09074..2af6f903678 100644 --- a/api/server/routes/admin/auth.js +++ b/api/server/routes/admin/auth.js @@ -156,6 +156,7 @@ function buildGoogleAdminRefreshDeps(sessionExpiry) { router.post( '/login/local', middleware.logHeaders, + middleware.requireSameOrigin, middleware.loginLimiter, middleware.checkBan, middleware.validateEmailLogin, diff --git a/api/server/routes/admin/auth.refresh.test.js b/api/server/routes/admin/auth.refresh.test.js index 3539c2d285e..c6e0287054c 100644 --- a/api/server/routes/admin/auth.refresh.test.js +++ b/api/server/routes/admin/auth.refresh.test.js @@ -98,6 +98,7 @@ jest.mock('~/strategies', () => ({ jest.mock('~/server/middleware', () => ({ logHeaders: jest.fn((req, res, next) => next()), + requireSameOrigin: jest.fn((req, res, next) => next()), loginLimiter: jest.fn((req, res, next) => next()), checkBan: jest.fn((req, res, next) => next()), validateEmailLogin: jest.fn((req, res, next) => next()), @@ -606,6 +607,33 @@ describe('admin local login route', () => { ); }); + it('rejects cross-site submissions before rate limiting or local auth', async () => { + const response = await request(app).post('/api/admin/login/local').send({ + email: 'admin@example.com', + password: 'password', + }); + + expect(response.status).toBe(200); + expect(middleware.requireSameOrigin).toHaveBeenCalledTimes(1); + expect(middleware.requireSameOrigin.mock.invocationCallOrder[0]).toBeLessThan( + middleware.loginLimiter.mock.invocationCallOrder[0], + ); + + jest.clearAllMocks(); + middleware.requireSameOrigin.mockImplementationOnce((req, res) => + res.status(403).json({ message: 'Cross-site request rejected' }), + ); + + const rejected = await request(app).post('/api/admin/login/local').send({ + email: 'admin@example.com', + password: 'password', + }); + + expect(rejected.status).toBe(403); + expect(middleware.loginLimiter).not.toHaveBeenCalled(); + expect(middleware.requireLocalAuth).not.toHaveBeenCalled(); + }); + it('stops before local auth when the email login gate rejects the request', async () => { middleware.validateEmailLogin.mockImplementationOnce((req, res) => res.status(403).json({ message: 'Email login is not allowed.' }), diff --git a/api/server/routes/auth.2fa-ratelimit.test.js b/api/server/routes/auth.2fa-ratelimit.test.js index 4867f78afec..c352fd64bfa 100644 --- a/api/server/routes/auth.2fa-ratelimit.test.js +++ b/api/server/routes/auth.2fa-ratelimit.test.js @@ -52,6 +52,7 @@ jest.mock('~/server/middleware', () => { const pass = (req, res, next) => next(); return { logHeaders: pass, + requireSameOrigin: pass, loginLimiter: pass, setTwoFactorTempUser: (...args) => mockSetTwoFactorTempUser(...args), twoFactorTempLimiter: (...args) => mockTwoFactorTempLimiter(...args), diff --git a/api/server/routes/auth.cloudfront.test.js b/api/server/routes/auth.cloudfront.test.js index 56f2acf06fb..ed7803d5731 100644 --- a/api/server/routes/auth.cloudfront.test.js +++ b/api/server/routes/auth.cloudfront.test.js @@ -49,6 +49,7 @@ jest.mock('~/server/middleware', () => { const pass = (req, res, next) => next(); return { logHeaders: pass, + requireSameOrigin: pass, loginLimiter: pass, setTwoFactorTempUser: pass, twoFactorTempLimiter: pass, diff --git a/api/server/routes/auth.cross-site.test.js b/api/server/routes/auth.cross-site.test.js new file mode 100644 index 00000000000..0613d66adb2 --- /dev/null +++ b/api/server/routes/auth.cross-site.test.js @@ -0,0 +1,188 @@ +const express = require('express'); +const request = require('supertest'); +const { ErrorTypes } = require('librechat-data-provider'); + +const mockLoginLimiter = jest.fn((req, res, next) => next()); +const mockRequireLocalAuth = jest.fn((req, res, next) => next()); +const mockSetTwoFactorTempUser = jest.fn((req, res, next) => next()); +const mockLoginController = jest.fn((req, res) => res.status(204).end()); +const mockVerify2FAWithTempToken = jest.fn((req, res) => res.status(204).end()); + +jest.mock('@librechat/data-schemas', () => ({ + ...jest.requireActual('@librechat/data-schemas'), + logger: { debug: jest.fn(), info: jest.fn(), warn: jest.fn(), error: jest.fn() }, +})); + +jest.mock('@librechat/api', () => ({ + ...jest.requireActual('@librechat/api'), + createSetBalanceConfig: jest.fn(() => (req, res, next) => next()), +})); + +jest.mock('~/server/controllers/AuthController', () => ({ + refreshController: jest.fn((req, res) => res.status(204).end()), + registrationController: jest.fn((req, res) => res.status(204).end()), + resetPasswordController: jest.fn((req, res) => res.status(204).end()), + resetPasswordRequestController: jest.fn((req, res) => res.status(204).end()), + graphTokenController: jest.fn((req, res) => res.status(204).end()), +})); + +jest.mock('~/server/controllers/TwoFactorController', () => ({ + enable2FA: jest.fn((req, res) => res.status(204).end()), + verify2FA: jest.fn((req, res) => res.status(204).end()), + confirm2FA: jest.fn((req, res) => res.status(204).end()), + disable2FA: jest.fn((req, res) => res.status(204).end()), + regenerateBackupCodes: jest.fn((req, res) => res.status(204).end()), +})); + +jest.mock('~/server/controllers/auth/TwoFactorAuthController', () => ({ + verify2FAWithTempToken: (...args) => mockVerify2FAWithTempToken(...args), +})); + +jest.mock('~/server/controllers/auth/LogoutController', () => ({ + logoutController: jest.fn((req, res) => res.status(204).end()), +})); + +jest.mock('~/server/controllers/auth/LoginController', () => ({ + loginController: (...args) => mockLoginController(...args), +})); + +jest.mock('~/models', () => ({ + findBalanceByUser: jest.fn(), + upsertBalanceFields: jest.fn(), +})); + +jest.mock('~/server/services/Config', () => ({ + getAppConfig: jest.fn(), +})); + +jest.mock('~/server/middleware', () => { + const pass = (req, res, next) => next(); + return { + logHeaders: pass, + requireSameOrigin: jest.requireActual('~/server/middleware/requireSameOrigin'), + loginLimiter: (...args) => mockLoginLimiter(...args), + setTwoFactorTempUser: (...args) => mockSetTwoFactorTempUser(...args), + twoFactorTempLimiter: pass, + checkBan: pass, + validateEmailLogin: pass, + requireLocalAuth: (...args) => mockRequireLocalAuth(...args), + requireLdapAuth: (...args) => mockRequireLocalAuth(...args), + registerLimiter: pass, + checkInviteUser: pass, + validateRegistration: pass, + resetPasswordLimiter: pass, + resetPasswordSubmissionLimiter: pass, + validatePasswordReset: pass, + requireJwtAuth: pass, + }; +}); + +const ORIGINAL_ENV = process.env; +const APP_ORIGIN = 'https://chat.example.com'; +const OTHER_ORIGIN = 'https://other-site.example.net'; +const ADMIN_PANEL_ORIGIN = 'https://admin.example.com'; + +describe('local login endpoints reject cross-site submissions', () => { + let app; + + beforeAll(() => { + process.env = { + ...ORIGINAL_ENV, + DOMAIN_CLIENT: APP_ORIGIN, + DOMAIN_SERVER: APP_ORIGIN, + ADMIN_PANEL_URL: `${ADMIN_PANEL_ORIGIN}/`, + }; + app = express(); + app.use(express.json()); + app.use(express.urlencoded({ extended: true })); + app.use('/api/auth', require('./auth')); + }); + + afterAll(() => { + process.env = ORIGINAL_ENV; + }); + + beforeEach(() => { + jest.clearAllMocks(); + }); + + it('rejects a login form submitted from another site before authenticating', async () => { + const response = await request(app) + .post('/api/auth/login') + .set('Host', 'chat.example.com') + .set('Sec-Fetch-Site', 'cross-site') + .set('Origin', OTHER_ORIGIN) + .type('form') + .send({ email: 'other@example.com', password: 'other-password' }) + .expect(403); + + expect(response.body).toEqual({ + message: 'Cross-site request rejected', + code: ErrorTypes.AUTH_CROSS_ORIGIN, + }); + expect(response.headers['set-cookie']).toBeUndefined(); + expect(mockLoginLimiter).not.toHaveBeenCalled(); + expect(mockRequireLocalAuth).not.toHaveBeenCalled(); + expect(mockLoginController).not.toHaveBeenCalled(); + }); + + it('rejects a cross-site temp-token 2FA submission before verifying it', async () => { + await request(app) + .post('/api/auth/2fa/verify-temp') + .set('Host', 'chat.example.com') + .set('Sec-Fetch-Site', 'cross-site') + .set('Origin', OTHER_ORIGIN) + .type('form') + .send({ tempToken: 'other-temp-token', token: '123456' }) + .expect(403); + + expect(mockSetTwoFactorTempUser).not.toHaveBeenCalled(); + expect(mockVerify2FAWithTempToken).not.toHaveBeenCalled(); + }); + + it('accepts the login form the app submits from its own origin', async () => { + await request(app) + .post('/api/auth/login') + .set('Host', 'chat.example.com') + .set('Sec-Fetch-Site', 'same-origin') + .set('Origin', APP_ORIGIN) + .send({ email: 'user@example.com', password: 'password' }) + .expect(204); + + expect(mockRequireLocalAuth).toHaveBeenCalledTimes(1); + expect(mockLoginController).toHaveBeenCalledTimes(1); + }); + + it('accepts a 2FA submission from the configured client origin', async () => { + await request(app) + .post('/api/auth/2fa/verify-temp') + .set('Host', 'api.example.com') + .set('Sec-Fetch-Site', 'same-site') + .set('Origin', APP_ORIGIN) + .send({ tempToken: 'temp-token', token: '123456' }) + .expect(204); + + expect(mockVerify2FAWithTempToken).toHaveBeenCalledTimes(1); + }); + + it('accepts a login posted from the configured admin panel origin', async () => { + await request(app) + .post('/api/auth/login') + .set('Host', 'chat.example.com') + .set('Sec-Fetch-Site', 'same-site') + .set('Origin', ADMIN_PANEL_ORIGIN) + .send({ email: 'admin@example.com', password: 'password' }) + .expect(204); + + expect(mockLoginController).toHaveBeenCalledTimes(1); + }); + + it('accepts a server-side client that sends no browser headers', async () => { + await request(app) + .post('/api/auth/login') + .send({ email: 'user@example.com', password: 'password' }) + .expect(204); + + expect(mockLoginController).toHaveBeenCalledTimes(1); + }); +}); diff --git a/api/server/routes/auth.js b/api/server/routes/auth.js index 6c942dff7bf..99a20944921 100644 --- a/api/server/routes/auth.js +++ b/api/server/routes/auth.js @@ -43,6 +43,7 @@ router.post('/logout', middleware.requireJwtAuth, logoutController); router.post( '/login', middleware.logHeaders, + middleware.requireSameOrigin, middleware.loginLimiter, middleware.checkBan, middleware.validateEmailLogin, @@ -91,6 +92,7 @@ router.post('/2fa/enable', middleware.requireJwtAuth, enable2FA); router.post('/2fa/verify', middleware.requireJwtAuth, verify2FA); router.post( '/2fa/verify-temp', + middleware.requireSameOrigin, middleware.setTwoFactorTempUser, middleware.twoFactorTempLimiter, middleware.checkBan, diff --git a/api/server/routes/auth.reset-password-ratelimit.test.js b/api/server/routes/auth.reset-password-ratelimit.test.js index 7d49576d7c1..5aeebef408c 100644 --- a/api/server/routes/auth.reset-password-ratelimit.test.js +++ b/api/server/routes/auth.reset-password-ratelimit.test.js @@ -52,6 +52,7 @@ jest.mock('~/server/middleware', () => { const pass = (req, res, next) => next(); return { logHeaders: pass, + requireSameOrigin: pass, loginLimiter: pass, setTwoFactorTempUser: pass, twoFactorTempLimiter: pass, diff --git a/api/server/routes/oauth.js b/api/server/routes/oauth.js index 2d87bb43a00..710fac1c632 100644 --- a/api/server/routes/oauth.js +++ b/api/server/routes/oauth.js @@ -81,7 +81,6 @@ router.get( '/google/callback', passport.authenticate('google', { failureRedirect: `${domains.client}/oauth/error`, - failureMessage: true, session: false, scope: ['openid', 'profile', 'email'], }), @@ -106,7 +105,6 @@ router.get( '/facebook/callback', passport.authenticate('facebook', { failureRedirect: `${domains.client}/oauth/error`, - failureMessage: true, session: false, scope: ['public_profile'], profileFields: ['id', 'email', 'name'], @@ -149,7 +147,6 @@ router.get( '/github/callback', passport.authenticate('github', { failureRedirect: `${domains.client}/oauth/error`, - failureMessage: true, session: false, scope: ['user:email', 'read:user'], }), @@ -173,7 +170,6 @@ router.get( '/discord/callback', passport.authenticate('discord', { failureRedirect: `${domains.client}/oauth/error`, - failureMessage: true, session: false, scope: ['identify', 'email'], }), @@ -196,7 +192,6 @@ router.post( '/apple/callback', passport.authenticate('apple', { failureRedirect: `${domains.client}/oauth/error`, - failureMessage: true, session: false, }), setBalanceConfig, diff --git a/api/server/routes/oauth.state.test.js b/api/server/routes/oauth.state.test.js new file mode 100644 index 00000000000..2f440598704 --- /dev/null +++ b/api/server/routes/oauth.state.test.js @@ -0,0 +1,363 @@ +const express = require('express'); +const request = require('supertest'); +const passport = require('passport'); +const cookieParser = require('cookie-parser'); + +jest.mock('@librechat/data-schemas', () => ({ + ...jest.requireActual('@librechat/data-schemas'), + logger: { debug: jest.fn(), info: jest.fn(), warn: jest.fn(), error: jest.fn() }, +})); + +jest.mock('~/server/services/Config', () => ({ + getAppConfig: jest.fn().mockResolvedValue({}), +})); + +jest.mock('~/models', () => ({ + findUser: jest.fn(), + updateUser: jest.fn(), + findBalanceByUser: jest.fn(), + upsertBalanceFields: jest.fn(), +})); + +jest.mock('~/strategies/process', () => ({ + createSocialUser: jest.fn(), + handleExistingUser: jest.fn().mockResolvedValue(undefined), +})); + +jest.mock('~/server/middleware', () => ({ + logHeaders: (req, res, next) => next(), + loginLimiter: (req, res, next) => next(), + markOAuthNavigation: (req, res, next) => next(), + checkDomainAllowed: (req, res, next) => next(), +})); + +jest.mock('~/server/controllers/auth/oauth', () => ({ + createOAuthHandler: () => (req, res) => res.status(200).json({ userId: req.user._id }), +})); + +const ORIGINAL_ENV = process.env; +const APP_URL = 'https://chat.example.com'; + +/** Splits a `Set-Cookie` header into its name, value and lower-cased attributes. */ +function parseSetCookie(header) { + const [pair, ...attributes] = header.split(';').map((part) => part.trim()); + const separator = pair.indexOf('='); + return { + name: pair.slice(0, separator), + value: decodeURIComponent(pair.slice(separator + 1)), + attributes: attributes.map((attribute) => attribute.toLowerCase()), + }; +} + +function getStateCookie(response, provider) { + const headers = response.headers['set-cookie'] ?? []; + return headers + .map(parseSetCookie) + .find(({ name }) => name.startsWith(`__Host-oauth_state_${provider}.`)); +} + +/** The `name=value` pair a browser sends back for a cookie the server set. */ +const asCookie = ({ name, value }) => `${name}=${value}`; + +const stateLocation = (response) => new URL(response.headers.location).searchParams.get('state'); +const BINDING_PATTERN = /^[A-Za-z0-9_-]{43}$/; + +describe('OAuth login state binding', () => { + let app; + let github; + let apple; + let findUser; + + beforeAll(() => { + process.env = { + ...ORIGINAL_ENV, + DOMAIN_CLIENT: APP_URL, + DOMAIN_SERVER: APP_URL, + GITHUB_CLIENT_ID: 'github-client', + GITHUB_CLIENT_SECRET: 'github-secret', + GITHUB_CALLBACK_URL: '/oauth/github/callback', + APPLE_CLIENT_ID: 'apple-client', + APPLE_TEAM_ID: 'apple-team', + APPLE_KEY_ID: 'apple-key', + APPLE_PRIVATE_KEY_PATH: '/nonexistent/apple.p8', + APPLE_CALLBACK_URL: '/oauth/apple/callback', + }; + + const githubStrategy = require('~/strategies/githubStrategy'); + const appleStrategy = require('~/strategies/appleStrategy'); + ({ findUser } = require('~/models')); + + const stateOptions = { secret: 'jwt-secret', secureCookie: true }; + github = githubStrategy(stateOptions); + apple = appleStrategy(stateOptions); + passport.use(github); + passport.use(apple); + + app = express(); + app.use(express.urlencoded({ extended: true })); + app.use(cookieParser()); + app.use(passport.initialize()); + app.use('/oauth', require('./oauth')); + app.use((err, req, res, _next) => res.status(500).json({ message: err.message })); + }); + + afterAll(() => { + passport.unuse('github'); + passport.unuse('apple'); + process.env = ORIGINAL_ENV; + }); + + beforeEach(() => { + jest.clearAllMocks(); + github._oauth2.getOAuthAccessToken = jest.fn((code, params, callback) => + callback(null, 'provider-access-token', 'provider-refresh-token', {}), + ); + github.userProfile = jest.fn((accessToken, done) => + done(null, { + id: '42', + username: 'user', + displayName: 'User', + emails: [{ value: 'user@example.com', verified: true }], + photos: [{ value: 'https://avatars.example.com/42' }], + }), + ); + apple._oauth2.getOAuthAccessToken = jest.fn((code, params, callback) => + callback(new Error('token exchange reached')), + ); + findUser.mockResolvedValue({ + _id: 'user-1', + provider: 'github', + githubId: '42', + email: 'user@example.com', + }); + }); + + describe('GitHub', () => { + it('signs the authorization state for a host-only binding cookie', async () => { + const response = await request(app).get('/oauth/github').expect(302); + + const location = new URL(response.headers.location); + const cookie = getStateCookie(response, 'github'); + + expect(location.origin).toBe('https://github.com'); + expect(location.searchParams.get('state')).toBeTruthy(); + expect(cookie.value).toMatch(BINDING_PATTERN); + expect(location.searchParams.get('state')).not.toContain(cookie.value); + expect(cookie.attributes).toEqual( + expect.arrayContaining(['httponly', 'secure', 'samesite=lax', 'path=/']), + ); + expect(cookie.attributes).toContain('max-age=600'); + }); + + it('issues a fresh state for every authorization request', async () => { + const first = await request(app).get('/oauth/github').expect(302); + const second = await request(app).get('/oauth/github').expect(302); + + expect(stateLocation(first)).not.toBe(stateLocation(second)); + }); + + it('rejects a callback link from a browser that never started the flow', async () => { + const otherFlow = await request(app).get('/oauth/github').expect(302); + const otherState = new URL(otherFlow.headers.location).searchParams.get('state'); + + const response = await request(app) + .get('/oauth/github/callback') + .query({ code: 'unrelated-code', state: otherState }) + .expect(302); + + expect(response.headers.location).toBe(`${APP_URL}/oauth/error`); + expect(github._oauth2.getOAuthAccessToken).not.toHaveBeenCalled(); + expect(findUser).not.toHaveBeenCalled(); + }); + + it('rejects a callback whose state belongs to a different flow', async () => { + const ownFlow = await request(app).get('/oauth/github').expect(302); + const otherFlow = await request(app).get('/oauth/github').expect(302); + const ownCookie = getStateCookie(ownFlow, 'github'); + const otherState = new URL(otherFlow.headers.location).searchParams.get('state'); + + const response = await request(app) + .get('/oauth/github/callback') + .set('Cookie', asCookie(ownCookie)) + .query({ code: 'unrelated-code', state: otherState }) + .expect(302); + + expect(response.headers.location).toBe(`${APP_URL}/oauth/error`); + expect(github._oauth2.getOAuthAccessToken).not.toHaveBeenCalled(); + }); + + it('rejects a callback that omits the state', async () => { + const ownFlow = await request(app).get('/oauth/github').expect(302); + const ownCookie = getStateCookie(ownFlow, 'github'); + + const response = await request(app) + .get('/oauth/github/callback') + .set('Cookie', asCookie(ownCookie)) + .query({ code: 'unrelated-code' }) + .expect(302); + + expect(response.headers.location).toBe(`${APP_URL}/oauth/error`); + expect(github._oauth2.getOAuthAccessToken).not.toHaveBeenCalled(); + }); + + it('completes the login when the callback returns to the browser that started it', async () => { + const flow = await request(app).get('/oauth/github').expect(302); + const cookie = getStateCookie(flow, 'github'); + const state = new URL(flow.headers.location).searchParams.get('state'); + + const response = await request(app) + .get('/oauth/github/callback') + .set('Cookie', asCookie(cookie)) + .query({ code: 'own-code', state }) + .expect(200); + + expect(response.body).toEqual({ userId: 'user-1' }); + expect(github._oauth2.getOAuthAccessToken).toHaveBeenCalledWith( + 'own-code', + expect.objectContaining({ redirect_uri: `${APP_URL}/oauth/github/callback` }), + expect.any(Function), + ); + + expect(getStateCookie(response, 'github')).toBeUndefined(); + }); + + it('still completes the pending login after an unrelated callback reaches the browser', async () => { + const ownFlow = await request(app).get('/oauth/github').expect(302); + const otherFlow = await request(app).get('/oauth/github').expect(302); + const ownCookie = getStateCookie(ownFlow, 'github'); + + const unrelated = await request(app) + .get('/oauth/github/callback') + .set('Cookie', asCookie(ownCookie)) + .query({ code: 'unrelated-code', state: stateLocation(otherFlow) }) + .expect(302); + + expect(unrelated.headers.location).toBe(`${APP_URL}/oauth/error`); + expect(getStateCookie(unrelated, 'github')).toBeUndefined(); + + await request(app) + .get('/oauth/github/callback') + .set('Cookie', asCookie(ownCookie)) + .query({ code: 'own-code', state: stateLocation(ownFlow) }) + .expect(200); + + expect(github._oauth2.getOAuthAccessToken).toHaveBeenCalledTimes(1); + expect(github._oauth2.getOAuthAccessToken).toHaveBeenCalledWith( + 'own-code', + expect.any(Object), + expect.any(Function), + ); + }); + + it('completes logins started in two tabs, in either order', async () => { + const firstTab = await request(app).get('/oauth/github').expect(302); + const secondTab = await request(app) + .get('/oauth/github') + .set('Cookie', asCookie(getStateCookie(firstTab, 'github'))) + .expect(302); + const binding = asCookie(getStateCookie(secondTab, 'github')); + + expect(binding).toBe(asCookie(getStateCookie(firstTab, 'github'))); + + await request(app) + .get('/oauth/github/callback') + .set('Cookie', binding) + .query({ code: 'first-code', state: stateLocation(firstTab) }) + .expect(200); + + await request(app) + .get('/oauth/github/callback') + .set('Cookie', binding) + .query({ code: 'second-code', state: stateLocation(secondTab) }) + .expect(200); + + expect(github._oauth2.getOAuthAccessToken).toHaveBeenCalledTimes(2); + }); + + it('completes both logins when two tabs start before either has a binding', async () => { + const [firstTab, secondTab] = await Promise.all([ + request(app).get('/oauth/github').expect(302), + request(app).get('/oauth/github').expect(302), + ]); + const firstCookie = getStateCookie(firstTab, 'github'); + const secondCookie = getStateCookie(secondTab, 'github'); + const browserCookies = `${asCookie(firstCookie)}; ${asCookie(secondCookie)}`; + + expect(firstCookie.name).not.toBe(secondCookie.name); + + for (const [tab, code] of [ + [secondTab, 'second-code'], + [firstTab, 'first-code'], + ]) { + await request(app) + .get('/oauth/github/callback') + .set('Cookie', browserCookies) + .query({ code, state: stateLocation(tab) }) + .expect(200); + } + + expect(github._oauth2.getOAuthAccessToken).toHaveBeenCalledTimes(2); + }); + }); + + describe('Apple', () => { + it('stores a SameSite=None state cookie that survives the cross-site form_post callback', async () => { + const response = await request(app).get('/oauth/apple').expect(302); + + const location = new URL(response.headers.location); + const cookie = getStateCookie(response, 'apple'); + + expect(location.origin).toBe('https://appleid.apple.com'); + expect(location.searchParams.get('response_mode')).toBe('form_post'); + expect(cookie.value).toMatch(BINDING_PATTERN); + expect(location.searchParams.get('state')).not.toContain(cookie.value); + expect(cookie.attributes).toEqual( + expect.arrayContaining(['httponly', 'secure', 'samesite=none', 'path=/']), + ); + }); + + it('issues a fresh state for every authorization request', async () => { + const first = await request(app).get('/oauth/apple').expect(302); + const second = await request(app).get('/oauth/apple').expect(302); + + const firstState = new URL(first.headers.location).searchParams.get('state'); + const secondState = new URL(second.headers.location).searchParams.get('state'); + + expect(firstState).not.toBe(secondState); + expect(getStateCookie(second, 'apple').value).toMatch(BINDING_PATTERN); + }); + + it('rejects a form_post callback from a browser that never started the flow', async () => { + const otherFlow = await request(app).get('/oauth/apple').expect(302); + const otherState = new URL(otherFlow.headers.location).searchParams.get('state'); + + const response = await request(app) + .post('/oauth/apple/callback') + .type('form') + .send({ code: 'unrelated-code', state: otherState }) + .expect(302); + + expect(response.headers.location).toBe(`${APP_URL}/oauth/error`); + expect(apple._oauth2.getOAuthAccessToken).not.toHaveBeenCalled(); + }); + + it('exchanges the code when the form_post returns to the browser that started the flow', async () => { + const flow = await request(app).get('/oauth/apple').expect(302); + const cookie = getStateCookie(flow, 'apple'); + const state = new URL(flow.headers.location).searchParams.get('state'); + + await request(app) + .post('/oauth/apple/callback') + .set('Cookie', asCookie(cookie)) + .type('form') + .send({ code: 'own-code', state }) + .expect(500); + + expect(apple._oauth2.getOAuthAccessToken).toHaveBeenCalledWith( + 'own-code', + expect.any(Object), + expect.any(Function), + ); + }); + }); +}); diff --git a/api/server/services/Schedules/index.js b/api/server/services/Schedules/index.js index 4e9c10101ce..50944efd019 100644 --- a/api/server/services/Schedules/index.js +++ b/api/server/services/Schedules/index.js @@ -27,18 +27,9 @@ function getService() { getAppConfig, findUserById: (userId) => mongoose.models.User.findById(userId).select('_id tenantId role').lean(), - findBalance: (userId) => mongoose.models.Balance.findOne({ user: userId }).lean(), + findBalance: (userId) => methods.findBalanceByUser(userId, { includeReservedCredits: true }), upsertBalance: (userId, { set, setOnInsert }) => - mongoose.models.Balance.findOneAndUpdate( - { user: userId }, - { - ...(set && Object.keys(set).length > 0 ? { $set: set } : {}), - ...(setOnInsert && Object.keys(setOnInsert).length > 0 - ? { $setOnInsert: setOnInsert } - : {}), - }, - { upsert: true, new: true }, - ).lean(), + methods.upsertBalanceFields(userId, set ?? {}, setOnInsert ?? {}), // Compare-and-set: only initialize an existing record while its credit is still null. // No upsert — a CAS miss must re-read the winner, not insert a fresh balance. `null` // in the filter also matches a legacy record whose `tokenCredits` field is absent. @@ -46,7 +37,7 @@ function getService() { mongoose.models.Balance.findOneAndUpdate( { user: userId, tokenCredits: null }, { $set: { tokenCredits, ...(sync && Object.keys(sync).length > 0 ? sync : {}) } }, - { new: true }, + { new: true, sort: { _id: 1 } }, ).lean(), enqueueAgentTrigger, // Reconciliation reads the durable delivery to tell a still-live admission (a diff --git a/api/server/socialLogins.js b/api/server/socialLogins.js index f4d088e6d03..138718fd2b5 100644 --- a/api/server/socialLogins.js +++ b/api/server/socialLogins.js @@ -1,7 +1,12 @@ const passport = require('passport'); const session = require('express-session'); const { CacheKeys } = require('librechat-data-provider'); -const { math, isEnabled, shouldUseSecureCookie } = require('@librechat/api'); +const { + math, + isEnabled, + shouldUseSecureCookie, + registerOpenIdWithRetry, +} = require('@librechat/api'); const { logger, DEFAULT_SESSION_EXPIRY } = require('@librechat/data-schemas'); const { openIdJwtLogin, @@ -21,7 +26,6 @@ const { const { getLogStores } = require('~/cache'); const DEFAULT_OPENID_REUSE_MAX_SESSION_AGE_MS = 15 * 60 * 1000; - const getSessionExpiry = () => math(process.env.SESSION_EXPIRY, DEFAULT_SESSION_EXPIRY); const getOpenIdSessionExpiry = () => { @@ -40,9 +44,10 @@ const getOpenIdSessionExpiry = () => { /** * Configures OpenID Connect for the application. * @param {Express.Application} app - The Express application instance. + * @param {AppConfig} [appConfig] - Base app config, read for OpenID discovery retry settings. * @returns {Promise} */ -async function configureOpenId(app) { +async function configureOpenId(app, appConfig) { logger.info('Configuring OpenID Connect...'); const sessionExpiry = getOpenIdSessionExpiry(); const sessionOptions = { @@ -58,44 +63,49 @@ async function configureOpenId(app) { app.use(session(sessionOptions)); app.use(passport.session()); - const config = await setupOpenId(); - if (!config) { - logger.error('OpenID Connect configuration failed - strategy not registered.'); - return; - } - - if (isEnabled(process.env.OPENID_REUSE_TOKENS)) { - logger.info('OpenID token reuse is enabled.'); - passport.use('openidJwt', openIdJwtLogin(config)); - } - logger.info('OpenID Connect configured successfully.'); + await registerOpenIdWithRetry({ + setupOpenId, + registerJwtStrategy: (config) => passport.use('openidJwt', openIdJwtLogin(config)), + reuseTokens: isEnabled(process.env.OPENID_REUSE_TOKENS), + discovery: appConfig?.registration?.openidDiscovery, + env: { + startupAttempts: process.env.OPENID_DISCOVERY_RETRY_ATTEMPTS, + retryDelayMs: process.env.OPENID_DISCOVERY_RETRY_DELAY_MS, + }, + }); } /** * * @param {Express.Application} app + * @param {AppConfig} [appConfig] - Base app config, read for the social login state lifetime. */ -const configureSocialLogins = async (app) => { +const configureSocialLogins = async (app, appConfig) => { logger.info('Configuring social logins...'); + const stateOptions = { + secret: process.env.JWT_SECRET, + secureCookie: shouldUseSecureCookie(), + maxAgeMs: appConfig?.registration?.oauthStateTtlMs, + }; if (process.env.GOOGLE_CLIENT_ID && process.env.GOOGLE_CLIENT_SECRET) { - passport.use(googleLogin()); + passport.use(googleLogin(stateOptions)); passport.use('googleAdmin', googleAdminLogin()); } if (process.env.FACEBOOK_CLIENT_ID && process.env.FACEBOOK_CLIENT_SECRET) { - passport.use(facebookLogin()); + passport.use(facebookLogin(stateOptions)); passport.use('facebookAdmin', facebookAdminLogin()); } if (process.env.GITHUB_CLIENT_ID && process.env.GITHUB_CLIENT_SECRET) { - passport.use(githubLogin()); + passport.use(githubLogin(stateOptions)); passport.use('githubAdmin', githubAdminLogin()); } if (process.env.DISCORD_CLIENT_ID && process.env.DISCORD_CLIENT_SECRET) { - passport.use(discordLogin()); + passport.use(discordLogin(stateOptions)); passport.use('discordAdmin', discordAdminLogin()); } if (process.env.APPLE_CLIENT_ID && process.env.APPLE_PRIVATE_KEY_PATH) { - passport.use(appleLogin()); + passport.use(appleLogin(stateOptions)); passport.use('appleAdmin', appleAdminLogin()); } if ( @@ -105,7 +115,7 @@ const configureSocialLogins = async (app) => { process.env.OPENID_SCOPE && process.env.OPENID_SESSION_SECRET ) { - await configureOpenId(app); + await configureOpenId(app, appConfig); } if ( process.env.SAML_ENTRY_POINT && diff --git a/api/server/socialLogins.spec.js b/api/server/socialLogins.spec.js index bf016a43ebb..cbf6e19a5ce 100644 --- a/api/server/socialLogins.spec.js +++ b/api/server/socialLogins.spec.js @@ -9,6 +9,14 @@ const mockSetupOpenId = jest.fn(); const mockSetupSaml = jest.fn(); const mockIsEnabled = jest.fn(); const mockShouldUseSecureCookie = jest.fn(() => true); +const mockRegisterOpenIdWithRetry = jest.fn( + async ({ setupOpenId, registerJwtStrategy, reuseTokens }) => { + const config = await setupOpenId(); + if (config && reuseTokens) { + registerJwtStrategy(config); + } + }, +); const mockMath = jest.fn((value, fallback) => { if (value == null || value === '') { return fallback; @@ -42,10 +50,11 @@ jest.mock('@librechat/api', () => ({ math: (...args) => mockMath(...args), isEnabled: (...args) => mockIsEnabled(...args), shouldUseSecureCookie: (...args) => mockShouldUseSecureCookie(...args), + registerOpenIdWithRetry: (...args) => mockRegisterOpenIdWithRetry(...args), })); jest.mock('@librechat/data-schemas', () => ({ DEFAULT_SESSION_EXPIRY: 900000, - logger: { error: jest.fn(), info: jest.fn() }, + logger: { error: jest.fn(), info: jest.fn(), warn: jest.fn() }, })); jest.mock('~/cache', () => ({ getLogStores: (...args) => mockGetLogStores(...args) })); jest.mock('~/strategies', () => ({ @@ -90,6 +99,10 @@ describe('configureSocialLogins OpenID session expiry', () => { process.env = ORIGINAL_ENV; }); + afterEach(() => { + jest.useRealTimers(); + }); + it('extends the OpenID session cookie to the reuse window when token reuse is enabled', async () => { process.env.SESSION_EXPIRY = '1000 * 60 * 15'; process.env.OPENID_REUSE_TOKENS = 'true'; @@ -140,4 +153,68 @@ describe('configureSocialLogins OpenID session expiry', () => { ); expect(mockPassportUse).not.toHaveBeenCalled(); }); + + it('passes OpenID strategy wiring and both retry sources to the API package', async () => { + process.env.OPENID_DISCOVERY_RETRY_ATTEMPTS = '2'; + process.env.OPENID_DISCOVERY_RETRY_DELAY_MS = '1000'; + const discovery = { startupAttempts: 0, retryDelayMs: 250 }; + const app = { use: jest.fn() }; + + await configureSocialLogins(app, { registration: { openidDiscovery: discovery } }); + + expect(mockRegisterOpenIdWithRetry).toHaveBeenCalledWith({ + setupOpenId: expect.any(Function), + registerJwtStrategy: expect.any(Function), + reuseTokens: false, + discovery, + env: { startupAttempts: '2', retryDelayMs: '1000' }, + }); + expect(mockSetupOpenId).toHaveBeenCalledTimes(1); + }); +}); + +describe('configureSocialLogins OAuth state options', () => { + const ORIGINAL_ENV = process.env; + + beforeEach(() => { + jest.clearAllMocks(); + process.env = { + JWT_SECRET: 'jwt-secret', + GITHUB_CLIENT_ID: 'github-client', + GITHUB_CLIENT_SECRET: 'github-secret', + APPLE_CLIENT_ID: 'apple-client', + APPLE_PRIVATE_KEY_PATH: '/keys/apple.p8', + }; + mockIsEnabled.mockReturnValue(false); + }); + + afterAll(() => { + process.env = ORIGINAL_ENV; + }); + + it('passes the cookie security setting and configured state lifetime to user strategies', async () => { + const strategies = require('~/strategies'); + const app = { use: jest.fn() }; + + await configureSocialLogins(app, { registration: { oauthStateTtlMs: 120000 } }); + + const expected = { secret: 'jwt-secret', secureCookie: true, maxAgeMs: 120000 }; + expect(strategies.githubLogin).toHaveBeenCalledWith(expected); + expect(strategies.appleLogin).toHaveBeenCalledWith(expected); + expect(strategies.githubAdminLogin).toHaveBeenCalledWith(); + expect(strategies.appleAdminLogin).toHaveBeenCalledWith(); + }); + + it('leaves the state lifetime to the store default when the config omits it', async () => { + const strategies = require('~/strategies'); + mockShouldUseSecureCookie.mockReturnValueOnce(false); + + await configureSocialLogins({ use: jest.fn() }); + + expect(strategies.githubLogin).toHaveBeenCalledWith({ + secret: 'jwt-secret', + secureCookie: false, + maxAgeMs: undefined, + }); + }); }); diff --git a/api/server/utils/import/fork.js b/api/server/utils/import/fork.js index 467aabc0dea..b777eed80da 100644 --- a/api/server/utils/import/fork.js +++ b/api/server/utils/import/fork.js @@ -1,5 +1,5 @@ const { v4: uuidv4 } = require('uuid'); -const { withoutTraceRefs } = require('@librechat/api'); +const { cloneLineage, withoutTraceRefs, getAllMessagesUpToParent } = require('@librechat/api'); const { logger, tenantStorage } = require('@librechat/data-schemas'); const { EModelEndpoint, Constants, ForkOptions } = require('librechat-data-provider'); const { getConvo, getMessages, getSharedMessages } = require('~/models'); @@ -21,53 +21,12 @@ function cloneMessagesWithTimestamps( importBatchBuilder, { detachSubagentRuntime = false } = {}, ) { - const idMapping = new Map(); - - // First pass: create ID mapping and sort messages by parentMessageId - const sortedMessages = [...messagesToClone].sort((a, b) => { - if (a.parentMessageId === Constants.NO_PARENT) { - return -1; - } - if (b.parentMessageId === Constants.NO_PARENT) { - return 1; - } - return 0; - }); - - // Helper function to ensure date object - const ensureDate = (dateValue) => { - if (!dateValue) { - return new Date(); - } - return dateValue instanceof Date ? dateValue : new Date(dateValue); - }; - - // Second pass: clone messages while maintaining proper timestamps - for (const message of sortedMessages) { - const newMessageId = uuidv4(); - idMapping.set(message.messageId, newMessageId); - - const parentId = - message.parentMessageId && message.parentMessageId !== Constants.NO_PARENT - ? idMapping.get(message.parentMessageId) - : Constants.NO_PARENT; - - // If this message has a parent, ensure its timestamp is after the parent's - let createdAt = ensureDate(message.createdAt); - if (parentId !== Constants.NO_PARENT) { - const parentMessage = importBatchBuilder.messages.find((msg) => msg.messageId === parentId); - if (parentMessage) { - const parentDate = ensureDate(parentMessage.createdAt); - if (createdAt <= parentDate) { - createdAt = new Date(parentDate.getTime() + 1); - } - } - } - + const { entries, idMapping } = cloneLineage(messagesToClone, uuidv4); + for (const { source, messageId, parentMessageId, createdAt } of entries) { const clonedMessage = { - ...withoutTraceRefs(message), - messageId: newMessageId, - parentMessageId: parentId, + ...withoutTraceRefs(source), + messageId, + parentMessageId, createdAt, }; if (detachSubagentRuntime) { @@ -189,48 +148,6 @@ async function forkConversation({ } } -/** - * Retrieves all messages up to the root from the target message. - * @param {TMessage[]} messages - The list of messages to search. - * @param {string} targetMessageId - The ID of the target message. - * @returns {TMessage[]} The list of messages up to the root from the target message. - */ -function getAllMessagesUpToParent(messages, targetMessageId) { - const targetMessage = messages.find((msg) => msg.messageId === targetMessageId); - if (!targetMessage) { - return []; - } - - const pathToRoot = new Set(); - const visited = new Set(); - let current = targetMessage; - - while (current) { - if (visited.has(current.messageId)) { - break; - } - - visited.add(current.messageId); - pathToRoot.add(current.messageId); - - const currentParentId = current.parentMessageId ?? Constants.NO_PARENT; - if (currentParentId === Constants.NO_PARENT) { - break; - } - - current = messages.find((msg) => msg.messageId === currentParentId); - } - - // Include all messages that are in the path or whose parent is in the path - // Exclude children of the target message - return messages.filter( - (msg) => - (pathToRoot.has(msg.messageId) && msg.messageId !== targetMessageId) || - (pathToRoot.has(msg.parentMessageId) && msg.parentMessageId !== targetMessageId) || - msg.messageId === targetMessageId, - ); -} - /** * Retrieves all messages up to the root from the target message and its neighbors. * @param {TMessage[]} messages - The list of messages to search. diff --git a/api/server/utils/import/importers-timestamp.spec.js b/api/server/utils/import/importers-timestamp.spec.js index 268cc74c0d8..90572edd2cb 100644 --- a/api/server/utils/import/importers-timestamp.spec.js +++ b/api/server/utils/import/importers-timestamp.spec.js @@ -530,4 +530,119 @@ describe('Import Timestamp Ordering', () => { ); }); }); + + describe('Large exports', () => { + const baseTime = 1700000000; + const chatGptNode = (parent, role) => ({ + parent, + children: [], + message: { + author: { role }, + create_time: baseTime, + content: { content_type: 'text', parts: [role] }, + metadata: {}, + }, + }); + + /** These sizes import in a few hundred milliseconds; a per-message scan takes ten seconds or more. */ + const importBudgetMs = 3000; + + const importJson = async (jsonData) => { + const importBatchBuilder = new ImportBatchBuilder('user-123'); + const startedAt = performance.now(); + await getImporter(jsonData)(jsonData, 'user-123', () => importBatchBuilder); + return { messages: importBatchBuilder.messages, elapsedMs: performance.now() - startedAt }; + }; + + test('imports a flat LibreChat export whose messages all name an absent parent', async () => { + const count = 60000; + const { messages, elapsedMs } = await importJson({ + conversationId: 'large-flat', + title: 'Large flat export', + messages: Array.from({ length: count }, (_, index) => ({ + messageId: `m${index}`, + parentMessageId: 'absent-root', + text: 'x', + sender: 'user', + isCreatedByUser: true, + })), + }); + + expect(elapsedMs).toBeLessThan(importBudgetMs); + expect(messages).toHaveLength(count); + }); + + test('orders a long LibreChat chain whose timestamps all collide', async () => { + const count = 60000; + const createdAt = '2024-01-01T00:00:00.000Z'; + const { messages, elapsedMs } = await importJson({ + conversationId: 'large-chain', + title: 'Large chain export', + messages: Array.from({ length: count }, (_, index) => ({ + messageId: `m${index}`, + parentMessageId: index === 0 ? Constants.NO_PARENT : `m${index - 1}`, + text: 'x', + sender: 'user', + isCreatedByUser: true, + createdAt, + })), + }); + + expect(elapsedMs).toBeLessThan(importBudgetMs); + expect(messages).toHaveLength(count); + expect(messages[count - 1].parentMessageId).toBe(messages[count - 2].messageId); + expect(messages[count - 1].createdAt.getTime()).toBe( + new Date(createdAt).getTime() + count - 1, + ); + }); + + test('orders a long ChatGPT branch listed deepest-first', async () => { + const count = 10000; + const mapping = {}; + for (let depth = count - 1; depth >= 0; depth--) { + mapping[`n${depth}`] = chatGptNode( + depth === 0 ? null : `n${depth - 1}`, + depth % 2 ? 'assistant' : 'user', + ); + } + + const { messages, elapsedMs } = await importJson([ + { title: 'Deep branch', create_time: baseTime, mapping }, + ]); + + const byId = new Map(messages.map((message) => [message.messageId, message])); + const root = messages.find((message) => message.parentMessageId === Constants.NO_PARENT); + const leaf = messages[0]; + expect(elapsedMs).toBeLessThan(importBudgetMs); + expect(messages).toHaveLength(count); + expect(root.createdAt.getTime()).toBe(baseTime * 1000); + expect(leaf.createdAt.getTime()).toBe(baseTime * 1000 + count - 1); + expect(leaf.createdAt.getTime()).toBeGreaterThan( + byId.get(leaf.parentMessageId).createdAt.getTime(), + ); + }); + + test('attaches many ChatGPT replies behind one long run of system messages', async () => { + const systemCount = 10000; + const replyCount = 10000; + const mapping = { root: chatGptNode(null, 'user') }; + for (let index = 0; index < systemCount; index++) { + mapping[`s${index}`] = chatGptNode(index === 0 ? 'root' : `s${index - 1}`, 'system'); + } + for (let index = 0; index < replyCount; index++) { + mapping[`r${index}`] = chatGptNode(`s${systemCount - 1}`, 'assistant'); + } + + const { messages, elapsedMs } = await importJson([ + { title: 'System run', create_time: baseTime, mapping }, + ]); + + const root = messages.find((message) => message.parentMessageId === Constants.NO_PARENT); + expect(elapsedMs).toBeLessThan(importBudgetMs); + expect(messages).toHaveLength(replyCount + 1); + expect(messages.filter((message) => message.parentMessageId === root.messageId)).toHaveLength( + replyCount, + ); + }); + }); }); diff --git a/api/server/utils/import/importers.js b/api/server/utils/import/importers.js index d1d2edd29df..d496749e5cd 100644 --- a/api/server/utils/import/importers.js +++ b/api/server/utils/import/importers.js @@ -7,7 +7,12 @@ const { stripMessageUIResourceMarkers, } = require('@librechat/data-schemas'); const { EModelEndpoint, Constants, Tools, openAISettings } = require('librechat-data-provider'); -const { withoutTraceRefs } = require('@librechat/api'); +const { + withoutTraceRefs, + orderMessageLineage, + createChatGptLineage, + linkChatGptCitations, +} = require('@librechat/api'); const { getEndpointsConfig } = require('~/server/services/Config'); const { createImportBatchBuilder } = require('./importBatchBuilder'); const { resolveImportDefaultModel } = require('./defaults'); @@ -445,86 +450,7 @@ function processConversation(conv, importBatchBuilder, requestUserId, defaultMod } } - /** - * Finds the nearest valid parent by traversing up through skippable messages - * (system, reasoning_recap, thoughts). Uses iterative traversal to avoid - * stack overflow on deep chains of skippable messages. - * - * @param {string} startId - The ID of the starting parent message. - * @returns {string} The ID of the nearest valid parent message. - */ - const findValidParent = (startId) => { - const visited = new Set(); - let parentId = startId; - - while (parentId) { - if (!messageMap.has(parentId) || visited.has(parentId)) { - return Constants.NO_PARENT; - } - visited.add(parentId); - - const parentMapping = conv.mapping[parentId]; - if (!parentMapping?.message) { - return Constants.NO_PARENT; - } - - const contentType = parentMapping.message.content?.content_type; - const shouldSkip = - parentMapping.message.author?.role === 'system' || - contentType === 'reasoning_recap' || - contentType === 'thoughts'; - - if (!shouldSkip) { - return messageMap.get(parentId); - } - - parentId = parentMapping.parent; - } - - return Constants.NO_PARENT; - }; - - /** - * Helper function to find thinking content from parent chain (thoughts messages) - * @param {string} parentId - The ID of the parent message. - * @param {Set} visited - Set of already-visited IDs to prevent cycles. - * @returns {Array} The thinking content array (empty if not found). - */ - const findThinkingContent = (parentId, visited = new Set()) => { - // Guard against circular references in malformed imports - if (!parentId || visited.has(parentId)) { - return []; - } - visited.add(parentId); - - const parentMapping = conv.mapping[parentId]; - if (!parentMapping?.message) { - return []; - } - - const contentType = parentMapping.message.content?.content_type; - - // If this is a thoughts message, extract the thinking content - if (contentType === 'thoughts') { - const thoughts = parentMapping.message.content.thoughts || []; - const thinkingText = thoughts - .map((t) => t.content || t.summary || '') - .filter(Boolean) - .join('\n\n'); - - if (thinkingText) { - return [{ type: 'think', think: thinkingText }]; - } - return []; - } - - // If this is reasoning_recap, look at its parent for thoughts - if (contentType === 'reasoning_recap') { - return findThinkingContent(parentMapping.parent, visited); - } - - return []; - }; + const lineage = createChatGptLineage(conv.mapping, messageMap); // Create and save messages using the mapped IDs const messages = []; @@ -554,7 +480,7 @@ function processConversation(conv, importBatchBuilder, requestUserId, defaultMod if (!newMessageId) { continue; } - const parentMessageId = findValidParent(mapping.parent); + const parentMessageId = lineage.findValidParent(mapping.parent); const messageText = formatMessageText(mapping.message); @@ -593,7 +519,7 @@ function processConversation(conv, importBatchBuilder, requestUserId, defaultMod // For assistant messages, check if there's thinking content in the parent chain if (!isCreatedByUser) { - const thinkingContent = findThinkingContent(mapping.parent); + const thinkingContent = lineage.findThinkingContent(mapping.parent); if (thinkingContent.length > 0) { // Combine thinking content with the text response message.content = [...thinkingContent, { type: 'text', text: messageText }]; @@ -603,10 +529,7 @@ function processConversation(conv, importBatchBuilder, requestUserId, defaultMod messages.push(message); } - const cycleDetected = adjustTimestampsForOrdering(messages); - if (cycleDetected) { - breakParentCycles(messages); - } + orderMessageLineage(messages); for (const message of messages) { importBatchBuilder.saveMessage(message); @@ -633,28 +556,7 @@ function processAssistantMessage(messageData, messageText) { return messageText; } - const citations = messageData.metadata?.citations ?? []; - - const sortedCitations = [...citations].sort((a, b) => b.start_ix - a.start_ix); - - let result = messageText; - for (const citation of sortedCitations) { - if ( - !citation.metadata?.type || - citation.metadata.type !== 'webpage' || - typeof citation.start_ix !== 'number' || - typeof citation.end_ix !== 'number' || - citation.start_ix >= citation.end_ix - ) { - continue; - } - - const replacement = ` ([${citation.metadata.title}](${citation.metadata.url}))`; - - result = result.slice(0, citation.start_ix) + replacement + result.slice(citation.end_ix); - } - - return result; + return linkChatGptCitations(messageText, messageData.metadata?.citations); } /** @@ -693,85 +595,4 @@ function formatMessageText(messageData) { return messageText; } -/** - * Adjusts message timestamps to ensure children always come after parents. - * Messages are sorted by createdAt and buildTree expects parents to appear before children. - * ChatGPT exports can have slight timestamp inversions (e.g., tool call results - * arriving a few ms before their parent). Uses multiple passes to handle cascading adjustments. - * Capped at N passes (where N = message count) to guarantee termination on cyclic graphs. - * - * @param {Array} messages - Array of message objects with messageId, parentMessageId, and createdAt. - * @returns {boolean} True if cyclic parent relationships were detected. - */ -function adjustTimestampsForOrdering(messages) { - if (messages.length === 0) { - return false; - } - - const timestampMap = new Map(); - for (const msg of messages) { - timestampMap.set(msg.messageId, msg.createdAt); - } - - let hasChanges = true; - let remainingPasses = messages.length; - while (hasChanges && remainingPasses > 0) { - hasChanges = false; - remainingPasses--; - for (const message of messages) { - if (message.parentMessageId && message.parentMessageId !== Constants.NO_PARENT) { - const parentTimestamp = timestampMap.get(message.parentMessageId); - if (parentTimestamp && message.createdAt <= parentTimestamp) { - message.createdAt = new Date(parentTimestamp.getTime() + 1); - timestampMap.set(message.messageId, message.createdAt); - hasChanges = true; - } - } - } - } - - const cycleDetected = remainingPasses === 0 && hasChanges; - if (cycleDetected) { - logger.warn( - '[importers] Detected cyclic parent relationships while adjusting import timestamps', - ); - } - return cycleDetected; -} - -/** - * Severs cyclic parentMessageId back-edges so saved messages form a valid tree. - * Walks each message's parent chain; if a message is visited twice, its parentMessageId - * is set to NO_PARENT to break the cycle. - * - * @param {Array} messages - Array of message objects with messageId and parentMessageId. - */ -function breakParentCycles(messages) { - const parentLookup = new Map(); - for (const msg of messages) { - parentLookup.set(msg.messageId, msg); - } - - const settled = new Set(); - for (const message of messages) { - const chain = new Set(); - let current = message; - while (current && !settled.has(current.messageId)) { - if (chain.has(current.messageId)) { - current.parentMessageId = Constants.NO_PARENT; - break; - } - chain.add(current.messageId); - const parentId = current.parentMessageId; - if (!parentId || parentId === Constants.NO_PARENT) { - break; - } - current = parentLookup.get(parentId); - } - for (const id of chain) { - settled.add(id); - } - } -} - module.exports = { getImporter, processAssistantMessage }; diff --git a/api/server/utils/import/importers.spec.js b/api/server/utils/import/importers.spec.js index aa71772de5f..b9d945af94c 100644 --- a/api/server/utils/import/importers.spec.js +++ b/api/server/utils/import/importers.spec.js @@ -1594,6 +1594,26 @@ describe('processAssistantMessage', () => { }); }); + test('should link tens of thousands of citations in one message', () => { + const count = 40000; + const marker = '【†】'; + const span = `${'word '.repeat(19)}${marker}`; + const citations = Array.from({ length: count }, (_, index) => ({ + start_ix: (index + 1) * span.length - marker.length, + end_ix: (index + 1) * span.length, + metadata: { type: 'webpage', title: 'Source', url: 'https://example.com' }, + })); + + const text = span.repeat(count); + const startedAt = performance.now(); + const result = processAssistantMessage({ metadata: { citations } }, text); + const elapsedMs = performance.now() - startedAt; + + expect(elapsedMs).toBeLessThan(1000); + expect(result).not.toContain(marker); + expect(result.split(' ([Source](https://example.com))')).toHaveLength(count + 1); + }); + test('should handle potential ReDoS attack payloads', () => { // Test with increasing input sizes to check for exponential behavior const sizes = [32, 33, 34]; // Adding more sizes would increase test time diff --git a/api/strategies/appleStrategy.js b/api/strategies/appleStrategy.js index 6eace87baea..398c3fd9ab3 100644 --- a/api/strategies/appleStrategy.js +++ b/api/strategies/appleStrategy.js @@ -1,6 +1,7 @@ const jwt = require('jsonwebtoken'); const { logger } = require('@librechat/data-schemas'); const { Strategy: AppleStrategy } = require('passport-apple'); +const { createOAuthStateStore, deferStateToStore } = require('@librechat/api'); const socialLogin = require('./socialLogin'); /** @@ -45,11 +46,21 @@ const getAppleConfig = (callbackURL) => ({ passReqToCallback: false, }); -const appleStrategy = () => - new AppleStrategy( - getAppleConfig(`${process.env.DOMAIN_SERVER}${process.env.APPLE_CALLBACK_URL}`), +/** + * Apple returns with a cross-site form POST, so its state cookie must be `SameSite=None`. + * @param {Omit} stateOptions + */ +const appleStrategy = (stateOptions) => { + const strategy = new AppleStrategy( + { + ...getAppleConfig(`${process.env.DOMAIN_SERVER}${process.env.APPLE_CALLBACK_URL}`), + store: createOAuthStateStore({ ...stateOptions, provider: 'apple', crossSiteCallback: true }), + }, appleLogin, ); + deferStateToStore(strategy); + return strategy; +}; const appleAdminStrategy = () => new AppleStrategy( diff --git a/api/strategies/discordStrategy.js b/api/strategies/discordStrategy.js index 7fb68280d5a..2f91996d37a 100644 --- a/api/strategies/discordStrategy.js +++ b/api/strategies/discordStrategy.js @@ -1,4 +1,5 @@ const { Strategy: DiscordStrategy } = require('passport-discord'); +const { createOAuthStateStore } = require('@librechat/api'); const socialLogin = require('./socialLogin'); const getProfileDetails = ({ profile }) => { @@ -32,9 +33,13 @@ const getDiscordConfig = (callbackURL) => ({ authorizationURL: 'https://discord.com/api/oauth2/authorize?prompt=none', }); -const discordStrategy = () => +/** @param {Omit} stateOptions */ +const discordStrategy = (stateOptions) => new DiscordStrategy( - getDiscordConfig(`${process.env.DOMAIN_SERVER}${process.env.DISCORD_CALLBACK_URL}`), + { + ...getDiscordConfig(`${process.env.DOMAIN_SERVER}${process.env.DISCORD_CALLBACK_URL}`), + store: createOAuthStateStore({ ...stateOptions, provider: 'discord' }), + }, discordLogin, ); diff --git a/api/strategies/facebookStrategy.js b/api/strategies/facebookStrategy.js index f638c3bfdb9..6191b13c6cb 100644 --- a/api/strategies/facebookStrategy.js +++ b/api/strategies/facebookStrategy.js @@ -1,4 +1,5 @@ const FacebookStrategy = require('passport-facebook').Strategy; +const { createOAuthStateStore } = require('@librechat/api'); const socialLogin = require('./socialLogin'); const getProfileDetails = ({ profile }) => ({ @@ -22,9 +23,13 @@ const getFacebookConfig = (callbackURL) => ({ profileFields: ['id', 'email', 'name'], }); -const facebookStrategy = () => +/** @param {Omit} stateOptions */ +const facebookStrategy = (stateOptions) => new FacebookStrategy( - getFacebookConfig(`${process.env.DOMAIN_SERVER}${process.env.FACEBOOK_CALLBACK_URL}`), + { + ...getFacebookConfig(`${process.env.DOMAIN_SERVER}${process.env.FACEBOOK_CALLBACK_URL}`), + store: createOAuthStateStore({ ...stateOptions, provider: 'facebook' }), + }, facebookLogin, ); diff --git a/api/strategies/githubStrategy.js b/api/strategies/githubStrategy.js index 363acbfcdb4..b0cb9ef6791 100644 --- a/api/strategies/githubStrategy.js +++ b/api/strategies/githubStrategy.js @@ -1,4 +1,5 @@ const { Strategy: GitHubStrategy } = require('passport-github2'); +const { createOAuthStateStore } = require('@librechat/api'); const socialLogin = require('./socialLogin'); const getProfileDetails = ({ profile }) => ({ @@ -30,9 +31,13 @@ const getGitHubConfig = (callbackURL) => ({ }), }); -const githubStrategy = () => +/** @param {Omit} stateOptions */ +const githubStrategy = (stateOptions) => new GitHubStrategy( - getGitHubConfig(`${process.env.DOMAIN_SERVER}${process.env.GITHUB_CALLBACK_URL}`), + { + ...getGitHubConfig(`${process.env.DOMAIN_SERVER}${process.env.GITHUB_CALLBACK_URL}`), + store: createOAuthStateStore({ ...stateOptions, provider: 'github' }), + }, githubLogin, ); diff --git a/api/strategies/googleStrategy.js b/api/strategies/googleStrategy.js index bee9a061a29..a616ef25026 100644 --- a/api/strategies/googleStrategy.js +++ b/api/strategies/googleStrategy.js @@ -1,4 +1,5 @@ const { Strategy: GoogleStrategy } = require('passport-google-oauth20'); +const { createOAuthStateStore } = require('@librechat/api'); const socialLogin = require('./socialLogin'); const getProfileDetails = ({ profile }) => ({ @@ -20,9 +21,13 @@ const getGoogleConfig = (callbackURL) => ({ proxy: true, }); -const googleStrategy = () => +/** @param {Omit} stateOptions */ +const googleStrategy = (stateOptions) => new GoogleStrategy( - getGoogleConfig(`${process.env.DOMAIN_SERVER}${process.env.GOOGLE_CALLBACK_URL}`), + { + ...getGoogleConfig(`${process.env.DOMAIN_SERVER}${process.env.GOOGLE_CALLBACK_URL}`), + store: createOAuthStateStore({ ...stateOptions, provider: 'google' }), + }, googleLogin, ); diff --git a/api/typedefs.js b/api/typedefs.js index 92839f48d4d..81bc2b2ea90 100644 --- a/api/typedefs.js +++ b/api/typedefs.js @@ -1114,6 +1114,18 @@ * @memberof typedefs */ +/** + * @exports BalanceReservation + * @typedef {import('@librechat/api').BalanceReservation} BalanceReservation + * @memberof typedefs + */ + +/** + * @exports BalanceReservations + * @typedef {import('@librechat/api').BalanceReservations} BalanceReservations + * @memberof typedefs + */ + /** * @exports Keyv * @typedef {import('keyv')} Keyv diff --git a/client/jest.config.cjs b/client/jest.config.cjs index 8c01907f403..21a86c42557 100644 --- a/client/jest.config.cjs +++ b/client/jest.config.cjs @@ -1,4 +1,6 @@ /** v0.8.8-rc3 */ +const { maxWorkers } = require('../config/jest.workers.cjs'); + module.exports = { roots: ['/src'], testEnvironment: 'jsdom', @@ -34,7 +36,7 @@ module.exports = { '^librechat-data-provider/react-query$': '/../node_modules/librechat-data-provider/src/react-query', }, - maxWorkers: '50%', + maxWorkers, /** Coverage maps accumulate for the life of a worker, so a long run can push * a worker past a gigabyte and get it killed by the OS, which fails whatever * suite it was holding. Recycling bloated workers also avoids swap thrash. */ diff --git a/client/src/Providers/ShareContext.tsx b/client/src/Providers/ShareContext.tsx index 74caf07cb98..7e1e4000602 100644 --- a/client/src/Providers/ShareContext.tsx +++ b/client/src/Providers/ShareContext.tsx @@ -1,5 +1,11 @@ import { createContext, useContext } from 'react'; -type TShareContext = { isSharedConvo?: boolean; shareId?: string }; +type TShareContext = { + isSharedConvo?: boolean; + shareId?: string; + /** Whether the link was published with a configured sender label. No conversation + * is in scope under a share link, so the header reads it here. */ + hasConfiguredSender?: boolean; +}; export const ShareContext = createContext({} as TShareContext); export const useShareContext = () => useContext(ShareContext); diff --git a/client/src/common/types.ts b/client/src/common/types.ts index d7e2a81a245..2c39712cbd1 100644 --- a/client/src/common/types.ts +++ b/client/src/common/types.ts @@ -491,7 +491,7 @@ export type ToolDialogProps = { }; export type TResError = { - response: { data: { message: string } }; + response: { data: { message: string; code?: string } }; message: string; }; @@ -676,8 +676,13 @@ export type TThread = { id: string; createdAt: string }; declare global { interface Window { google_tag_manager?: unknown; + /** Answers the server emits with the document, ahead of the app's own + * scripts, for questions the first render must not guess at. */ __LIBRECHAT_CONFIG__?: { enableQueryDevtools?: boolean; + /** Whether this deployment configured footer content of its own, so the + * composer reserves the footer bar's band on its first frame. */ + hasConfiguredFooter?: boolean; }; } } diff --git a/client/src/components/Auth/TwoFactorScreen.tsx b/client/src/components/Auth/TwoFactorScreen.tsx index 0166f897709..7cd073eca5c 100644 --- a/client/src/components/Auth/TwoFactorScreen.tsx +++ b/client/src/components/Auth/TwoFactorScreen.tsx @@ -1,6 +1,7 @@ import React, { useState, useCallback } from 'react'; import { useSearchParams } from 'react-router-dom'; import { useToastContext } from '@librechat/client'; +import { ErrorTypes } from 'librechat-data-provider'; import { useForm, Controller } from 'react-hook-form'; import { REGEXP_ONLY_DIGITS, REGEXP_ONLY_DIGITS_AND_CHARS } from 'input-otp'; import { @@ -50,11 +51,13 @@ const TwoFactorScreen: React.FC = React.memo(() => { }, onError: (error: unknown) => { setIsLoading(false); - const err = error as { response?: { data?: { message?: unknown } } }; - const errorMsg = - typeof err.response?.data?.message === 'string' - ? err.response.data.message - : 'Error verifying 2FA'; + const data = (error as { response?: { data?: { message?: unknown; code?: unknown } } }) + .response?.data; + if (data?.code === ErrorTypes.AUTH_CROSS_ORIGIN) { + showToast({ message: localize('com_auth_error_login_cross_origin'), status: 'error' }); + return; + } + const errorMsg = typeof data?.message === 'string' ? data.message : 'Error verifying 2FA'; showToast({ message: errorMsg, status: 'error' }); }, }); diff --git a/client/src/components/Auth/__tests__/TwoFactorScreen.spec.tsx b/client/src/components/Auth/__tests__/TwoFactorScreen.spec.tsx new file mode 100644 index 00000000000..215df2ced87 --- /dev/null +++ b/client/src/components/Auth/__tests__/TwoFactorScreen.spec.tsx @@ -0,0 +1,70 @@ +import { MemoryRouter } from 'react-router-dom'; +import { act, render } from '@testing-library/react'; +import { ErrorTypes } from 'librechat-data-provider'; +import TwoFactorScreen from '../TwoFactorScreen'; + +const mockShowToast = jest.fn(); +let mockVerifyOptions: { onError: (error: unknown) => void } | undefined; + +jest.mock('@librechat/client', () => ({ + ...jest.requireActual('@librechat/client'), + useToastContext: () => ({ showToast: mockShowToast }), +})); + +jest.mock('~/hooks', () => ({ + useLocalize: () => (key: string) => key, +})); + +jest.mock('~/data-provider', () => ({ + useVerifyTwoFactorTempMutation: (options: { onError: (error: unknown) => void }) => { + mockVerifyOptions = options; + return { mutate: jest.fn() }; + }, +})); + +function renderScreen() { + render( + + + , + ); +} + +describe('TwoFactorScreen verification errors', () => { + beforeEach(() => { + jest.clearAllMocks(); + mockVerifyOptions = undefined; + }); + + it('shows a localized message when the server rejects a cross-origin submission', () => { + renderScreen(); + + act(() => { + mockVerifyOptions?.onError({ + response: { + data: { message: 'Cross-site request rejected', code: ErrorTypes.AUTH_CROSS_ORIGIN }, + }, + }); + }); + + expect(mockShowToast).toHaveBeenCalledWith({ + message: 'com_auth_error_login_cross_origin', + status: 'error', + }); + }); + + it('keeps showing the server message for other verification failures', () => { + renderScreen(); + + act(() => { + mockVerifyOptions?.onError({ + response: { data: { message: 'Invalid 2FA code or backup code' } }, + }); + }); + + expect(mockShowToast).toHaveBeenCalledWith({ + message: 'Invalid 2FA code or backup code', + status: 'error', + }); + }); +}); diff --git a/client/src/components/Chat/ChatView.tsx b/client/src/components/Chat/ChatView.tsx index 24186fa708f..8e127931f5b 100644 --- a/client/src/components/Chat/ChatView.tsx +++ b/client/src/components/Chat/ChatView.tsx @@ -59,7 +59,7 @@ function ChatView({ index = 0, project }: { index?: number; project?: TChatProje /** A conversation carries a footer only for configured content, and the * composer's clearance has to account for the bar when it does — including - * while the config is still in flight, so a cold load does not jump. */ + * before the config answers, so a cold load does not jump. */ const configuredFooter = useConfiguredFooter(); const methods = useForm({ @@ -114,10 +114,10 @@ function ChatView({ index = 0, project }: { index?: number; project?: TChatProje (conversationId === Constants.NEW_CONVO || !conversationId); /** A footer bar renders beneath the composer on the welcome screen always, and - * in a conversation when the deployment configured one. `present` already - * carries the remembered answer while the config is in flight, so this is the - * same value before and after it resolves. */ - const footerBelow = isLandingPage || configuredFooter.present; + * in a conversation when the deployment configured one. The shell already + * carried that answer, so this is the same value before and after the config + * resolves. */ + const footerBelow = isLandingPage || configuredFooter; const isNavigating = (!messagesTree || messagesTree.length === 0) && conversationId != null; const isProjectLandingPage = isLandingPage && project != null; @@ -214,7 +214,7 @@ function ChatView({ index = 0, project }: { index?: number; project?: TChatProje {/* The generic disclaimer is the welcome screen's; a deployment's own footer, privacy policy and terms stay with the conversation that always showed them. */} - {!isLandingPage && configuredFooter.present &&