mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-09-06 08:18:39 +00:00
feat(callbackServer): used for OAuth 2.1 callbacks
This commit is contained in:
parent
fb8f827127
commit
ae2aab2aa4
2 changed files with 251 additions and 0 deletions
95
src/services/mcp/utils/__tests__/callbackServer.spec.ts
Normal file
95
src/services/mcp/utils/__tests__/callbackServer.spec.ts
Normal 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()
|
||||
})
|
||||
})
|
||||
156
src/services/mcp/utils/callbackServer.ts
Normal file
156
src/services/mcp/utils/callbackServer.ts
Normal 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())
|
||||
})
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue