feat(callbackServer): used for OAuth 2.1 callbacks

This commit is contained in:
Elliott de Launay 2026-03-06 14:18:05 +00:00 committed by Elliott de Launay
parent fb8f827127
commit ae2aab2aa4
No known key found for this signature in database
GPG key ID: BB899BED766D1806
2 changed files with 251 additions and 0 deletions

View file

@ -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()
})
})

View file

@ -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<CallbackResult>}>
*/
export function startCallbackServer(
port?: number,
expectedState?: string,
): Promise<{
server: http.Server
port: number
result: Promise<CallbackResult>
}> {
// 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<CallbackResult>((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(`
<!DOCTYPE html>
<html>
<head>
<title>OAuth Callback</title>
</head>
<body>
<h1>OAuth Authentication Failed</h1>
<p>Error: Invalid state parameter</p>
</body>
</html>
`)
rejectResult(new Error("Invalid state parameter"))
return
}
// Send HTML response
res.writeHead(200, { "Content-Type": "text/html" })
res.end(`
<!DOCTYPE html>
<html>
<head>
<title>OAuth Callback</title>
</head>
<body>
<h1>OAuth Authentication ${error ? "Failed" : "Successful"}</h1>
<p>
${error ? `Error: ${error}${errorDescription ? ` - ${errorDescription}` : ""}` : "You can close this window."}
</p>
</body>
</html>
`)
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<void> {
return new Promise((resolve) => {
server.close(() => resolve())
})
}