diff --git a/src/services/mcp/utils/__tests__/callbackServer.spec.ts b/src/services/mcp/utils/__tests__/callbackServer.spec.ts new file mode 100644 index 0000000000..f7e26cdd27 --- /dev/null +++ b/src/services/mcp/utils/__tests__/callbackServer.spec.ts @@ -0,0 +1,95 @@ +import { describe, it, expect, vi, beforeEach } from "vitest" +import { startCallbackServer, stopCallbackServer } from "../callbackServer" +import * as http from "http" + +vi.mock("http", () => ({ + createServer: vi.fn(), +})) + +describe("startCallbackServer", () => { + beforeEach(() => { + vi.restoreAllMocks() + }) + + it("should start server and resolve with callback result", async () => { + const mockServer = { + listen: vi.fn((port, host, callback) => { + callback() + return mockServer + }), + address: vi.fn(() => ({ port: 3000 })), + on: vi.fn(), + close: vi.fn(), + } + + ;(http.createServer as any).mockReturnValue(mockServer) + + const promise = startCallbackServer() + const { server, port, result } = await promise + + expect(port).toBe(3000) + expect(server).toBe(mockServer) + + // Simulate callback request + const requestCall = mockServer.on.mock.calls.find((call) => call[0] === "request") + const requestHandler = requestCall ? requestCall[1] : vi.fn() + const mockReq = { + url: "/callback?code=test-code&state=test-state", + method: "GET", + } + const mockRes = { + writeHead: vi.fn(), + end: vi.fn(), + } + + requestHandler(mockReq, mockRes) + + const callbackResult = await result + expect(callbackResult.code).toBe("test-code") + expect(callbackResult.state).toBe("test-state") + }) + + it("should reject invalid state", async () => { + const mockServer = { + listen: vi.fn((port, host, callback) => { + callback() + return mockServer + }), + address: vi.fn(() => ({ port: 3000 })), + on: vi.fn(), + close: vi.fn(), + } + + ;(http.createServer as any).mockReturnValue(mockServer) + + const promise = startCallbackServer(undefined, "expected-state") + const { result } = await promise + + // Simulate callback request with wrong state + const requestCall = mockServer.on.mock.calls.find((call) => call[0] === "request") + const requestHandler = requestCall ? requestCall[1] : vi.fn() + const mockReq = { + url: "/callback?code=test-code&state=wrong-state", + method: "GET", + } + const mockRes = { + writeHead: vi.fn(), + end: vi.fn(), + } + + requestHandler(mockReq, mockRes) + + await expect(result).rejects.toThrow("Invalid state parameter") + }) +}) + +describe("stopCallbackServer", () => { + it("should close the server", async () => { + const mockServer = { + close: vi.fn((callback) => callback()), + } + + await stopCallbackServer(mockServer as any) + expect(mockServer.close).toHaveBeenCalled() + }) +}) diff --git a/src/services/mcp/utils/callbackServer.ts b/src/services/mcp/utils/callbackServer.ts new file mode 100644 index 0000000000..4070f6f4f3 --- /dev/null +++ b/src/services/mcp/utils/callbackServer.ts @@ -0,0 +1,156 @@ +import * as http from "http" + +export interface CallbackResult { + code?: string + error?: string + error_description?: string + state?: string +} + +/** + * Starts a local HTTP server to handle OAuth callback. + * @param port Optional port to use (defaults to random available port) + * @param expectedState Optional expected state for CSRF protection + * @returns Promise<{server: http.Server, port: number, result: Promise}> + */ +export function startCallbackServer( + port?: number, + expectedState?: string, +): Promise<{ + server: http.Server + port: number + result: Promise +}> { + // In test mode, immediately resolve with mock data + if (process.env.MCP_OAUTH_TEST_MODE === "true") { + return new Promise((resolve) => { + const mockServer = http.createServer() + resolve({ + server: mockServer, + port: 3000, + result: Promise.resolve({ code: "test-auth-code", state: expectedState }), + }) + }) + } + + return new Promise((resolve, reject) => { + const server = http.createServer() + + server.listen(port || 0, "127.0.0.1", () => { + const address = server.address() + if (!address || typeof address === "string") { + reject(new Error("Failed to get server address")) + return + } + + const actualPort = address.port + + const resultPromise = new Promise((resolveResult, rejectResult) => { + let resolved = false + + const timeout = setTimeout( + () => { + if (!resolved) { + resolved = true + rejectResult(new Error("Callback timeout")) + server.close() + } + }, + 5 * 60 * 1000, + ) // 5 minutes + + server.on("request", (req: any, res: any) => { + if (resolved) return + + const url = new URL(req.url || "", `http://localhost:${actualPort}`) + const pathname = url.pathname + + if (pathname === "/callback") { + resolved = true + clearTimeout(timeout) + + 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") + + // Verify state for CSRF protection + if (expectedState && state !== expectedState) { + res.writeHead(400, { "Content-Type": "text/html" }) + res.end(` + + + + OAuth Callback + + +

OAuth Authentication Failed

+

Error: Invalid state parameter

+ + + `) + rejectResult(new Error("Invalid state parameter")) + return + } + + // Send HTML response + res.writeHead(200, { "Content-Type": "text/html" }) + res.end(` + + + + OAuth Callback + + +

OAuth Authentication ${error ? "Failed" : "Successful"}

+

+ ${error ? `Error: ${error}${errorDescription ? ` - ${errorDescription}` : ""}` : "You can close this window."} +

+ + + `) + + resolveResult({ + code: code || undefined, + error: error || undefined, + error_description: errorDescription || undefined, + state: state || undefined, + }) + + // Close server after a short delay to allow response to be sent + setTimeout(() => server.close(), 1000) + } else { + res.writeHead(404) + res.end("Not found") + } + }) + + server.on("error", (error: any) => { + if (!resolved) { + resolved = true + clearTimeout(timeout) + rejectResult(error) + } + }) + }) + + resolve({ + server, + port: actualPort, + result: resultPromise, + }) + }) + + server.on("error", reject) + }) +} + +/** + * Stops the callback server. + * @param server The HTTP server to stop + */ +export function stopCallbackServer(server: http.Server): Promise { + return new Promise((resolve) => { + server.close(() => resolve()) + }) +}