📻 fix: Replay MCP OAuth Prompts for Coalesced Connections (#13565)

* fix: Replay MCP OAuth URL for Joined Connections

* chore: Sort MCP OAuth Imports

* test: Restore MCP OAuth Registry Spies

* fix: Replay pending MCP OAuth prompts

* fix: Replay MCP OAuth on Stream Resume

* fix: Preserve MCP OAuth Replay Context

* chore: Format MCP OAuth Replay Context

* test: Expect MCP OAuth Replay Expiry

* fix: Render pending MCP OAuth prompts

* chore: Clean MCP OAuth Replay Type Narrowing

* fix: Stabilize new MCP OAuth chats

* fix: Re-emit cached MCP OAuth prompts

* fix: Replay pending OAuth for selected MCP tools

* fix: Avoid stalling pending MCP OAuth replay

* test: Clean MCP OAuth review findings

* test: Restore MCP OAuth registry spy

* fix: Resolve OAuth Typecheck Regressions

* fix: Harden MCP OAuth replay edge cases

* test: Cover MCP OAuth joined prompt expiry

* test: Mark joined OAuth replay fixture

* test: Use OAuth fixture for joined replay expiry

* fix: Anchor resumed MCP OAuth prompts

* fix: Seed resumable turn metadata before MCP init

* test: Format resume metadata regression

* fix: Prioritize resumable stream routes

* fix: Preserve MCP OAuth resume message tree

* test: Fix MCP OAuth Resume Test Types

* fix: Replay MCP OAuth Regenerate Prompts

* fix: Skip OAuth-only Abort Persistence

* fix: Stabilize OAuth Resume Replay

* fix: Target Non-Tail Regenerate Responses

* fix: Scope Regenerate Step Updates

* fix: Clean Up OAuth Abort State

* fix: Preserve Regenerate Branch Siblings

* fix: Preserve OAuth Resume Branch State

* fix: Preserve OAuth Branch Resume State

* chore: Sort OAuth Resume Imports

* fix: Address OAuth Resume Review Findings

* test: Fix Abort Fixture Typing
This commit is contained in:
Danny Avila 2026-06-07 10:45:54 -04:00 • committed by GitHub
parent 6950448d03
commit cb1d536874
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
48 changed files with 4837 additions and 299 deletions

View file

@ -0,0 +1,197 @@
const mockLogger = {
debug: jest.fn(),
warn: jest.fn(),
error: jest.fn(),
info: jest.fn(),
};
const mockGenerationJobManager = {
createJob: jest.fn(),
emitError: jest.fn(),
completeJob: jest.fn(),
getResumeState: jest.fn(),
updateMetadata: jest.fn(),
};
const mockCheckAndIncrementPendingRequest = jest.fn();
const mockDecrementPendingRequest = jest.fn();
const mockFilterPersistableAbortContent = jest.fn((content) =>
content.filter((part) => part?.type !== 'tool_call'),
);
const mockGetConvo = jest.fn();
const mockSaveMessage = jest.fn();
jest.mock('@librechat/data-schemas', () => ({
logger: mockLogger,
}));
jest.mock('@librechat/api', () => ({
sendEvent: jest.fn(),
getViolationInfo: jest.fn(),
buildMessageFiles: jest.fn(() => []),
resolveTitleTiming: jest.fn(() => 'immediate'),
GenerationJobManager: mockGenerationJobManager,
filterPersistableAbortContent: (...args) => mockFilterPersistableAbortContent(...args),
decrementPendingRequest: (...args) => mockDecrementPendingRequest(...args),
sanitizeMessageForTransmit: jest.fn((message) => message),
checkAndIncrementPendingRequest: (...args) => mockCheckAndIncrementPendingRequest(...args),
}));
jest.mock('~/server/cleanup', () => ({
disposeClient: jest.fn(),
clientRegistry: null,
requestDataMap: {
set: jest.fn(),
},
}));
jest.mock('~/server/middleware', () => ({
handleAbortError: jest.fn(() => Promise.resolve()),
}));
jest.mock('~/cache', () => ({
logViolation: jest.fn(),
}));
jest.mock('~/models', () => ({
saveMessage: (...args) => mockSaveMessage(...args),
getConvo: (...args) => mockGetConvo(...args),
}));
const AgentController = require('../request');
describe('ResumableAgentController resume metadata', () => {
beforeEach(() => {
jest.clearAllMocks();
mockCheckAndIncrementPendingRequest.mockResolvedValue({ allowed: true });
mockDecrementPendingRequest.mockResolvedValue(undefined);
mockGetConvo.mockResolvedValue({ createdAt: '2026-06-07T00:00:00.000Z' });
mockGenerationJobManager.createJob.mockResolvedValue({
createdAt: 1000,
readyPromise: Promise.resolve(),
abortController: new AbortController(),
emitter: { on: jest.fn() },
});
mockGenerationJobManager.getResumeState.mockResolvedValue(null);
mockGenerationJobManager.updateMetadata.mockResolvedValue(undefined);
mockGenerationJobManager.emitError.mockResolvedValue(undefined);
mockSaveMessage.mockResolvedValue({});
});
it('stores the in-flight turn before MCP initialization can emit OAuth', async () => {
const conversationId = 'conversation-123';
const initializeClient = jest.fn().mockRejectedValue(new Error('stop before tool loading'));
const req = {
user: { id: 'user-123' },
body: {
text: 'Check Google Workspace availability.',
messageId: 'follow-up-user',
parentMessageId: 'original-response',
conversationId,
endpointOption: {
endpoint: 'agents',
modelOptions: { model: 'gpt-3.5-turbo' },
},
},
config: {},
};
const res = {
headersSent: true,
json: jest.fn(() => {
res.headersSent = true;
}),
status: jest.fn(() => res),
};
await AgentController(req, res, jest.fn(), initializeClient, null);
expect(mockGenerationJobManager.updateMetadata).toHaveBeenCalledWith(conversationId, {
conversationId,
responseMessageId: 'follow-up-user_',
userMessage: {
messageId: 'follow-up-user',
parentMessageId: 'original-response',
conversationId,
text: 'Check Google Workspace availability.',
},
});
expect(mockGenerationJobManager.updateMetadata.mock.invocationCallOrder[0]).toBeLessThan(
initializeClient.mock.invocationCallOrder[0],
);
});
it('filters OAuth prompts before saving partial responses on disconnect', async () => {
const conversationId = 'conversation-123';
let allSubscribersLeftHandler;
mockGenerationJobManager.createJob.mockResolvedValue({
createdAt: 1000,
readyPromise: Promise.resolve(),
abortController: new AbortController(),
emitter: {
on: jest.fn((event, handler) => {
if (event === 'allSubscribersLeft') {
allSubscribersLeftHandler = handler;
}
}),
},
});
mockGenerationJobManager.getResumeState.mockResolvedValue({
conversationId,
responseMessageId: 'response-message',
userMessage: {
messageId: 'user-message',
parentMessageId: 'parent-message',
conversationId,
text: 'Use Google Workspace',
},
});
const initializeClient = jest.fn().mockRejectedValue(new Error('stop after setup'));
const req = {
user: { id: 'user-123' },
body: {
text: 'Use Google Workspace',
messageId: 'user-message',
parentMessageId: 'parent-message',
conversationId,
endpointOption: {
endpoint: 'agents',
modelOptions: { model: 'gpt-3.5-turbo' },
},
},
config: {},
};
const res = {
headersSent: true,
json: jest.fn(() => {
res.headersSent = true;
}),
status: jest.fn(() => res),
};
await AgentController(req, res, jest.fn(), initializeClient, null);
expect(allSubscribersLeftHandler).toEqual(expect.any(Function));
const oauthPart = {
type: 'tool_call',
tool_call: {
name: 'oauth_mcp_Google-Workspace',
auth: 'https://auth.example.com/oauth',
},
};
const textPart = { type: 'text', text: 'Partial response...' };
await allSubscribersLeftHandler([oauthPart, textPart]);
expect(mockFilterPersistableAbortContent).toHaveBeenCalledWith([oauthPart, textPart]);
expect(mockSaveMessage).toHaveBeenCalledWith(
expect.objectContaining({ userId: 'user-123' }),
expect.objectContaining({
content: [textPart],
messageId: 'response-message',
parentMessageId: 'user-message',
}),
expect.any(Object),
);
});
});

View file

@ -6,6 +6,7 @@ const {
buildMessageFiles,
resolveTitleTiming,
GenerationJobManager,
filterPersistableAbortContent,
decrementPendingRequest,
sanitizeMessageForTransmit,
checkAndIncrementPendingRequest,
@ -75,6 +76,31 @@ async function attachConversationCreatedAt(req, { userId, conversationId, isNewC
}
}
function getPreliminaryResponseMessageId({ messageId, responseMessageId }) {
if (typeof responseMessageId === 'string' && responseMessageId.length > 0) {
return responseMessageId;
}
if (typeof messageId !== 'string' || messageId.length === 0) {
return null;
}
return `${messageId.replace(/_+$/, '')}_`;
}
function getPreliminaryUserMessage({ messageId, parentMessageId, text }, conversationId) {
if (typeof messageId !== 'string' || messageId.length === 0) {
return null;
}
return {
messageId,
parentMessageId,
conversationId,
text,
};
}
/**
* Resumable Agent Controller - Generation runs independently of HTTP connection.
* Returns streamId immediately, client subscribes separately via SSE.
@ -134,6 +160,16 @@ const ResumableAgentController = async (req, res, next, initializeClient, addTit
await attachConversationCreatedAt(req, { userId, conversationId, isNewConvo });
const preliminaryUserMessage = getPreliminaryUserMessage(req.body, conversationId);
const preliminaryResponseMessageId = getPreliminaryResponseMessageId(req.body);
if (preliminaryUserMessage || preliminaryResponseMessageId) {
await GenerationJobManager.updateMetadata(streamId, {
conversationId,
responseMessageId: preliminaryResponseMessageId,
userMessage: preliminaryUserMessage,
});
}
// Note: We no longer use res.on('close') to abort since we send JSON immediately.
// The response closes normally after res.json(), which is not an abort condition.
// Abort handling is done through GenerationJobManager via the SSE stream connection.
@ -155,6 +191,12 @@ const ResumableAgentController = async (req, res, next, initializeClient, addTit
return;
}
const persistableContent = filterPersistableAbortContent(aggregatedContent);
if (persistableContent.length === 0) {
logger.debug('[ResumableAgentController] No persistable content to save partial response');
return;
}
const resumeState = await GenerationJobManager.getResumeState(streamId);
if (!resumeState?.userMessage) {
logger.debug('[ResumableAgentController] No user message to save partial response for');
@ -170,7 +212,7 @@ const ResumableAgentController = async (req, res, next, initializeClient, addTit
conversationId: responseConversationId,
parentMessageId: resumeState.userMessage.messageId,
sender: client?.sender ?? 'AI',
content: aggregatedContent,
content: persistableContent,
unfinished: true,
error: false,
isCreatedByUser: false,
@ -194,7 +236,7 @@ const ResumableAgentController = async (req, res, next, initializeClient, addTit
);
logger.debug(
`[ResumableAgentController] Saved partial response for ${streamId}, content parts: ${aggregatedContent.length}`,
`[ResumableAgentController] Saved partial response for ${streamId}, content parts: ${persistableContent.length}`,
);
} catch (error) {
logger.error('[ResumableAgentController] Error saving partial response:', error);

View file

@ -195,6 +195,43 @@ describe('Agent Abort Endpoint', () => {
expect(response.status).toBe(200);
expect(mockSaveMessage).not.toHaveBeenCalled();
});
it('should skip message saving when abort content is only an OAuth prompt', async () => {
const jobStreamId = 'test-stream-123';
mockGenerationJobManager.getJob.mockResolvedValue({
metadata: { userId: 'test-user-123' },
});
mockGenerationJobManager.abortJob.mockResolvedValue({
success: true,
jobData: {
userMessage: { messageId: 'user-msg-123' },
responseMessageId: 'response-msg-456',
conversationId: jobStreamId,
},
content: [
{
type: 'tool_call',
tool_call: {
type: 'tool_call',
id: 'oauth-call-1',
name: 'oauth_mcp_Google-Workspace',
args: '',
auth: 'https://auth.example.com/oauth',
},
},
],
text: '',
});
const response = await request(app)
.post('/api/agents/chat/abort')
.send({ conversationId: jobStreamId });
expect(response.status).toBe(200);
expect(mockSaveMessage).not.toHaveBeenCalled();
});
});
describe('Partial Response Saving', () => {
@ -264,8 +301,8 @@ describe('Agent Abort Endpoint', () => {
responseMessageId: 'response-msg-456',
conversationId: jobStreamId,
},
content: [],
text: '',
content: [{ type: 'text', text: 'Partial response...' }],
text: 'Partial response...',
});
mockSaveMessage.mockRejectedValue(new Error('Database error'));

View file

@ -45,9 +45,11 @@ jest.mock('~/server/middleware', () => ({
}));
jest.mock('~/server/routes/agents/chat', () => require('express').Router());
jest.mock('~/server/routes/agents/v1', () => ({
v1: require('express').Router(),
}));
jest.mock('~/server/routes/agents/v1', () => {
const router = require('express').Router();
router.use((req, res) => res.status(418).json({ error: 'v1 caught stream route' }));
return { v1: router };
});
jest.mock('~/server/routes/agents/openai', () => require('express').Router());
jest.mock('~/server/routes/agents/responses', () => require('express').Router());

View file

@ -1,5 +1,5 @@
const express = require('express');
const { isEnabled, GenerationJobManager } = require('@librechat/api');
const { isEnabled, GenerationJobManager, hasPersistableAbortContent } = require('@librechat/api');
const { createSseStreamTelemetry } = require('@librechat/api/telemetry');
const { logger } = require('@librechat/data-schemas');
const {
@ -43,8 +43,6 @@ router.use(requireJwtAuth);
router.use(checkBan);
router.use(uaParser);
router.use('/', v1);
/**
* Stream endpoints - mounted before chatRouter to bypass rate limiters
* These are GET requests and don't need message body validation or rate limiting
@ -273,7 +271,8 @@ router.post('/chat/abort', async (req, res) => {
if (
abortResult.success &&
abortResult.jobData?.userMessage?.messageId &&
abortResult.jobData?.responseMessageId
abortResult.jobData?.responseMessageId &&
hasPersistableAbortContent(abortResult.content)
) {
const { jobData, content, text } = abortResult;
const responseMessage = {
@ -314,6 +313,8 @@ router.post('/chat/abort', async (req, res) => {
return res.status(404).json({ error: 'Job not found', streamId: jobStreamId });
});
router.use('/', v1);
const chatRouter = express.Router();
chatRouter.use(configMiddleware);

View file

@ -1,11 +1,6 @@
const { tool } = require('@librechat/agents/langchain/tools');
const { logger, getTenantId } = require('@librechat/data-schemas');
const {
Providers,
StepTypes,
GraphEvents,
Constants: AgentConstants,
} = require('@librechat/agents');
const { Providers, Constants: AgentConstants } = require('@librechat/agents');
const {
sendEvent,
MCPOAuthHandler,
@ -14,7 +9,11 @@ const {
normalizeJsonSchema,
GenerationJobManager,
resolveJsonSchemaRefs,
buildOAuthToolCallName,
buildMCPAuthStepId,
buildMCPAuthToolCall,
buildMCPAuthRunStepEvent,
buildMCPAuthRunStepDeltaEvent,
buildMCPAuthRunStepEndDeltaEvent,
checkAccessWithRequestCache,
} = require('@librechat/api');
const {
@ -201,20 +200,11 @@ function isEmptyObjectSchema(jsonSchema) {
function createRunStepDeltaEmitter({ res, stepId, toolCall, streamId = null }) {
/**
* @param {string} authURL - The URL to redirect the user for OAuth authentication.
* @param {{ expiresAt?: number }} [options]
* @returns {Promise<void>}
*/
return async function (authURL) {
/** @type {{ id: string; delta: AgentToolCallDelta }} */
const data = {
id: stepId,
delta: {
type: StepTypes.TOOL_CALLS,
tool_calls: [{ ...toolCall, args: '' }],
auth: authURL,
expires_at: Date.now() + Time.TWO_MINUTES,
},
};
const eventData = { event: GraphEvents.ON_RUN_STEP_DELTA, data };
return async function (authURL, options) {
const eventData = buildMCPAuthRunStepDeltaEvent({ authURL, stepId, toolCall, options });
if (streamId) {
await GenerationJobManager.emitChunk(streamId, eventData);
} else {
@ -235,18 +225,7 @@ function createRunStepDeltaEmitter({ res, stepId, toolCall, streamId = null }) {
*/
function createRunStepEmitter({ res, runId, stepId, toolCall, index, streamId = null }) {
return async function () {
/** @type {import('@librechat/agents').RunStep} */
const data = {
runId: runId ?? Constants.USE_PRELIM_RESPONSE_MESSAGE_ID,
id: stepId,
type: StepTypes.TOOL_CALLS,
index: index ?? 0,
stepDetails: {
type: StepTypes.TOOL_CALLS,
tool_calls: [toolCall],
},
};
const eventData = { event: GraphEvents.ON_RUN_STEP, data };
const eventData = buildMCPAuthRunStepEvent({ runId, stepId, toolCall, index });
if (streamId) {
await GenerationJobManager.emitChunk(streamId, eventData);
} else {
@ -260,20 +239,43 @@ function createRunStepEmitter({ res, runId, stepId, toolCall, index, streamId =
* @param {object} params
* @param {string} params.flowId - The ID of the login flow.
* @param {FlowStateManager<any>} params.flowManager - The flow manager instance.
* @param {(authURL: string) => void} [params.callback]
* @param {(authURL: string, options?: { expiresAt?: number }) => void | Promise<void>} [params.callback]
*/
function createOAuthStart({ flowId, flowManager, callback }) {
/**
* Creates a function to handle OAuth login requests.
* @param {string} authURL - The URL to redirect the user for OAuth authentication.
* @param {{ expiresAt?: number }} [options]
* @returns {Promise<boolean>} Returns true to indicate the event was sent successfully.
*/
return async function (authURL) {
return async function (authURL, options) {
let emitted = false;
const emitOAuthStart = async (message) => {
if (options) {
await callback?.(authURL, options);
} else {
await callback?.(authURL);
}
emitted = true;
logger.debug(message);
};
const existingFlow = await flowManager.getFlowState(flowId, 'oauth_login');
if (existingFlow) {
await emitOAuthStart('Re-sent OAuth login request to client');
return true;
}
await flowManager.createFlowWithHandler(flowId, 'oauth_login', async () => {
callback?.(authURL);
logger.debug('Sent OAuth login request to client');
await emitOAuthStart('Sent OAuth login request to client');
return true;
});
if (!emitted) {
await emitOAuthStart('Re-sent OAuth login request to client');
}
return true;
};
}
@ -286,15 +288,7 @@ function createOAuthStart({ flowId, flowManager, callback }) {
*/
function createOAuthEnd({ res, stepId, toolCall, streamId = null }) {
return async function () {
/** @type {{ id: string; delta: AgentToolCallDelta }} */
const data = {
id: stepId,
delta: {
type: StepTypes.TOOL_CALLS,
tool_calls: [{ ...toolCall }],
},
};
const eventData = { event: GraphEvents.ON_RUN_STEP_DELTA, data };
const eventData = buildMCPAuthRunStepEndDeltaEvent({ stepId, toolCall });
if (streamId) {
await GenerationJobManager.emitChunk(streamId, eventData);
} else {
@ -323,14 +317,14 @@ function createAbortHandler({ userId, serverName, toolName, flowManager }) {
/**
* @param {Object} params
* @param {() => void} params.runStepEmitter
* @param {(authURL: string) => void} params.runStepDeltaEmitter
* @returns {(authURL: string) => void}
* @param {() => Promise<void>} params.runStepEmitter
* @param {(authURL: string, options?: { expiresAt?: number }) => Promise<void>} params.runStepDeltaEmitter
* @returns {(authURL: string, options?: { expiresAt?: number }) => Promise<void>}
*/
function createOAuthCallback({ runStepEmitter, runStepDeltaEmitter }) {
return function (authURL) {
runStepEmitter();
runStepDeltaEmitter(authURL);
return async function (authURL, options) {
await runStepEmitter();
await runStepDeltaEmitter(authURL, options);
};
}
@ -373,12 +367,11 @@ async function reconnectServer({
const runId = Constants.USE_PRELIM_RESPONSE_MESSAGE_ID;
const flowId = `${user.id}:${serverName}:${Date.now()}`;
const flowManager = getFlowStateManager(getLogStores(CacheKeys.FLOWS));
const stepId = 'step_oauth_login_' + serverName;
const toolCall = {
const stepId = buildMCPAuthStepId(serverName);
const toolCall = buildMCPAuthToolCall({
id: flowId,
name: buildOAuthToolCallName(serverName),
type: 'tool_call_chunk',
};
serverName,
});
// Set up abort handler to clean up OAuth flows if request is aborted
const oauthFlowId = MCPOAuthHandler.generateFlowId(user.id, serverName);
@ -949,6 +942,7 @@ module.exports = {
resolveConfigServers,
resolveMcpConfigNames,
resolveAllMcpConfigs,
createOAuthStart,
checkOAuthFlowStatus,
getServerConnectionStatus,
createUnavailableToolStub,

View file

@ -45,6 +45,7 @@ const {
createMCPTools,
createMCPPermissionContext,
getMCPSetupData,
createOAuthStart,
checkOAuthFlowStatus,
getServerConnectionStatus,
createUnavailableToolStub,
@ -100,6 +101,85 @@ describe('tests for the new helper functions used by the MCP connection status e
mockGetOAuthReconnectionManager = require('~/config').getOAuthReconnectionManager;
});
describe('createOAuthStart', () => {
const flowId = 'test-server:oauth_login:thread-1:run-1';
const authUrl = 'https://auth.example.com/oauth?state=test';
it('should create a login flow and emit the OAuth URL for the first request', async () => {
const callback = jest.fn();
const mockFlowManager = {
getFlowState: jest.fn().mockResolvedValue(null),
createFlowWithHandler: jest.fn(async (_flowId, _type, handler) => handler()),
};
const oauthStart = createOAuthStart({
flowId,
flowManager: mockFlowManager,
callback,
});
await expect(oauthStart(authUrl)).resolves.toBe(true);
expect(mockFlowManager.getFlowState).toHaveBeenCalledWith(flowId, 'oauth_login');
expect(mockFlowManager.createFlowWithHandler).toHaveBeenCalledWith(
flowId,
'oauth_login',
expect.any(Function),
);
expect(callback).toHaveBeenCalledWith(authUrl);
expect(logger.debug).toHaveBeenCalledWith('Sent OAuth login request to client');
});
it('should replay the OAuth URL when the login flow already exists', async () => {
const callback = jest.fn();
const mockFlowManager = {
getFlowState: jest.fn().mockResolvedValue({
status: 'COMPLETED',
result: true,
}),
createFlowWithHandler: jest.fn(),
};
const oauthStart = createOAuthStart({
flowId,
flowManager: mockFlowManager,
callback,
});
await expect(oauthStart(authUrl)).resolves.toBe(true);
expect(mockFlowManager.getFlowState).toHaveBeenCalledWith(flowId, 'oauth_login');
expect(mockFlowManager.createFlowWithHandler).not.toHaveBeenCalled();
expect(callback).toHaveBeenCalledWith(authUrl);
expect(logger.debug).toHaveBeenCalledWith('Re-sent OAuth login request to client');
});
it('should replay the OAuth URL when flow creation is deduped internally', async () => {
const callback = jest.fn();
const mockFlowManager = {
getFlowState: jest.fn().mockResolvedValue(null),
createFlowWithHandler: jest.fn().mockResolvedValue(true),
};
const oauthStart = createOAuthStart({
flowId,
flowManager: mockFlowManager,
callback,
});
await expect(oauthStart(authUrl)).resolves.toBe(true);
expect(mockFlowManager.getFlowState).toHaveBeenCalledWith(flowId, 'oauth_login');
expect(mockFlowManager.createFlowWithHandler).toHaveBeenCalledWith(
flowId,
'oauth_login',
expect.any(Function),
);
expect(callback).toHaveBeenCalledWith(authUrl);
expect(logger.debug).toHaveBeenCalledWith('Re-sent OAuth login request to client');
});
});
describe('getMCPSetupData', () => {
const mockUserId = 'user-123';
const mockConfig = {

View file

@ -2,8 +2,6 @@ const { logger } = require('@librechat/data-schemas');
const { tool: toolFn, DynamicStructuredTool } = require('@librechat/agents/langchain/tools');
const {
sleep,
StepTypes,
GraphEvents,
createToolSearch,
createBashExecutionTool,
Constants: AgentConstants,
@ -18,11 +16,17 @@ const {
isActionDomainAllowed,
buildWebSearchContext,
buildImageToolContext,
buildOAuthToolCallName,
buildToolClassification,
getMissingCustomUserVars,
buildWebSearchDynamicContext,
getCodeApiAuthHeaders,
getReplayablePendingMCPOAuthStart,
getMCPServerNamesFromTools,
buildMCPAuthToolCall,
buildMCPAuthStepId,
buildMCPAuthRunStepEvent,
buildMCPAuthRunStepDeltaEvent,
buildMCPAuthRunStepCompletedEvent,
isFileAuthoringToolDefinition,
} = require('@librechat/api');
const {
@ -596,44 +600,35 @@ async function loadToolDefinitionsWrapper({ req, res, agent, streamId = null, to
const flowManager = getFlowStateManager(flowsCache);
const configServers = await resolveConfigServers(req);
const pendingOAuthServers = new Set();
const pendingOAuthStarts = new Map();
const emittedOAuthStarts = new Map();
const oauthToolCallIds = new Map();
const oauthStepIndexes = new Map();
const createOAuthEmitter = (serverName, index) => {
return async (authURL) => {
const flowId = `${req.user.id}:${serverName}:${Date.now()}`;
const stepId = 'step_oauth_login_' + serverName;
return async (authURL, options) => {
if (emittedOAuthStarts.get(serverName) === authURL) {
return;
}
emittedOAuthStarts.set(serverName, authURL);
const flowId =
oauthToolCallIds.get(serverName) ?? `${req.user.id}:${serverName}:${Date.now()}`;
const stepId = buildMCPAuthStepId(serverName);
oauthToolCallIds.set(serverName, flowId);
oauthStepIndexes.set(serverName, index);
const toolCall = {
const toolCall = buildMCPAuthToolCall({
id: flowId,
name: buildOAuthToolCallName(serverName),
type: 'tool_call_chunk',
};
serverName,
});
const runStepData = {
runId: Constants.USE_PRELIM_RESPONSE_MESSAGE_ID,
id: stepId,
type: StepTypes.TOOL_CALLS,
index,
stepDetails: {
type: StepTypes.TOOL_CALLS,
tool_calls: [toolCall],
},
};
const runStepDeltaData = {
id: stepId,
delta: {
type: StepTypes.TOOL_CALLS,
tool_calls: [{ ...toolCall, args: '' }],
auth: authURL,
expires_at: Date.now() + Time.TWO_MINUTES,
},
};
const runStepEvent = { event: GraphEvents.ON_RUN_STEP, data: runStepData };
const runStepDeltaEvent = { event: GraphEvents.ON_RUN_STEP_DELTA, data: runStepDeltaData };
const runStepEvent = buildMCPAuthRunStepEvent({ stepId, toolCall, index });
const runStepDeltaEvent = buildMCPAuthRunStepDeltaEvent({
authURL,
stepId,
toolCall,
options,
});
if (streamId) {
await GenerationJobManager.emitChunk(streamId, runStepEvent);
@ -651,25 +646,19 @@ async function loadToolDefinitionsWrapper({ req, res, agent, streamId = null, to
const createOAuthEndEmitter = (serverName) => {
return async () => {
const stepId = 'step_oauth_login_' + serverName;
const toolCall = {
const stepId = buildMCPAuthStepId(serverName);
const toolCall = buildMCPAuthToolCall({
id: oauthToolCallIds.get(serverName),
name: buildOAuthToolCallName(serverName),
args: '',
output: 'OAuth authentication completed',
serverName,
type: 'tool_call',
};
const runStepCompletedEvent = {
event: GraphEvents.ON_RUN_STEP_COMPLETED,
data: {
result: {
id: stepId,
index: oauthStepIndexes.get(serverName) ?? 0,
tool_call: toolCall,
},
},
};
});
const runStepCompletedEvent = buildMCPAuthRunStepCompletedEvent({
stepId,
toolCall,
index: oauthStepIndexes.get(serverName) ?? 0,
});
if (streamId) {
await GenerationJobManager.emitChunk(streamId, runStepCompletedEvent);
@ -683,7 +672,45 @@ async function loadToolDefinitionsWrapper({ req, res, agent, streamId = null, to
};
};
const getPendingOAuthStartForEmit = async (serverName) => {
const cachedOAuthStart = pendingOAuthStarts.get(serverName);
if (cachedOAuthStart?.options?.expiresAt != null) {
return cachedOAuthStart;
}
const pendingOAuthStart = await getReplayablePendingMCPOAuthStart({
flowManager,
userId: req.user.id,
serverName,
});
if (!pendingOAuthStart) {
return cachedOAuthStart;
}
if (!cachedOAuthStart || pendingOAuthStart.authURL === cachedOAuthStart.authURL) {
pendingOAuthStarts.set(serverName, pendingOAuthStart);
return pendingOAuthStart;
}
return cachedOAuthStart;
};
const getOrFetchMCPServerTools = async (userId, serverName) => {
const addPendingOAuthServer = async () => {
const pendingOAuthStart = await getReplayablePendingMCPOAuthStart({
flowManager,
userId,
serverName,
});
if (!pendingOAuthStart) {
return false;
}
pendingOAuthServers.add(serverName);
pendingOAuthStarts.set(serverName, pendingOAuthStart);
return true;
};
let serverConfig;
try {
serverConfig =
@ -718,11 +745,19 @@ async function loadToolDefinitionsWrapper({ req, res, agent, streamId = null, to
const cached = await getMCPServerTools(userId, serverName);
if (cached) {
await addPendingOAuthServer();
return cached;
}
const oauthStart = async () => {
if (await addPendingOAuthServer()) {
return null;
}
const oauthStart = async (authURL, options) => {
pendingOAuthServers.add(serverName);
if (typeof authURL === 'string' && authURL.length > 0) {
pendingOAuthStarts.set(serverName, { authURL, options });
}
};
const result = await reinitMCPServer({
@ -813,6 +848,22 @@ async function loadToolDefinitionsWrapper({ req, res, agent, streamId = null, to
},
);
for (const serverName of getMCPServerNamesFromTools(filteredTools)) {
if (pendingOAuthServers.has(serverName)) {
continue;
}
const pendingOAuthStart = await getReplayablePendingMCPOAuthStart({
flowManager,
userId: req.user.id,
serverName,
});
if (pendingOAuthStart) {
pendingOAuthServers.add(serverName);
pendingOAuthStarts.set(serverName, pendingOAuthStart);
}
}
if (pendingOAuthServers.size > 0 && (res || streamId)) {
const serverNames = Array.from(pendingOAuthServers);
logger.info(
@ -821,6 +872,12 @@ async function loadToolDefinitionsWrapper({ req, res, agent, streamId = null, to
const oauthWaitPromises = serverNames.map(async (serverName, index) => {
try {
const pendingOAuthStart = await getPendingOAuthStartForEmit(serverName);
const oauthStart = createOAuthEmitter(serverName, index);
if (pendingOAuthStart) {
await oauthStart(pendingOAuthStart.authURL, pendingOAuthStart.options);
}
const result = await reinitMCPServer({
user: req.user,
serverName,
@ -828,7 +885,7 @@ async function loadToolDefinitionsWrapper({ req, res, agent, streamId = null, to
userMCPAuthMap,
flowManager,
returnOnOAuth: false,
oauthStart: createOAuthEmitter(serverName, index),
oauthStart,
oauthEnd: createOAuthEndEmitter(serverName),
connectionTimeout: Time.TWO_MINUTES,
});

View file

@ -20,7 +20,7 @@ const { getLogStores } = require('~/cache');
* @param {boolean} [params.forceNew]
* @param {number} [params.connectionTimeout]
* @param {FlowStateManager<any>} [params.flowManager]
* @param {(authURL: string) => Promise<void>} [params.oauthStart]
* @param {(authURL: string, options?: { expiresAt?: number }) => Promise<void>} [params.oauthStart]
* @param {() => Promise<void>} [params.oauthEnd]
* @param {Record<string, Record<string, string>>} [params.userMCPAuthMap]
*/

View file

@ -42,6 +42,7 @@ const mockLegacyDomainEncode = jest.fn();
const mockDecryptMetadata = jest.fn();
const mockCreateActionTool = jest.fn();
const mockGetServerConfig = jest.fn();
const mockFlowManager = { getFlowState: jest.fn() };
const mockResolveConfigServers = jest.fn();
const mockUserCanUseMCPServers = jest.fn().mockResolvedValue(true);
jest.mock('~/server/services/Tools/credentials', () => ({
@ -77,7 +78,7 @@ jest.mock('~/models', () => ({
findPluginAuthsByKeys: jest.fn(),
}));
jest.mock('~/config', () => ({
getFlowStateManager: jest.fn(() => ({})),
getFlowStateManager: jest.fn(() => mockFlowManager),
getMCPServersRegistry: jest.fn(() => ({
getServerConfig: (...args) => mockGetServerConfig(...args),
})),
@ -100,6 +101,7 @@ const {
resolveAgentCapabilities,
} = require('../ToolService');
const { reinitMCPServer } = require('~/server/services/Tools/mcp');
const { PENDING_STALE_MS } = require('@librechat/api');
function createMockReq(capabilities) {
return {
@ -134,6 +136,7 @@ describe('ToolService - Action Capability Gating', () => {
mockGetCachedTools.mockResolvedValue(null);
mockGetUserMCPAuthMap.mockResolvedValue({});
mockGetServerConfig.mockResolvedValue(undefined);
mockFlowManager.getFlowState.mockResolvedValue(undefined);
mockResolveConfigServers.mockResolvedValue({});
});
@ -421,6 +424,349 @@ describe('ToolService - Action Capability Gating', () => {
expect(mockGetMCPServerTools).not.toHaveBeenCalled();
});
it('should re-emit pending MCP OAuth prompts when cached tool definitions exist', async () => {
const serverName = 'Google-Workspace';
const authorizationUrl = 'https://auth.example.com/Google-Workspace';
const mcpTool = `search${Constants.mcp_delimiter}${serverName}`;
const capabilities = [AgentCapabilities.tools];
const req = createMockReq(capabilities);
const res = { writableEnded: false };
mockGetEndpointsConfig.mockResolvedValue(createEndpointsConfig(capabilities));
mockGetServerConfig.mockResolvedValue({
type: 'streamable-http',
url: 'https://demo.librechat.ai/mcp',
requiresOAuth: true,
});
mockGetMCPServerTools.mockResolvedValue({
[mcpTool]: {
function: {
name: mcpTool,
description: 'Cached search',
parameters: {},
},
},
});
mockFlowManager.getFlowState.mockResolvedValue({
status: 'PENDING',
createdAt: Date.now(),
metadata: { authorizationUrl },
});
mockLoadToolDefinitions.mockImplementation(async (params, deps) => {
const serverTools = await deps.getOrFetchMCPServerTools(params.userId, serverName);
return {
toolDefinitions: serverTools ? Object.keys(serverTools) : [],
toolRegistry: new Map(),
hasDeferredTools: false,
};
});
reinitMCPServer.mockImplementation(async ({ oauthStart }) => {
await oauthStart(authorizationUrl);
return { availableTools: { [mcpTool]: {} } };
});
const result = await loadAgentTools({
req,
res,
agent: { id: 'agent_123', tools: [mcpTool] },
definitionsOnly: true,
});
expect(result.toolDefinitions).toEqual([mcpTool]);
expect(mockGetMCPServerTools).toHaveBeenCalledWith(req.user.id, serverName);
expect(reinitMCPServer).toHaveBeenCalledWith(
expect.objectContaining({
serverName,
returnOnOAuth: false,
oauthStart: expect.any(Function),
}),
);
expect(mockSendEvent).toHaveBeenCalledWith(
res,
expect.objectContaining({
event: 'on_run_step',
data: expect.objectContaining({
id: `step_oauth_login_${serverName}`,
}),
}),
);
expect(mockSendEvent).toHaveBeenCalledWith(
res,
expect.objectContaining({
event: 'on_run_step_delta',
data: expect.objectContaining({
id: `step_oauth_login_${serverName}`,
delta: expect.objectContaining({
auth: authorizationUrl,
}),
}),
}),
);
});
it('should not join in-flight MCP initialization before replaying pending OAuth prompts', async () => {
const serverName = 'Google-Workspace';
const authorizationUrl = 'https://auth.example.com/Google-Workspace';
const mcpTool = `${Constants.mcp_all}${Constants.mcp_delimiter}${serverName}`;
const capabilities = [AgentCapabilities.tools];
const req = createMockReq(capabilities);
const res = { writableEnded: false };
mockGetEndpointsConfig.mockResolvedValue(createEndpointsConfig(capabilities));
mockGetServerConfig.mockResolvedValue({
type: 'streamable-http',
url: 'https://demo.librechat.ai/mcp',
requiresOAuth: true,
});
mockGetMCPServerTools.mockResolvedValue(null);
mockFlowManager.getFlowState.mockResolvedValue({
status: 'PENDING',
createdAt: Date.now(),
metadata: { authorizationUrl },
});
mockLoadToolDefinitions.mockImplementation(async (params, deps) => {
await deps.getOrFetchMCPServerTools(params.userId, serverName);
return {
toolDefinitions: [],
toolRegistry: new Map(),
hasDeferredTools: false,
};
});
reinitMCPServer.mockImplementation(async ({ oauthStart }) => {
await oauthStart(authorizationUrl);
return { availableTools: null };
});
await loadAgentTools({
req,
res,
agent: { id: 'agent_123', tools: [mcpTool] },
definitionsOnly: true,
});
expect(mockGetMCPServerTools).toHaveBeenCalledWith(req.user.id, serverName);
expect(reinitMCPServer).toHaveBeenCalledTimes(1);
expect(reinitMCPServer).toHaveBeenCalledWith(
expect.objectContaining({
serverName,
returnOnOAuth: false,
oauthStart: expect.any(Function),
}),
);
expect(mockSendEvent).toHaveBeenCalledWith(
res,
expect.objectContaining({
event: 'on_run_step_delta',
data: expect.objectContaining({
id: `step_oauth_login_${serverName}`,
delta: expect.objectContaining({
auth: authorizationUrl,
}),
}),
}),
);
});
it('should re-emit pending MCP OAuth prompts when selected MCP tools are already concrete', async () => {
const serverName = `Google${Constants.mcp_delimiter}Workspace`;
const authorizationUrl = 'https://auth.example.com/Google-Workspace';
const mcpTool = `search${Constants.mcp_delimiter}${serverName}`;
const capabilities = [AgentCapabilities.tools];
const req = createMockReq(capabilities);
const res = { writableEnded: false };
mockGetEndpointsConfig.mockResolvedValue(createEndpointsConfig(capabilities));
mockFlowManager.getFlowState.mockResolvedValue({
status: 'PENDING',
createdAt: Date.now(),
metadata: { authorizationUrl },
});
mockLoadToolDefinitions.mockResolvedValue({
toolDefinitions: [mcpTool],
toolRegistry: new Map(),
hasDeferredTools: false,
});
reinitMCPServer.mockImplementation(async ({ oauthStart }) => {
await oauthStart(authorizationUrl);
return { availableTools: { [mcpTool]: {} } };
});
const result = await loadAgentTools({
req,
res,
agent: { id: 'agent_123', tools: [mcpTool] },
definitionsOnly: true,
});
expect(result.toolDefinitions).toEqual([mcpTool]);
expect(mockGetMCPServerTools).not.toHaveBeenCalled();
expect(reinitMCPServer).toHaveBeenCalledWith(
expect.objectContaining({
serverName,
returnOnOAuth: false,
oauthStart: expect.any(Function),
}),
);
expect(mockSendEvent).toHaveBeenCalledWith(
res,
expect.objectContaining({
event: 'on_run_step_delta',
data: expect.objectContaining({
id: `step_oauth_login_${serverName}`,
delta: expect.objectContaining({
auth: authorizationUrl,
}),
}),
}),
);
});
it('should emit stored pending MCP OAuth prompts before waiting on a silent in-flight join', async () => {
const serverName = 'Google-Workspace';
const authorizationUrl = 'https://auth.example.com/Google-Workspace';
const mcpTool = `search${Constants.mcp_delimiter}${serverName}`;
const capabilities = [AgentCapabilities.tools];
const req = createMockReq(capabilities);
const res = { writableEnded: false };
mockGetEndpointsConfig.mockResolvedValue(createEndpointsConfig(capabilities));
mockFlowManager.getFlowState.mockResolvedValue({
status: 'PENDING',
createdAt: Date.now(),
metadata: { authorizationUrl },
});
mockLoadToolDefinitions.mockResolvedValue({
toolDefinitions: [mcpTool],
toolRegistry: new Map(),
hasDeferredTools: false,
});
reinitMCPServer.mockResolvedValue({ availableTools: null });
const result = await loadAgentTools({
req,
res,
agent: { id: 'agent_123', tools: [mcpTool] },
definitionsOnly: true,
});
expect(result.toolDefinitions).toEqual([mcpTool]);
expect(reinitMCPServer).toHaveBeenCalledWith(
expect.objectContaining({
serverName,
returnOnOAuth: false,
oauthStart: expect.any(Function),
}),
);
expect(mockSendEvent).toHaveBeenCalledWith(
res,
expect.objectContaining({
event: 'on_run_step_delta',
data: expect.objectContaining({
id: `step_oauth_login_${serverName}`,
delta: expect.objectContaining({
auth: authorizationUrl,
}),
}),
}),
);
});
it('should preserve OAuth URLs emitted while discovering MCP tools before a silent wait join', async () => {
const serverName = 'Google-Workspace';
const authorizationUrl = 'https://auth.example.com/Google-Workspace';
const mcpTool = `search${Constants.mcp_delimiter}${serverName}`;
const capabilities = [AgentCapabilities.tools];
const req = createMockReq(capabilities);
const res = { writableEnded: false };
mockGetEndpointsConfig.mockResolvedValue(createEndpointsConfig(capabilities));
mockGetServerConfig.mockResolvedValue({
type: 'streamable-http',
url: 'https://demo.librechat.ai/mcp',
requiresOAuth: true,
});
mockGetMCPServerTools.mockResolvedValue(null);
mockFlowManager.getFlowState.mockResolvedValue(null);
mockLoadToolDefinitions.mockImplementation(async (params, deps) => {
await deps.getOrFetchMCPServerTools(params.userId, serverName);
return {
toolDefinitions: [],
toolRegistry: new Map(),
hasDeferredTools: false,
};
});
reinitMCPServer
.mockImplementationOnce(async ({ oauthStart }) => {
await oauthStart(authorizationUrl, { expiresAt: Date.now() + 60_000 });
return { availableTools: null };
})
.mockResolvedValue({ availableTools: null });
await loadAgentTools({
req,
res,
agent: { id: 'agent_123', tools: [mcpTool] },
definitionsOnly: true,
});
expect(reinitMCPServer).toHaveBeenCalledTimes(2);
expect(mockSendEvent).toHaveBeenCalledWith(
res,
expect.objectContaining({
event: 'on_run_step_delta',
data: expect.objectContaining({
id: `step_oauth_login_${serverName}`,
delta: expect.objectContaining({
auth: authorizationUrl,
}),
}),
}),
);
});
it('should preserve pending-flow expiry for OAuth URLs captured during discovery', async () => {
const serverName = 'Google-Workspace';
const authorizationUrl = 'https://auth.example.com/Google-Workspace';
const mcpTool = `search${Constants.mcp_delimiter}${serverName}`;
const capabilities = [AgentCapabilities.tools];
const req = createMockReq(capabilities);
const res = { writableEnded: false };
const createdAt = Date.now() - 45_000;
mockGetEndpointsConfig.mockResolvedValue(createEndpointsConfig(capabilities));
mockGetServerConfig.mockResolvedValue({
type: 'streamable-http',
url: 'https://demo.librechat.ai/mcp',
requiresOAuth: true,
});
mockGetMCPServerTools.mockResolvedValue(null);
mockFlowManager.getFlowState.mockResolvedValueOnce(null).mockResolvedValueOnce({
status: 'PENDING',
createdAt,
metadata: { authorizationUrl },
});
mockLoadToolDefinitions.mockImplementation(async (params, deps) => {
await deps.getOrFetchMCPServerTools(params.userId, serverName);
return {
toolDefinitions: [],
toolRegistry: new Map(),
hasDeferredTools: false,
};
});
reinitMCPServer
.mockImplementationOnce(async ({ oauthStart }) => {
await oauthStart(authorizationUrl);
return { availableTools: null };
})
.mockResolvedValue({ availableTools: null });
await loadAgentTools({
req,
res,
agent: { id: 'agent_123', tools: [mcpTool] },
definitionsOnly: true,
});
const authDeltaEvent = mockSendEvent.mock.calls
.map(([, event]) => event)
.find((event) => event.data?.delta?.auth === authorizationUrl);
expect(authDeltaEvent?.data.delta.expires_at).toBe(createdAt + PENDING_STALE_MS);
});
it('should use request-scoped MCP config before falling back to the registry', async () => {
const serverName = 'config-server';
const mcpTool = `search${Constants.mcp_delimiter}${serverName}`;