refactor(callbackServer): make cancellable

This commit is contained in:
Elliott de Launay 2026-04-24 09:51:19 -04:00
parent b9d2852cc3
commit cde464eb48
No known key found for this signature in database
GPG key ID: BB899BED766D1806
2 changed files with 116 additions and 73 deletions

View file

@ -9,6 +9,7 @@ vi.mock("http", () => ({
describe("startCallbackServer", () => { describe("startCallbackServer", () => {
beforeEach(() => { beforeEach(() => {
vi.restoreAllMocks() vi.restoreAllMocks()
delete process.env.MCP_OAUTH_TEST_MODE
}) })
it("should start server and resolve with callback result", async () => { it("should start server and resolve with callback result", async () => {
@ -95,7 +96,32 @@ describe("stopCallbackServer", () => {
close: vi.fn((callback) => callback()), 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() 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
}
})
})

View file

@ -22,6 +22,7 @@ export function startCallbackServer(
server: http.Server server: http.Server
port: number port: number
result: Promise<CallbackResult> result: Promise<CallbackResult>
cancel: () => void
}> { }> {
// In test mode, immediately resolve with mock data // In test mode, immediately resolve with mock data
if (process.env.MCP_OAUTH_TEST_MODE === "true") { if (process.env.MCP_OAUTH_TEST_MODE === "true") {
@ -31,6 +32,7 @@ export function startCallbackServer(
server: mockServer, server: mockServer,
port: 3000, port: 3000,
result: Promise.resolve({ code: "test-auth-code", state: expectedState }), result: Promise.resolve({ code: "test-auth-code", state: expectedState }),
cancel: () => {},
}) })
}) })
} }
@ -47,63 +49,77 @@ export function startCallbackServer(
const actualPort = address.port const actualPort = address.port
const resultPromise = new Promise<CallbackResult>((resolveResult, rejectResult) => { let resolveResult!: (value: CallbackResult) => void
let resolved = false let rejectResult!: (reason: unknown) => void
const resultPromise = new Promise<CallbackResult>((res, rej) => {
resolveResult = res
rejectResult = rej
})
const timeout = setTimeout(() => { let resolved = false
if (!resolved) {
resolved = true
rejectResult(new Error("Callback timeout"))
server.close()
}
}, OAUTH_FLOW_TIMEOUT_MS)
server.on("request", (req: any, res: any) => { const timeout = setTimeout(() => {
if (resolved) return 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 cancel = () => {
const pathname = url.pathname if (!resolved) {
resolved = true
clearTimeout(timeout)
rejectResult(new Error("Callback cancelled"))
}
}
if (pathname === "/callback") { server.on("request", (req: any, res: any) => {
resolved = true if (resolved) return
clearTimeout(timeout)
const code = url.searchParams.get("code") const url = new URL(req.url || "", `http://localhost:${actualPort}`)
const error = url.searchParams.get("error") const pathname = url.pathname
const errorDescription = url.searchParams.get("error_description")
const state = url.searchParams.get("state")
const hasError = !!error
// Verify state for CSRF protection if (pathname === "/callback") {
if (expectedState && state !== expectedState) { resolved = true
res.writeHead(400, { clearTimeout(timeout)
"Content-Type": "text/html",
"Content-Security-Policy": "default-src 'none'; style-src 'unsafe-inline'",
})
res.end(`
<!DOCTYPE html>
<html>
<head>
<title>${t("mcp:oauth.callback.title")}</title>
</head>
<body>
<h1>${t("mcp:oauth.callback.failed")}</h1>
<p>${t("mcp:oauth.callback.invalid_state")}</p>
</body>
</html>
`)
rejectResult(new Error("Invalid state parameter"))
server.close()
return
}
// Send HTML response const code = url.searchParams.get("code")
res.writeHead(200, { 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-Type": "text/html",
"Content-Security-Policy": "Content-Security-Policy": "default-src 'none'; style-src 'unsafe-inline'",
"default-src 'none'; style-src 'unsafe-inline'; script-src 'unsafe-inline'",
}) })
res.end(` res.end(`
<!DOCTYPE html>
<html>
<head>
<title>${t("mcp:oauth.callback.title")}</title>
</head>
<body>
<h1>${t("mcp:oauth.callback.failed")}</h1>
<p>${t("mcp:oauth.callback.invalid_state")}</p>
</body>
</html>
`)
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(`
<!DOCTYPE html> <!DOCTYPE html>
<html> <html>
<head> <head>
@ -152,36 +168,36 @@ export function startCallbackServer(
</html> </html>
`) `)
resolveResult({ resolveResult({
code: code || undefined, code: code || undefined,
error: error || undefined, error: error || undefined,
error_description: errorDescription || undefined, error_description: errorDescription || undefined,
state: state || undefined, state: state || undefined,
}) })
// Close server immediately after response drains // Close server immediately after response drains
res.on("finish", () => { res.on("finish", () => {
server.close() server.close()
}) })
} else { } else {
res.writeHead(404) res.writeHead(404)
res.end("Not found") res.end("Not found")
} }
}) })
server.on("error", (error: any) => { server.on("error", (error: any) => {
if (!resolved) { if (!resolved) {
resolved = true resolved = true
clearTimeout(timeout) clearTimeout(timeout)
rejectResult(error) rejectResult(error)
} }
})
}) })
resolve({ resolve({
server, server,
port: actualPort, port: actualPort,
result: resultPromise, result: resultPromise,
cancel,
}) })
}) })
@ -190,10 +206,11 @@ export function startCallbackServer(
} }
/** /**
* Stops the callback server. * Stops the callback server and cancels any pending result promise so its
* @param server The HTTP server to stop * timeout doesn't fire after the provider is already closed.
*/ */
export function stopCallbackServer(server: http.Server): Promise<void> { export function stopCallbackServer(server: http.Server, cancel: () => void): Promise<void> {
cancel()
return new Promise((resolve) => { return new Promise((resolve) => {
server.close(() => resolve()) server.close(() => resolve())
}) })