mirror of
https://github.com/danny-avila/LibreChat.git
synced 2026-10-09 06:56:30 +00:00
📻 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:
parent
6950448d03
commit
cb1d536874
48 changed files with 4837 additions and 299 deletions
|
|
@ -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),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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'));
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
*/
|
||||
|
|
|
|||
|
|
@ -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}`;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue