diff --git a/src/services/mcp/McpHub.ts b/src/services/mcp/McpHub.ts index 95f7e55f37..a27053b564 100644 --- a/src/services/mcp/McpHub.ts +++ b/src/services/mcp/McpHub.ts @@ -1111,20 +1111,32 @@ export class McpHub { return new Promise((resolve) => { let disposed = false + let cancellationDisposable: vscode.Disposable | undefined const cleanup = () => { if (disposed) return disposed = true clearTimeout(timeoutHandle) unsubscribe() + cancellationDisposable?.dispose() this._oauthWatchers.delete(watcherKey) } // --- Cross-window token watcher --- const onTokensChanged = async () => { if (disposed || this.isDisposed) return + // Don't reconnect if the server is already connected (e.g. a token refresh + // from another window fired after _completeOAuthFlow already succeeded). + const conn = this.findConnection(name, source) + if (!conn || conn.server.status === "connected") { + cleanup() + return + } try { const data = await this.secretStorage?.getOAuthData(serverUrl) + // Re-check after the async yield — cancellation or completion may have + // fired while getOAuthData was in flight. + if (disposed || this.isDisposed) return if (data && Date.now() < data.expires_at - TOKEN_EXPIRY_BUFFER_MS) { cleanup() await authProvider.close() @@ -1144,7 +1156,10 @@ export class McpHub { }) // --- Cancellation --- - cancellationToken.onCancellationRequested(() => { + cancellationDisposable = cancellationToken.onCancellationRequested(() => { + // Guard: flow may have already completed (e.g. tokens arrived via + // onTokensChanged) by the time VS Code fires this callback. + if (disposed) return cleanup() void authProvider.close() const conn = this.findConnection(name, source) @@ -1189,6 +1204,7 @@ export class McpHub { if (choice === authenticateLabel) { // Guard: another window may have authed while the toast was showing const tokens = await this.secretStorage!.getOAuthData(serverUrl) + if (disposed) return if (tokens && Date.now() < tokens.expires_at - TOKEN_EXPIRY_BUFFER_MS) { cleanup() await authProvider.close() @@ -1206,6 +1222,10 @@ export class McpHub { } catch { // _completeOAuthFlow handles its own error state } + // Cancellation may have fired while _completeOAuthFlow was running. + // If so, the cancellation handler already cleaned up and resolved — + // don't overwrite that state with a stale error. + if (disposed) return cleanup() resolve() return @@ -2414,12 +2434,13 @@ export class McpHub { this.isProgrammaticUpdate = false - // Cancel all active OAuth token watchers + // Cancel all active OAuth token watchers and in-flight reauth promises for (const { unsubscribe, abortHandle } of this._oauthWatchers.values()) { unsubscribe() clearTimeout(abortHandle) } this._oauthWatchers.clear() + this.reauthPromises.clear() this.removeAllFileWatchers() diff --git a/src/services/mcp/McpOAuthClientProvider.ts b/src/services/mcp/McpOAuthClientProvider.ts index 0b45146f2d..fe070ffd3c 100644 --- a/src/services/mcp/McpOAuthClientProvider.ts +++ b/src/services/mcp/McpOAuthClientProvider.ts @@ -129,7 +129,7 @@ export class McpOAuthClientProvider implements OAuthClientProvider { const scopes: string[] = authServerMeta?.scopes_supported ?? [] // Generate a CSRF state token for the OAuth flow. - const state = Array.from(crypto.getRandomValues(new Uint8Array(8))) + const state = Array.from(crypto.getRandomValues(new Uint8Array(16))) .map((b) => b.toString(16).padStart(2, "0")) .join("") @@ -333,6 +333,17 @@ export class McpOAuthClientProvider implements OAuthClientProvider { if (this._authServerMeta?.authorization_endpoint) { try { const fixed = new URL(this._authServerMeta.authorization_endpoint as string) + // Validate the authorization_endpoint origin matches the issuer to prevent + // a compromised metadata document from redirecting users to a phishing page. + const expectedOrigin = this._authServerMeta.issuer + ? new URL(this._authServerMeta.issuer as string).origin + : new URL(this._serverUrl).origin + if (fixed.origin !== expectedOrigin) { + // Fall through and use the SDK-supplied URL unchanged + throw new Error( + `authorization_endpoint origin mismatch: expected ${expectedOrigin}, got ${fixed.origin}`, + ) + } // Copy all query params generated by the SDK authorizationUrl.searchParams.forEach((value, key) => { fixed.searchParams.set(key, value) diff --git a/src/services/mcp/__tests__/McpHub.spec.ts b/src/services/mcp/__tests__/McpHub.spec.ts index 41cfb5ea07..e3dc23b032 100644 --- a/src/services/mcp/__tests__/McpHub.spec.ts +++ b/src/services/mcp/__tests__/McpHub.spec.ts @@ -2558,7 +2558,9 @@ describe("McpHub", () => { ) // _initiateOAuthFlow awaits getOAuthData() before calling withProgress. - // Flush that microtask so withProgress runs and capturedCancellationToken is set. + // Two ticks: tick 1 resolves getOAuthData, tick 2 runs the continuation + // that calls withProgress, setting capturedCancellationToken. + await Promise.resolve() await Promise.resolve() capturedCancellationToken._fire() @@ -2591,11 +2593,13 @@ describe("McpHub", () => { }) it("should resolve without calling _completeOAuthFlow when tokens exist at click time", async () => { - // Tokens are present when Authenticate is clicked (cross-window guard in the loop) + // Tokens are present when Authenticate is clicked (click-time guard in the loop). + // First call (pre-withProgress early-return check) returns null so withProgress runs. + // Second call (after click) returns valid tokens, exercising the click-time guard. vsc.window.showInformationMessage.mockResolvedValueOnce(t("mcp:oauth.flow.authenticateButton") as any) - mockSecretStorage.getOAuthData.mockResolvedValue({ - expires_at: Date.now() + 10 * 60 * 1000, - }) + mockSecretStorage.getOAuthData + .mockResolvedValueOnce(null) // pre-check: no tokens yet, flow proceeds to withProgress + .mockResolvedValue({ expires_at: Date.now() + 10 * 60 * 1000 }) // at click time const completeOAuthSpy = vi.spyOn(mcpHub as any, "_completeOAuthFlow") @@ -2629,5 +2633,80 @@ describe("McpHub", () => { ), ).resolves.toBeUndefined() }) + + it("should not reconnect when hub is disposed while cross-window tokens arrive", async () => { + vi.useFakeTimers() + vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {})) + + mockSecretStorage.onDidChange.mockImplementation((_key: string, cb: () => void) => { + Promise.resolve().then(() => { + // Dispose the hub before the token callback runs + ;(mcpHub as any).isDisposed = true + mockSecretStorage.getOAuthData.mockResolvedValue({ + expires_at: Date.now() + 10 * 60 * 1000, + }) + cb() + }) + return vi.fn() + }) + + const flowPromise = (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + // Timeout to unblock the flow (watcher bailed due to isDisposed) + await vi.advanceTimersByTimeAsync(OAUTH_FLOW_TIMEOUT_MS) + await flowPromise + + expect((mcpHub as any).connectToServer).not.toHaveBeenCalled() + }) + + it("should cancel the previous watcher when called again for the same server", async () => { + vi.useFakeTimers() + vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {})) + + // Start first flow (intentionally not awaited — the second call orphans it) + ;(mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + // Two ticks: getOAuthData resolves, then withProgress registers the watcher + await Promise.resolve() + await Promise.resolve() + expect((mcpHub as any)._oauthWatchers.size).toBe(1) + const firstEntry = (mcpHub as any)._oauthWatchers.get(`${serverName}:${source}`) + + // Start second flow for the same server — should tear down the first watcher + const secondFlow = (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + await Promise.resolve() + await Promise.resolve() + // Still exactly one watcher for this server key + expect((mcpHub as any)._oauthWatchers.size).toBe(1) + // Watcher entry was replaced (second flow's entry, not first) + const secondEntry = (mcpHub as any)._oauthWatchers.get(`${serverName}:${source}`) + expect(secondEntry).not.toBe(firstEntry) + + // Advance past timeout so the second flow resolves + await vi.advanceTimersByTimeAsync(OAUTH_FLOW_TIMEOUT_MS) + await secondFlow + }) }) }) diff --git a/src/services/mcp/utils/callbackServer.ts b/src/services/mcp/utils/callbackServer.ts index e447222b0d..2f12b79b96 100644 --- a/src/services/mcp/utils/callbackServer.ts +++ b/src/services/mcp/utils/callbackServer.ts @@ -76,7 +76,10 @@ export function startCallbackServer( // Verify state for CSRF protection if (expectedState && state !== expectedState) { - res.writeHead(400, { "Content-Type": "text/html" }) + res.writeHead(400, { + "Content-Type": "text/html", + "Content-Security-Policy": "default-src 'none'; style-src 'unsafe-inline'", + }) res.end(` @@ -90,11 +93,16 @@ export function startCallbackServer( `) rejectResult(new Error("Invalid state parameter")) + server.close() return } // Send HTML response - res.writeHead(200, { "Content-Type": "text/html" }) + res.writeHead(200, { + "Content-Type": "text/html", + "Content-Security-Policy": + "default-src 'none'; style-src 'unsafe-inline'; script-src 'unsafe-inline'", + }) res.end(` @@ -122,7 +130,6 @@ export function startCallbackServer( const isError = ${hasError ? "true" : "false"}; if (!isError) { let count = 5; - const countEl = document.getElementById('count'); const countdownEl = document.getElementById('countdown'); if (countdownEl) { countdownEl.innerHTML = \`${t("mcp:oauth.callback.tab_closing_in", { count: 5 })}\`; diff --git a/src/services/mcp/utils/oauth.ts b/src/services/mcp/utils/oauth.ts index 8bdf2452ce..e205626c00 100644 --- a/src/services/mcp/utils/oauth.ts +++ b/src/services/mcp/utils/oauth.ts @@ -44,11 +44,19 @@ export interface OAuthDiscoveryResult { * * Returns an {@link OAuthDiscoveryResult} on success, or `null` if any step fails. */ +const DISCOVERY_TIMEOUT_MS = 5_000 + export async function fetchOAuthAuthServerMetadata(serverUrl: string): Promise { try { // Step 1 – RFC 9728: resolve the authorization server issuer URL and // capture the resource indicator for RFC 8707. - const resourceMeta = await discoverOAuthProtectedResourceMetadata(serverUrl) + // The SDK does not accept an AbortSignal, so we race it against a timeout. + const resourceMeta = await Promise.race([ + discoverOAuthProtectedResourceMetadata(serverUrl), + new Promise((_, reject) => + setTimeout(() => reject(new Error("OAuth discovery timeout")), DISCOVERY_TIMEOUT_MS), + ), + ]) const authServers = resourceMeta.authorization_servers if (!authServers?.length) return null @@ -68,6 +76,7 @@ export async function fetchOAuthAuthServerMetadata(serverUrl: string): Promise