mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-06 02:47:56 +00:00
refactor(callbackServer): make cancellable
This commit is contained in:
parent
b9d2852cc3
commit
cde464eb48
2 changed files with 116 additions and 73 deletions
|
|
@ -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
|
||||||
|
}
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
|
||||||
|
|
@ -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())
|
||||||
})
|
})
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue