diff --git a/src/services/mcp/utils/__tests__/callbackServer.spec.ts b/src/services/mcp/utils/__tests__/callbackServer.spec.ts index 3159ffbae2..b5af501a3b 100644 --- a/src/services/mcp/utils/__tests__/callbackServer.spec.ts +++ b/src/services/mcp/utils/__tests__/callbackServer.spec.ts @@ -9,6 +9,7 @@ vi.mock("http", () => ({ describe("startCallbackServer", () => { beforeEach(() => { vi.restoreAllMocks() + delete process.env.MCP_OAUTH_TEST_MODE }) it("should start server and resolve with callback result", async () => { @@ -95,7 +96,32 @@ describe("stopCallbackServer", () => { close: vi.fn((callback) => callback()), } - await stopCallbackServer(mockServer as any) + await stopCallbackServer(mockServer as any, () => {}) + expect(mockServer.close).toHaveBeenCalled() + }) + + it("should call the cancel function before closing", async () => { + const mockServer = { close: vi.fn((callback) => callback()) } + const cancel = vi.fn() + + await stopCallbackServer(mockServer as any, cancel) + expect(cancel).toHaveBeenCalledTimes(1) expect(mockServer.close).toHaveBeenCalled() }) }) + +describe("startCallbackServer in test mode", () => { + it("should resolve immediately with mock auth code when MCP_OAUTH_TEST_MODE is set", async () => { + process.env.MCP_OAUTH_TEST_MODE = "true" + try { + const { port, result, cancel } = await startCallbackServer(undefined, "test-state") + expect(port).toBe(3000) + expect(typeof cancel).toBe("function") + const callbackResult = await result + expect(callbackResult.code).toBe("test-auth-code") + expect(callbackResult.state).toBe("test-state") + } finally { + delete process.env.MCP_OAUTH_TEST_MODE + } + }) +}) diff --git a/src/services/mcp/utils/callbackServer.ts b/src/services/mcp/utils/callbackServer.ts index 2f12b79b96..7610e09541 100644 --- a/src/services/mcp/utils/callbackServer.ts +++ b/src/services/mcp/utils/callbackServer.ts @@ -22,6 +22,7 @@ export function startCallbackServer( server: http.Server port: number result: Promise + cancel: () => void }> { // In test mode, immediately resolve with mock data if (process.env.MCP_OAUTH_TEST_MODE === "true") { @@ -31,6 +32,7 @@ export function startCallbackServer( server: mockServer, port: 3000, result: Promise.resolve({ code: "test-auth-code", state: expectedState }), + cancel: () => {}, }) }) } @@ -47,63 +49,77 @@ export function startCallbackServer( const actualPort = address.port - const resultPromise = new Promise((resolveResult, rejectResult) => { - let resolved = false + let resolveResult!: (value: CallbackResult) => void + let rejectResult!: (reason: unknown) => void + const resultPromise = new Promise((res, rej) => { + resolveResult = res + rejectResult = rej + }) - const timeout = setTimeout(() => { - if (!resolved) { - resolved = true - rejectResult(new Error("Callback timeout")) - server.close() - } - }, OAUTH_FLOW_TIMEOUT_MS) + let resolved = false - server.on("request", (req: any, res: any) => { - if (resolved) return + const timeout = setTimeout(() => { + if (!resolved) { + resolved = true + rejectResult(new Error("Callback timeout")) + server.close() + } + }, OAUTH_FLOW_TIMEOUT_MS) - const url = new URL(req.url || "", `http://localhost:${actualPort}`) - const pathname = url.pathname + const cancel = () => { + if (!resolved) { + resolved = true + clearTimeout(timeout) + rejectResult(new Error("Callback cancelled")) + } + } - if (pathname === "/callback") { - resolved = true - clearTimeout(timeout) + server.on("request", (req: any, res: any) => { + if (resolved) return - const code = url.searchParams.get("code") - const error = url.searchParams.get("error") - const errorDescription = url.searchParams.get("error_description") - const state = url.searchParams.get("state") - const hasError = !!error + const url = new URL(req.url || "", `http://localhost:${actualPort}`) + const pathname = url.pathname - // Verify state for CSRF protection - if (expectedState && state !== expectedState) { - res.writeHead(400, { - "Content-Type": "text/html", - "Content-Security-Policy": "default-src 'none'; style-src 'unsafe-inline'", - }) - res.end(` - - - - ${t("mcp:oauth.callback.title")} - - -

${t("mcp:oauth.callback.failed")}

-

${t("mcp:oauth.callback.invalid_state")}

- - - `) - rejectResult(new Error("Invalid state parameter")) - server.close() - return - } + if (pathname === "/callback") { + resolved = true + clearTimeout(timeout) - // Send HTML response - res.writeHead(200, { + const code = url.searchParams.get("code") + const error = url.searchParams.get("error") + const errorDescription = url.searchParams.get("error_description") + const state = url.searchParams.get("state") + const hasError = !!error + + // Verify state for CSRF protection + if (expectedState && state !== expectedState) { + res.writeHead(400, { "Content-Type": "text/html", - "Content-Security-Policy": - "default-src 'none'; style-src 'unsafe-inline'; script-src 'unsafe-inline'", + "Content-Security-Policy": "default-src 'none'; style-src 'unsafe-inline'", }) res.end(` + + + + ${t("mcp:oauth.callback.title")} + + +

${t("mcp:oauth.callback.failed")}

+

${t("mcp:oauth.callback.invalid_state")}

+ + + `) + rejectResult(new Error("Invalid state parameter")) + server.close() + return + } + + // Send HTML response + res.writeHead(200, { + "Content-Type": "text/html", + "Content-Security-Policy": + "default-src 'none'; style-src 'unsafe-inline'; script-src 'unsafe-inline'", + }) + res.end(` @@ -152,36 +168,36 @@ export function startCallbackServer( `) - resolveResult({ - code: code || undefined, - error: error || undefined, - error_description: errorDescription || undefined, - state: state || undefined, - }) + resolveResult({ + code: code || undefined, + error: error || undefined, + error_description: errorDescription || undefined, + state: state || undefined, + }) - // Close server immediately after response drains - res.on("finish", () => { - server.close() - }) - } else { - res.writeHead(404) - res.end("Not found") - } - }) + // Close server immediately after response drains + res.on("finish", () => { + server.close() + }) + } else { + res.writeHead(404) + res.end("Not found") + } + }) - server.on("error", (error: any) => { - if (!resolved) { - resolved = true - clearTimeout(timeout) - rejectResult(error) - } - }) + server.on("error", (error: any) => { + if (!resolved) { + resolved = true + clearTimeout(timeout) + rejectResult(error) + } }) resolve({ server, port: actualPort, result: resultPromise, + cancel, }) }) @@ -190,10 +206,11 @@ export function startCallbackServer( } /** - * Stops the callback server. - * @param server The HTTP server to stop + * Stops the callback server and cancels any pending result promise so its + * timeout doesn't fire after the provider is already closed. */ -export function stopCallbackServer(server: http.Server): Promise { +export function stopCallbackServer(server: http.Server, cancel: () => void): Promise { + cancel() return new Promise((resolve) => { server.close(() => resolve()) })