From 571510bdf32ded4788c124efd6c2fbfcdce30874 Mon Sep 17 00:00:00 2001 From: maximhar Date: Tue, 17 Mar 2026 15:49:23 +0200 Subject: [PATCH] fix: MCP OAuth exact redirect URIs for Slack-style providers (#454) * Add exact MCP OAuth redirect URI support * Allow proxied HTTPS loopback MCP redirects * Fix MCP OAuth busy-port test determinism * Make MCP OAuth redirect test deterministic * fix(coding-agent): honor exact loopback redirect ports * fix(coding-agent): expand env vars in standalone MCP oauth config --------- Co-authored-by: can1357 --- .../ai/src/utils/oauth/callback-server.ts | 48 +++- packages/coding-agent/CHANGELOG.md | 4 + packages/coding-agent/src/capability/mcp.ts | 8 +- .../coding-agent/src/discovery/builtin.ts | 20 +- .../coding-agent/src/discovery/mcp-json.ts | 14 +- packages/coding-agent/src/mcp/oauth-flow.ts | 88 +++++- packages/coding-agent/src/mcp/types.ts | 3 + .../controllers/mcp-command-controller.ts | 35 +-- .../test/discovery/mcp-json.test.ts | 123 +++++++++ packages/coding-agent/test/oauth-flow.test.ts | 251 +++++++++++++++++- 10 files changed, 560 insertions(+), 34 deletions(-) create mode 100644 packages/coding-agent/test/discovery/mcp-json.test.ts diff --git a/packages/ai/src/utils/oauth/callback-server.ts b/packages/ai/src/utils/oauth/callback-server.ts index 683d08ea6..acd9f888b 100644 --- a/packages/ai/src/utils/oauth/callback-server.ts +++ b/packages/ai/src/utils/oauth/callback-server.ts @@ -19,6 +19,14 @@ const CALLBACK_PATH = "/callback"; export type CallbackResult = { code: string; state: string }; +export interface OAuthCallbackFlowOptions { + preferredPort: number; + callbackPath?: string; + callbackHostname?: string; + /** Exact redirect URI advertised to the provider; disables port fallback. */ + redirectUri?: string; +} + /** * Abstract base class for OAuth flows with local callback servers. */ @@ -26,13 +34,28 @@ export abstract class OAuthCallbackFlow { ctrl: OAuthController; preferredPort: number; callbackPath: string; + callbackHostname: string; + redirectUri?: string; #callbackResolve?: (result: CallbackResult) => void; #callbackReject?: (error: string) => void; - constructor(ctrl: OAuthController, preferredPort: number, callbackPath: string = CALLBACK_PATH) { + constructor( + ctrl: OAuthController, + preferredPortOrOptions: number | OAuthCallbackFlowOptions, + callbackPath: string = CALLBACK_PATH, + ) { this.ctrl = ctrl; - this.preferredPort = preferredPort; - this.callbackPath = callbackPath; + if (typeof preferredPortOrOptions === "number") { + this.preferredPort = preferredPortOrOptions; + this.callbackPath = callbackPath; + this.callbackHostname = DEFAULT_HOSTNAME; + return; + } + + this.preferredPort = preferredPortOrOptions.preferredPort; + this.callbackPath = preferredPortOrOptions.callbackPath ?? CALLBACK_PATH; + this.callbackHostname = preferredPortOrOptions.callbackHostname ?? DEFAULT_HOSTNAME; + this.redirectUri = preferredPortOrOptions.redirectUri; } /** @@ -95,17 +118,22 @@ export abstract class OAuthCallbackFlow { * Start callback server, trying preferred port first, falling back to random. */ async #startCallbackServer(expectedState: string): Promise<{ server: Bun.Server; redirectUri: string }> { - // Try preferred port first try { - const redirectUri = `http://${DEFAULT_HOSTNAME}:${this.preferredPort}${this.callbackPath}`; const server = this.#createServer(this.preferredPort, expectedState); + if (this.redirectUri) { + return { server, redirectUri: this.redirectUri }; + } + const redirectUri = `http://${this.callbackHostname}:${this.preferredPort}${this.callbackPath}`; return { server, redirectUri }; } catch { - // Port busy or unavailable, try random port - const randomPort = 0; // Let OS assign - const server = this.#createServer(randomPort, expectedState); + if (this.redirectUri) { + throw new Error( + `OAuth callback port ${this.preferredPort} unavailable; cannot fall back to a random port when oauth.redirectUri is set`, + ); + } + const server = this.#createServer(0, expectedState); const actualPort = server.port; - const redirectUri = `http://${DEFAULT_HOSTNAME}:${actualPort}${this.callbackPath}`; + const redirectUri = `http://${this.callbackHostname}:${actualPort}${this.callbackPath}`; this.ctrl.onProgress?.(`Preferred port ${this.preferredPort} unavailable, using port ${actualPort}`); return { server, redirectUri }; } @@ -116,7 +144,7 @@ export abstract class OAuthCallbackFlow { */ #createServer(port: number, expectedState: string): Bun.Server { return Bun.serve({ - hostname: DEFAULT_HOSTNAME, + hostname: this.callbackHostname, port, reusePort: false, fetch: req => this.#handleCallback(req, expectedState), diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 7a2a976de..912a62143 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -2,6 +2,10 @@ ## [Unreleased] +### Fixed + +- Added `oauth.redirectUri`, `oauth.clientSecret`, and `oauth.callbackPath` support for MCP server OAuth config so providers can use exact registered redirect URIs while preserving local callback listener settings ([#445](https://github.com/can1357/oh-my-pi/issues/445)) + ## [13.12.8] - 2026-03-16 ### Breaking Changes diff --git a/packages/coding-agent/src/capability/mcp.ts b/packages/coding-agent/src/capability/mcp.ts index d4cc52a64..f0d2b97c3 100644 --- a/packages/coding-agent/src/capability/mcp.ts +++ b/packages/coding-agent/src/capability/mcp.ts @@ -31,11 +31,17 @@ export interface MCPServer { auth?: { type: "oauth" | "apikey"; credentialId?: string; + tokenUrl?: string; + clientId?: string; + clientSecret?: string; }; - /** OAuth configuration (clientId, callbackPort) for servers requiring explicit client credentials */ + /** OAuth configuration (clientId, clientSecret, redirectUri, callbackPort, callbackPath) for servers requiring explicit client credentials */ oauth?: { clientId?: string; + clientSecret?: string; + redirectUri?: string; callbackPort?: number; + callbackPath?: string; }; /** Transport type */ transport?: "stdio" | "sse" | "http"; diff --git a/packages/coding-agent/src/discovery/builtin.ts b/packages/coding-agent/src/discovery/builtin.ts index 7b8cb35f8..745142da9 100644 --- a/packages/coding-agent/src/discovery/builtin.ts +++ b/packages/coding-agent/src/discovery/builtin.ts @@ -160,8 +160,24 @@ async function loadMCPServers(ctx: LoadContext): Promise> env: serverConfig.env as Record | undefined, url: serverConfig.url as string | undefined, headers: serverConfig.headers as Record | undefined, - auth: serverConfig.auth as { type: "oauth" | "apikey"; credentialId?: string } | undefined, - oauth: serverConfig.oauth as { clientId?: string; callbackPort?: number } | undefined, + auth: serverConfig.auth as + | { + type: "oauth" | "apikey"; + credentialId?: string; + tokenUrl?: string; + clientId?: string; + clientSecret?: string; + } + | undefined, + oauth: serverConfig.oauth as + | { + clientId?: string; + clientSecret?: string; + redirectUri?: string; + callbackPort?: number; + callbackPath?: string; + } + | undefined, transport: serverConfig.type as "stdio" | "sse" | "http" | undefined, _source: createSourceMeta(PROVIDER_ID, path, level), }); diff --git a/packages/coding-agent/src/discovery/mcp-json.ts b/packages/coding-agent/src/discovery/mcp-json.ts index 4bee99abd..5b27f594a 100644 --- a/packages/coding-agent/src/discovery/mcp-json.ts +++ b/packages/coding-agent/src/discovery/mcp-json.ts @@ -34,9 +34,18 @@ interface MCPConfigFile { auth?: { type: "oauth" | "apikey"; credentialId?: string; + tokenUrl?: string; + clientId?: string; + clientSecret?: string; }; type?: "stdio" | "sse" | "http"; - oauth?: { clientId?: string; callbackPort?: number }; + oauth?: { + clientId?: string; + clientSecret?: string; + redirectUri?: string; + callbackPort?: number; + callbackPath?: string; + }; } >; } @@ -93,7 +102,8 @@ function transformMCPConfig(config: MCPConfigFile, source: SourceMeta): MCPServe if (server.env) server.env = expandEnvVarsDeep(server.env); if (server.url) server.url = expandEnvVarsDeep(server.url); if (server.headers) server.headers = expandEnvVarsDeep(server.headers); - + if (server.auth) server.auth = expandEnvVarsDeep(server.auth); + if (server.oauth) server.oauth = expandEnvVarsDeep(server.oauth); servers.push(server); } } diff --git a/packages/coding-agent/src/mcp/oauth-flow.ts b/packages/coding-agent/src/mcp/oauth-flow.ts index 7441b2c90..a4c9d513d 100644 --- a/packages/coding-agent/src/mcp/oauth-flow.ts +++ b/packages/coding-agent/src/mcp/oauth-flow.ts @@ -6,11 +6,93 @@ */ import type { OAuthController, OAuthCredentials } from "@oh-my-pi/pi-ai"; +import type { OAuthCallbackFlowOptions } from "@oh-my-pi/pi-ai/utils/oauth/callback-server"; import { OAuthCallbackFlow } from "@oh-my-pi/pi-ai/utils/oauth/callback-server"; const DEFAULT_PORT = 3000; const CALLBACK_PATH = "/callback"; +function isLoopbackHostname(hostname: string): boolean { + return hostname === "localhost" || hostname === "127.0.0.1"; +} + +function resolveRedirectUri(redirectUri: string | undefined): string | undefined { + const trimmed = redirectUri?.trim(); + if (!trimmed) return undefined; + + const parsed = new URL(trimmed); + if (parsed.protocol !== "http:" && parsed.protocol !== "https:") { + throw new Error("OAuth redirect URI must use http or https"); + } + return parsed.toString(); +} + +function parseRedirectUri(redirectUri: string | undefined): URL | undefined { + return redirectUri ? new URL(redirectUri) : undefined; +} + +function getUriPort(uri: URL): number { + if (uri.port !== "") return Number(uri.port); + return uri.protocol === "https:" ? 443 : 80; +} + +function validateRedirectConfig(config: MCPOAuthConfig, redirectUri: string | undefined): void { + const parsed = parseRedirectUri(redirectUri); + if (!parsed || parsed.protocol !== "https:" || !isLoopbackHostname(parsed.hostname)) { + return; + } + + if (config.callbackPort === undefined) { + throw new Error( + "HTTPS loopback redirect URIs require oauth.callbackPort to point at the local HTTP callback listener behind your TLS terminator", + ); + } + + if (config.callbackPort === getUriPort(parsed)) { + throw new Error( + "HTTPS loopback redirect URIs cannot reuse the same local port; terminate TLS separately and forward to oauth.callbackPort", + ); + } +} + +function resolveCallbackPort(callbackPort: number | undefined, redirectUri: string | undefined): number { + if (callbackPort !== undefined) return callbackPort; + + const parsed = parseRedirectUri(redirectUri); + if (!parsed || parsed.protocol !== "http:" || !isLoopbackHostname(parsed.hostname)) { + return DEFAULT_PORT; + } + + const port = getUriPort(parsed); + return Number.isFinite(port) && port > 0 ? port : DEFAULT_PORT; +} + +function resolveCallbackPath(callbackPath: string | undefined, redirectUri: string | undefined): string { + const trimmed = callbackPath?.trim(); + if (trimmed) return trimmed.startsWith("/") ? trimmed : `/${trimmed}`; + + const parsed = parseRedirectUri(redirectUri); + if (parsed?.pathname) return parsed.pathname; + return CALLBACK_PATH; +} + +function resolveCallbackHostname(redirectUri: string | undefined): string | undefined { + const parsed = parseRedirectUri(redirectUri); + if (!parsed || !isLoopbackHostname(parsed.hostname)) return undefined; + return parsed.hostname; +} + +function resolveCallbackOptions(config: MCPOAuthConfig): OAuthCallbackFlowOptions { + const redirectUri = resolveRedirectUri(config.redirectUri); + validateRedirectConfig(config, redirectUri); + return { + preferredPort: resolveCallbackPort(config.callbackPort, redirectUri), + callbackPath: resolveCallbackPath(config.callbackPath, redirectUri), + callbackHostname: resolveCallbackHostname(redirectUri), + redirectUri, + }; +} + export interface MCPOAuthConfig { /** Authorization endpoint URL */ authorizationUrl: string; @@ -22,8 +104,12 @@ export interface MCPOAuthConfig { clientSecret?: string; /** OAuth scopes (space-separated) */ scopes?: string; + /** Exact redirect URI to advertise to the provider */ + redirectUri?: string; /** Custom callback port (default: 3000) */ callbackPort?: number; + /** Custom callback path (default: /callback or redirectUri pathname) */ + callbackPath?: string; } /** @@ -39,7 +125,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { private config: MCPOAuthConfig, ctrl: OAuthController, ) { - super(ctrl, config.callbackPort ?? DEFAULT_PORT, CALLBACK_PATH); + super(ctrl, resolveCallbackOptions(config)); this.#resolvedClientId = this.#resolveClientId(config); } diff --git a/packages/coding-agent/src/mcp/types.ts b/packages/coding-agent/src/mcp/types.ts index f2b786213..01a126153 100644 --- a/packages/coding-agent/src/mcp/types.ts +++ b/packages/coding-agent/src/mcp/types.ts @@ -68,7 +68,10 @@ interface MCPServerConfigBase { /** OAuth configuration for servers requiring explicit client credentials */ oauth?: { clientId?: string; + clientSecret?: string; + redirectUri?: string; callbackPort?: number; + callbackPath?: string; }; } diff --git a/packages/coding-agent/src/modes/controllers/mcp-command-controller.ts b/packages/coding-agent/src/modes/controllers/mcp-command-controller.ts index ba3ec9689..47e96595c 100644 --- a/packages/coding-agent/src/modes/controllers/mcp-command-controller.ts +++ b/packages/coding-agent/src/modes/controllers/mcp-command-controller.ts @@ -32,7 +32,7 @@ import { searchSmitheryRegistry, toConfigName, } from "../../mcp/smithery-registry"; -import type { MCPServerConfig, MCPServerConnection } from "../../mcp/types"; +import type { MCPAuthConfig, MCPServerConfig, MCPServerConnection } from "../../mcp/types"; import type { OAuthCredential } from "../../session/auth-storage"; import { shortenPath } from "../../tools/render-utils"; import { openPath } from "../../utils/open"; @@ -400,13 +400,16 @@ export class MCPCommandController { } try { + const oauthClientSecret = finalConfig.oauth?.clientSecret ?? ""; const credentialId = await this.#handleOAuthFlow( oauth.authorizationUrl, oauth.tokenUrl, oauth.clientId ?? finalConfig.oauth?.clientId ?? "", - "", + oauthClientSecret, oauth.scopes ?? "", finalConfig.oauth?.callbackPort, + finalConfig.oauth?.callbackPath, + finalConfig.oauth?.redirectUri, ); finalConfig = { ...finalConfig, @@ -415,7 +418,7 @@ export class MCPCommandController { credentialId, tokenUrl: oauth.tokenUrl, clientId: oauth.clientId ?? finalConfig.oauth?.clientId, - clientSecret: undefined, + clientSecret: finalConfig.oauth?.clientSecret, }, }; } catch (oauthError) { @@ -478,6 +481,8 @@ export class MCPCommandController { clientSecret: string, scopes: string, callbackPort?: number, + callbackPath?: string, + redirectUri?: string, ): Promise { const authStorage = this.ctx.session.modelRegistry.authStorage; let parsedAuthUrl: URL; @@ -493,6 +498,7 @@ export class MCPCommandController { } const resolvedClientId = clientId.trim() || parsedAuthUrl.searchParams.get("client_id") || undefined; + const resolvedClientSecret = clientSecret.trim() || undefined; try { // Create OAuth flow @@ -501,9 +507,11 @@ export class MCPCommandController { authorizationUrl: authUrl, tokenUrl: tokenUrl, clientId: resolvedClientId, - clientSecret: clientSecret || undefined, + clientSecret: resolvedClientSecret, scopes: scopes || undefined, + redirectUri, callbackPort, + callbackPath, }, { onAuth: (info: { url: string; instructions?: string }) => { @@ -653,7 +661,7 @@ export class MCPCommandController { } #stripOAuthAuth(config: MCPServerConfig): MCPServerConfig { - const next = { ...config } as MCPServerConfig & { auth?: { type: "oauth" | "apikey"; credentialId?: string } }; + const next = { ...config } as MCPServerConfig & { auth?: MCPAuthConfig }; delete next.auth; return next; } @@ -1261,9 +1269,7 @@ export class MCPCommandController { return; } - const currentAuth = ( - found.config as MCPServerConfig & { auth?: { type: "oauth" | "apikey"; credentialId?: string } } - ).auth; + const currentAuth = (found.config as MCPServerConfig & { auth?: MCPAuthConfig }).auth; if (currentAuth?.type === "oauth") { await this.#removeManagedOAuthCredential(currentAuth.credentialId); } @@ -1298,17 +1304,14 @@ export class MCPCommandController { return; } - const currentAuth = ( - found.config as MCPServerConfig & { - auth?: { type: "oauth" | "apikey"; credentialId?: string; clientSecret?: string }; - } - ).auth; + const currentAuth = (found.config as MCPServerConfig & { auth?: MCPAuthConfig }).auth; if (currentAuth?.type === "oauth") { await this.#removeManagedOAuthCredential(currentAuth.credentialId); } const baseConfig = this.#stripOAuthAuth(found.config); const oauth = await this.#resolveOAuthEndpointsFromServer(baseConfig); + const oauthClientSecret = found.config.oauth?.clientSecret ?? currentAuth?.clientSecret ?? ""; this.#showMessage(["", theme.fg("muted", `Reauthorizing "${name}"...`), ""].join("\n")); @@ -1316,9 +1319,11 @@ export class MCPCommandController { oauth.authorizationUrl, oauth.tokenUrl, oauth.clientId ?? found.config.oauth?.clientId ?? "", - "", + oauthClientSecret, oauth.scopes ?? "", found.config.oauth?.callbackPort, + found.config.oauth?.callbackPath, + found.config.oauth?.redirectUri, ); const updated: MCPServerConfig = { @@ -1328,7 +1333,7 @@ export class MCPCommandController { credentialId, tokenUrl: oauth.tokenUrl, clientId: oauth.clientId ?? found.config.oauth?.clientId, - clientSecret: currentAuth?.clientSecret, + clientSecret: oauthClientSecret || undefined, }, }; await updateMCPServer(found.filePath, name, updated); diff --git a/packages/coding-agent/test/discovery/mcp-json.test.ts b/packages/coding-agent/test/discovery/mcp-json.test.ts new file mode 100644 index 000000000..48b59ee9c --- /dev/null +++ b/packages/coding-agent/test/discovery/mcp-json.test.ts @@ -0,0 +1,123 @@ +import { afterEach, beforeEach, describe, expect, test } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { mcpCapability, type MCPServer } from "@oh-my-pi/pi-coding-agent/capability/mcp"; +import { loadCapability } from "@oh-my-pi/pi-coding-agent/discovery"; + +async function loadStandaloneMcpConfig(cwd: string): Promise { + const result = await loadCapability(mcpCapability.id, { + cwd, + providers: ["mcp-json"], + }); + return result.items; +} + +describe("standalone mcp.json oauth env expansion", () => { + let tempDir = ""; + const originalEnv = { + PI_OAUTH_TOKEN_URL: process.env.PI_OAUTH_TOKEN_URL, + PI_OAUTH_CLIENT_ID: process.env.PI_OAUTH_CLIENT_ID, + PI_OAUTH_CLIENT_SECRET: process.env.PI_OAUTH_CLIENT_SECRET, + PI_OAUTH_REDIRECT_URI: process.env.PI_OAUTH_REDIRECT_URI, + PI_OAUTH_CALLBACK_PATH: process.env.PI_OAUTH_CALLBACK_PATH, + PI_MCP_HEADER: process.env.PI_MCP_HEADER, + PI_MCP_URL: process.env.PI_MCP_URL, + PI_MCP_ENV: process.env.PI_MCP_ENV, + }; + + beforeEach(async () => { + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "omp-mcp-json-")); + process.env.PI_OAUTH_TOKEN_URL = "https://provider.example/token"; + process.env.PI_OAUTH_CLIENT_ID = "oauth-client-id"; + process.env.PI_OAUTH_CLIENT_SECRET = "oauth-client-secret"; + process.env.PI_OAUTH_REDIRECT_URI = "https://public.example/oauth/callback"; + process.env.PI_OAUTH_CALLBACK_PATH = "/oauth/callback"; + process.env.PI_MCP_HEADER = "Bearer test-token"; + process.env.PI_MCP_URL = "https://mcp.example.com"; + process.env.PI_MCP_ENV = "env-value"; + }); + + afterEach(async () => { + await fs.rm(tempDir, { recursive: true, force: true }); + for (const [key, value] of Object.entries(originalEnv)) { + if (value === undefined) { + delete process.env[key]; + } else { + process.env[key] = value; + } + } + }); + + test("expands standalone auth and oauth fields alongside existing env-expanded fields", async () => { + await fs.writeFile( + path.join(tempDir, "mcp.json"), + JSON.stringify({ + mcpServers: { + figma: { + url: "${PI_MCP_URL}/mcp", + headers: { Authorization: "${PI_MCP_HEADER}" }, + env: { MCP_VALUE: "${PI_MCP_ENV}" }, + auth: { + type: "oauth", + tokenUrl: "${PI_OAUTH_TOKEN_URL}", + clientId: "${PI_OAUTH_CLIENT_ID}", + clientSecret: "${PI_OAUTH_CLIENT_SECRET}", + }, + oauth: { + clientId: "${PI_OAUTH_CLIENT_ID}", + clientSecret: "${PI_OAUTH_CLIENT_SECRET}", + redirectUri: "${PI_OAUTH_REDIRECT_URI}", + callbackPort: 4317, + callbackPath: "${PI_OAUTH_CALLBACK_PATH}", + }, + }, + }, + }), + ); + + const [server] = await loadStandaloneMcpConfig(tempDir); + expect(server).toBeDefined(); + expect(server?.url).toBe("https://mcp.example.com/mcp"); + expect(server?.headers).toEqual({ Authorization: "Bearer test-token" }); + expect(server?.env).toEqual({ MCP_VALUE: "env-value" }); + expect(server?.auth).toEqual({ + type: "oauth", + tokenUrl: "https://provider.example/token", + clientId: "oauth-client-id", + clientSecret: "oauth-client-secret", + }); + expect(server?.oauth).toEqual({ + clientId: "oauth-client-id", + clientSecret: "oauth-client-secret", + redirectUri: "https://public.example/oauth/callback", + callbackPort: 4317, + callbackPath: "/oauth/callback", + }); + }); + + test("expands only the standalone oauth fields that are present", async () => { + await fs.writeFile( + path.join(tempDir, ".mcp.json"), + JSON.stringify({ + mcpServers: { + slack: { + url: "https://slack.example.com/mcp", + oauth: { + redirectUri: "${PI_OAUTH_REDIRECT_URI}", + callbackPath: "${PI_OAUTH_CALLBACK_PATH}", + }, + }, + }, + }), + ); + + const [server] = await loadStandaloneMcpConfig(tempDir); + expect(server).toBeDefined(); + expect(server?.oauth).toEqual({ + redirectUri: "https://public.example/oauth/callback", + callbackPath: "/oauth/callback", + }); + expect(server?.auth).toBeUndefined(); + }); +}); diff --git a/packages/coding-agent/test/oauth-flow.test.ts b/packages/coding-agent/test/oauth-flow.test.ts index fe1e9cc14..e5f9dd98b 100644 --- a/packages/coding-agent/test/oauth-flow.test.ts +++ b/packages/coding-agent/test/oauth-flow.test.ts @@ -1,6 +1,13 @@ -import { describe, expect, it } from "bun:test"; -import { MCPOAuthFlow } from "@oh-my-pi/pi-coding-agent/mcp/oauth-flow"; -import { hookFetch } from "@oh-my-pi/pi-utils"; +import { afterEach, describe, expect, it, vi } from "bun:test"; +import { MCPOAuthFlow } from "../src/mcp/oauth-flow"; +import { hookFetch } from "../../utils/src/hook-fetch"; + +const originalFetch = global.fetch; + +afterEach(() => { + vi.restoreAllMocks(); + global.fetch = originalFetch; +}); describe("mcp oauth flow", () => { it("uses Codex client name for dynamic client registration", async () => { @@ -45,4 +52,242 @@ describe("mcp oauth flow", () => { expect(authUrl.searchParams.get("client_id")).toBe("registered-client-id"); expect(authUrl.searchParams.get("state")).toBe("test-state"); }); + + it("uses configured callbackPath for the local redirect URI", async () => { + let observedRedirectUri = ""; + let tokenRequestBody = ""; + + using _hook = hookFetch((input, init) => { + const url = String(input); + if (url === "https://provider.example/token") { + tokenRequestBody = String(init?.body ?? ""); + return new Response( + JSON.stringify({ + access_token: "access-token", + refresh_token: "refresh-token", + expires_in: 3600, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); + } + + throw new Error(`Unexpected fetch: ${url}`); + }); + + const flow = new MCPOAuthFlow( + { + authorizationUrl: "https://provider.example/authorize", + tokenUrl: "https://provider.example/token", + clientId: "client-id", + callbackPort: 14567, + callbackPath: "slack/oauth_redirect", + }, + { + onAuth: info => { + const authUrl = new URL(info.url); + observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? ""; + const state = authUrl.searchParams.get("state") ?? ""; + queueMicrotask(() => { + void originalFetch(`${observedRedirectUri}?code=test-code&state=${state}`); + }); + }, + signal: AbortSignal.timeout(1_000), + }, + ); + + const credentials = await flow.login(); + const redirectUrl = new URL(observedRedirectUri); + const tokenParams = new URLSearchParams(tokenRequestBody); + + expect(redirectUrl.pathname).toBe("/slack/oauth_redirect"); + expect(tokenParams.get("redirect_uri")).toBe(observedRedirectUri); + expect(credentials).toMatchObject({ + access: "access-token", + refresh: "refresh-token", + }); + }); + + it("uses exact redirectUri and clientSecret for provider requests", async () => { + let observedRedirectUri = ""; + let tokenRequestBody = ""; + + using _hook = hookFetch((input, init) => { + const url = String(input); + if (url === "https://provider.example/token") { + tokenRequestBody = String(init?.body ?? ""); + return new Response( + JSON.stringify({ + access_token: "access-token", + refresh_token: "refresh-token", + expires_in: 3600, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); + } + + throw new Error(`Unexpected fetch: ${url}`); + }); + + const flow = new MCPOAuthFlow( + { + authorizationUrl: "https://provider.example/authorize", + tokenUrl: "https://provider.example/token", + clientId: "client-id", + clientSecret: "client-secret", + redirectUri: "https://public.example/slack/oauth_redirect", + callbackPort: 14568, + callbackPath: "slack/oauth_redirect", + }, + { + onAuth: info => { + const authUrl = new URL(info.url); + observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? ""; + const state = authUrl.searchParams.get("state") ?? ""; + queueMicrotask(() => { + void originalFetch(`http://localhost:14568/slack/oauth_redirect?code=test-code&state=${state}`); + }); + }, + signal: AbortSignal.timeout(1_000), + }, + ); + + const credentials = await flow.login(); + const tokenParams = new URLSearchParams(tokenRequestBody); + + expect(observedRedirectUri).toBe("https://public.example/slack/oauth_redirect"); + expect(tokenParams.get("redirect_uri")).toBe("https://public.example/slack/oauth_redirect"); + expect(tokenParams.get("client_secret")).toBe("client-secret"); + expect(credentials).toMatchObject({ + access: "access-token", + refresh: "refresh-token", + }); + }); + + it("supports https loopback redirectUri values behind a separate local callback port", async () => { + let observedRedirectUri = ""; + let tokenRequestBody = ""; + + using _hook = hookFetch((input, init) => { + const url = String(input); + if (url === "https://provider.example/token") { + tokenRequestBody = String(init?.body ?? ""); + return new Response( + JSON.stringify({ + access_token: "access-token", + refresh_token: "refresh-token", + expires_in: 3600, + }), + { status: 200, headers: { "Content-Type": "application/json" } }, + ); + } + + throw new Error(`Unexpected fetch: ${url}`); + }); + + const flow = new MCPOAuthFlow( + { + authorizationUrl: "https://provider.example/authorize", + tokenUrl: "https://provider.example/token", + redirectUri: "https://localhost:3443/slack/oauth_redirect", + callbackPort: 14570, + }, + { + onAuth: info => { + const authUrl = new URL(info.url); + observedRedirectUri = authUrl.searchParams.get("redirect_uri") ?? ""; + const state = authUrl.searchParams.get("state") ?? ""; + queueMicrotask(() => { + void originalFetch(`http://localhost:14570/slack/oauth_redirect?code=test-code&state=${state}`); + }); + }, + signal: AbortSignal.timeout(1_000), + }, + ); + + const credentials = await flow.login(); + const tokenParams = new URLSearchParams(tokenRequestBody); + + expect(observedRedirectUri).toBe("https://localhost:3443/slack/oauth_redirect"); + expect(tokenParams.get("redirect_uri")).toBe("https://localhost:3443/slack/oauth_redirect"); + expect(credentials).toMatchObject({ + access: "access-token", + refresh: "refresh-token", + }); + }); + + it("rejects https loopback redirectUri values without a separate callback port", () => { + expect( + () => + new MCPOAuthFlow( + { + authorizationUrl: "https://provider.example/authorize", + tokenUrl: "https://provider.example/token", + redirectUri: "https://localhost:3000/slack/oauth_redirect", + }, + {}, + ), + ).toThrow("HTTPS loopback redirect URIs require oauth.callbackPort"); + }); + + it("listens on the implied port for exact HTTP loopback redirectUri values", async () => { + const serveSpy = vi.spyOn(Bun, "serve").mockImplementation(options => { + expect(options.port).toBe(80); + throw new Error("EADDRINUSE"); + }); + + const flow = new MCPOAuthFlow( + { + authorizationUrl: "https://provider.example/authorize", + tokenUrl: "https://provider.example/token", + redirectUri: "http://localhost/callback", + }, + { signal: AbortSignal.timeout(1_000) }, + ); + + await expect(flow.login()).rejects.toThrow( + "OAuth callback port 80 unavailable; cannot fall back to a random port when oauth.redirectUri is set", + ); + expect(serveSpy).toHaveBeenCalledTimes(1); + }); + + it("listens on the explicit port for exact HTTP loopback redirectUri values", async () => { + const serveSpy = vi.spyOn(Bun, "serve").mockImplementation(options => { + expect(options.port).toBe(3000); + throw new Error("EADDRINUSE"); + }); + + const flow = new MCPOAuthFlow( + { + authorizationUrl: "https://provider.example/authorize", + tokenUrl: "https://provider.example/token", + redirectUri: "http://localhost:3000/callback", + }, + { signal: AbortSignal.timeout(1_000) }, + ); + + await expect(flow.login()).rejects.toThrow( + "OAuth callback port 3000 unavailable; cannot fall back to a random port when oauth.redirectUri is set", + ); + expect(serveSpy).toHaveBeenCalledTimes(1); + }); + + + it("fails instead of falling back to a random port when redirectUri is exact", async () => { + vi.spyOn(Bun, "serve").mockImplementation(() => { + throw new Error("EADDRINUSE"); + }); + + const flow = new MCPOAuthFlow( + { + authorizationUrl: "https://provider.example/authorize", + tokenUrl: "https://provider.example/token", + redirectUri: "https://public.example/slack/oauth_redirect", + callbackPort: 14569, + callbackPath: "/slack/oauth_redirect", + }, + { signal: AbortSignal.timeout(1_000) }, + ); + + await expect(flow.login()).rejects.toThrow("cannot fall back to a random port when oauth.redirectUri is set"); + }); });