From 8eb8be93d3f64e4a95a5247f5a5a2d6c7ed99d3f Mon Sep 17 00:00:00 2001 From: Elliott de Launay Date: Thu, 2 Apr 2026 23:45:07 -0400 Subject: [PATCH] refactor(McpHub): making prompt long running or dismissable --- src/i18n/locales/en/mcp.json | 9 + src/services/mcp/McpHub.ts | 261 ++++++++++++--------- src/services/mcp/__tests__/McpHub.spec.ts | 262 ++++++++++++++++++++++ src/services/mcp/constants.ts | 1 + src/services/mcp/utils/callbackServer.ts | 18 +- 5 files changed, 435 insertions(+), 116 deletions(-) diff --git a/src/i18n/locales/en/mcp.json b/src/i18n/locales/en/mcp.json index 2735b03238..9325484ca8 100644 --- a/src/i18n/locales/en/mcp.json +++ b/src/i18n/locales/en/mcp.json @@ -36,6 +36,15 @@ "tab_closing_in": "This tab will attempt to close in {{count}}s...", "safe_to_close": "If the tab did not close, you can safely close it manually.", "invalid_state": "Error: Invalid state parameter" + }, + "flow": { + "authenticating": "MCP server \"{{name}}\" requires authentication", + "waitingForBrowser": "Complete sign-in in your browser...", + "clickAuthenticate": "MCP server \"{{name}}\" is waiting for authentication.", + "authenticateButton": "Authenticate", + "cancelled": "OAuth authentication was cancelled", + "timedOut": "OAuth authentication timed out", + "connected": "MCP server \"{{name}}\" connected successfully after OAuth authentication." } } } diff --git a/src/services/mcp/McpHub.ts b/src/services/mcp/McpHub.ts index 18ca2568fa..95f7e55f37 100644 --- a/src/services/mcp/McpHub.ts +++ b/src/services/mcp/McpHub.ts @@ -36,7 +36,7 @@ import { GlobalFileNames } from "../../shared/globalFileNames" import { UnauthorizedError } from "@modelcontextprotocol/sdk/client/auth.js" import { fileExistsAtPath } from "../../utils/fs" -import { TOKEN_EXPIRY_BUFFER_MS } from "./constants" +import { TOKEN_EXPIRY_BUFFER_MS, OAUTH_FLOW_TIMEOUT_MS } from "./constants" import { SecretStorageService } from "./SecretStorageService" import { McpOAuthClientProvider } from "./McpOAuthClientProvider" import { arePathsEqual, getWorkspacePath } from "../../utils/path" @@ -1064,37 +1064,162 @@ export class McpHub { return } - // Show a confirmation toast so the user can decide whether to authenticate. - const choice = await vscode.window.showInformationMessage( - `MCP server "${name}" requires authentication.`, - "Authenticate", + // Show a persistent progress notification for the duration of the OAuth flow. + // Inside, a looping showInformationMessage gives the user an "Authenticate" + // button that reappears if dismissed, so they can trigger browser auth at any time. + await vscode.window.withProgress( + { + location: vscode.ProgressLocation.Notification, + title: t("mcp:oauth.flow.authenticating", { name }), + cancellable: true, + }, + (progress, cancellationToken) => + this._runOAuthFlowWithProgress( + name, + source, + config, + serverUrl, + authProvider, + transport, + connection, + progress, + cancellationToken, + ), ) + } - if (choice === "Authenticate") { - // Check tokens again — another window may have authed while toast was showing - const tokens = await this.secretStorage.getOAuthData(serverUrl) - if (tokens && Date.now() < tokens.expires_at - TOKEN_EXPIRY_BUFFER_MS) { - await authProvider.close() - await this.deleteConnection(name, source) - await this.connectToServer(name, config, source) - await this.notifyWebviewOfServerChanges() - return - } - await this._completeOAuthFlow(authProvider, transport, connection, name, source) - } else { - // Toast was dismissed or auto-timed-out. - const tokens = await this.secretStorage.getOAuthData(serverUrl) - if (tokens && Date.now() < tokens.expires_at - TOKEN_EXPIRY_BUFFER_MS) { - await authProvider.close() - await this.deleteConnection(name, source) - await this.connectToServer(name, config, source) - await this.notifyWebviewOfServerChanges() - return - } - // Tokens not available yet — start watching for them - await authProvider.close() - this._watchForOAuthTokens(name, source, serverUrl, config) + private _runOAuthFlowWithProgress( + name: string, + source: "global" | "project", + config: z.infer, + serverUrl: string, + authProvider: McpOAuthClientProvider, + transport: StreamableHTTPClientTransport, + connection: ConnectedMcpConnection, + progress: vscode.Progress<{ message?: string; increment?: number }>, + cancellationToken: vscode.CancellationToken, + ): Promise { + const watcherKey = `${name}:${source}` + + // Cancel any existing watcher for this connection before starting a new one + const existing = this._oauthWatchers.get(watcherKey) + if (existing) { + existing.unsubscribe() + clearTimeout(existing.abortHandle) + this._oauthWatchers.delete(watcherKey) } + + return new Promise((resolve) => { + let disposed = false + + const cleanup = () => { + if (disposed) return + disposed = true + clearTimeout(timeoutHandle) + unsubscribe() + this._oauthWatchers.delete(watcherKey) + } + + // --- Cross-window token watcher --- + const onTokensChanged = async () => { + if (disposed || this.isDisposed) return + try { + const data = await this.secretStorage?.getOAuthData(serverUrl) + if (data && Date.now() < data.expires_at - TOKEN_EXPIRY_BUFFER_MS) { + cleanup() + await authProvider.close() + await this.deleteConnection(name, source) + const validatedConfig = this.validateServerConfig(config, name) + await this.connectToServer(name, validatedConfig, source) + await this.notifyWebviewOfServerChanges() + resolve() + } + } catch (err) { + console.error(`[McpHub] OAuth token watcher failed for "${name}":`, err) + } + } + + const unsubscribe = this.secretStorage!.onDidChange(serverUrl, () => { + void onTokensChanged() + }) + + // --- Cancellation --- + cancellationToken.onCancellationRequested(() => { + cleanup() + void authProvider.close() + const conn = this.findConnection(name, source) + if (conn && conn.server.status !== "connected") { + conn.server.status = "disconnected" + this.appendErrorMessage(conn, t("mcp:oauth.flow.cancelled")) + void this.notifyWebviewOfServerChanges() + } + resolve() + }) + + // --- Timeout --- + const timeoutHandle = setTimeout(() => { + if (disposed) return + cleanup() + void authProvider.close() + const conn = this.findConnection(name, source) + if (conn && conn.server.status === "connecting") { + conn.server.status = "disconnected" + this.appendErrorMessage(conn, t("mcp:oauth.flow.timedOut")) + void this.notifyWebviewOfServerChanges() + } + resolve() + }, OAUTH_FLOW_TIMEOUT_MS) + + // Register in _oauthWatchers so deleteConnection() and dispose() can clean up + this._oauthWatchers.set(watcherKey, { unsubscribe, abortHandle: timeoutHandle }) + + const authenticateLabel = t("mcp:oauth.flow.authenticateButton") + // --- Looping "Authenticate" toast --- + const showAuthPromptLoop = async () => { + while (!disposed) { + progress.report({ message: t("mcp:oauth.flow.waitingForBrowser") }) + + const choice = await vscode.window.showInformationMessage( + t("mcp:oauth.flow.clickAuthenticate", { name }), + authenticateLabel, + ) + + if (disposed) return + + if (choice === authenticateLabel) { + // Guard: another window may have authed while the toast was showing + const tokens = await this.secretStorage!.getOAuthData(serverUrl) + if (tokens && Date.now() < tokens.expires_at - TOKEN_EXPIRY_BUFFER_MS) { + cleanup() + await authProvider.close() + await this.deleteConnection(name, source) + const validatedFastPathConfig = this.validateServerConfig(config, name) + await this.connectToServer(name, validatedFastPathConfig, source) + await this.notifyWebviewOfServerChanges() + resolve() + return + } + + progress.report({ message: t("mcp:oauth.flow.waitingForBrowser") }) + try { + await this._completeOAuthFlow(authProvider, transport, connection, name, source) + } catch { + // _completeOAuthFlow handles its own error state + } + cleanup() + resolve() + return + } + + // Toast was dismissed or auto-timed-out — wait briefly then re-show + if (!disposed) { + await delay(3000) + } + } + } + + void showAuthPromptLoop() + }) } private async _completeOAuthFlow( @@ -1130,9 +1255,7 @@ export class McpHub { await this.connectToServer(name, validatedConfig, source) await this.notifyWebviewOfServerChanges() - void vscode.window.showInformationMessage( - `MCP server "${name}" connected successfully after OAuth authentication.`, - ) + void vscode.window.showInformationMessage(t("mcp:oauth.flow.connected", { name })) } catch (error) { await authProvider.close() const conn = this.findConnection(name, source) @@ -1144,80 +1267,6 @@ export class McpHub { } } - private _watchForOAuthTokens( - name: string, - source: "global" | "project", - serverUrl: string, - config: z.infer, - ): void { - if (!this.secretStorage) return - - const watcherKey = `${name}:${source}` - - // Cancel any existing watcher for this connection before starting a new one - const existing = this._oauthWatchers.get(watcherKey) - if (existing) { - existing.unsubscribe() - clearTimeout(existing.abortHandle) - this._oauthWatchers.delete(watcherKey) - } - - // Called when SecretStorage fires onDidChange for this server's key. - // Runs in all VS Code windows the instant tokens are saved — no polling delay. - const onTokensChanged = async () => { - try { - if (this.isDisposed) return - - const conn = this.findConnection(name, source) - if (!conn || conn.server.status === "connected") { - cleanup() - return - } - - const data = await this.secretStorage?.getOAuthData(serverUrl) - if (data && Date.now() < data.expires_at - TOKEN_EXPIRY_BUFFER_MS) { - cleanup() - await this.deleteConnection(name, source) - const validatedConfig = this.validateServerConfig(config, name) - await this.connectToServer(name, validatedConfig, source) - await this.notifyWebviewOfServerChanges() - } - // If tokens aren't valid yet (e.g. a delete event fired), keep listening. - } catch (err) { - console.error(`[McpHub] OAuth token watcher failed for "${name}":`, err) - } - } - - const cleanup = () => { - const entry = this._oauthWatchers.get(watcherKey) - if (entry) { - entry.unsubscribe() - clearTimeout(entry.abortHandle) - this._oauthWatchers.delete(watcherKey) - } - } - - const unsubscribe = this.secretStorage.onDidChange(serverUrl, () => { - void onTokensChanged() - }) - - // Give up after 6 minutes if no token ever arrives - const abortHandle = setTimeout( - () => { - cleanup() - const conn = this.findConnection(name, source) - if (conn && conn.server.status === "connecting") { - conn.server.status = "disconnected" - this.appendErrorMessage(conn, "OAuth authentication timed out waiting for another window") - void this.notifyWebviewOfServerChanges() - } - }, - 6 * 60 * 1000, - ) - - this._oauthWatchers.set(watcherKey, { unsubscribe, abortHandle }) - } - private appendErrorMessage(connection: McpConnection, error: string, level: "error" | "warn" | "info" = "error") { const MAX_ERROR_LENGTH = 1000 const truncatedError = diff --git a/src/services/mcp/__tests__/McpHub.spec.ts b/src/services/mcp/__tests__/McpHub.spec.ts index 3f06627cc1..41cfb5ea07 100644 --- a/src/services/mcp/__tests__/McpHub.spec.ts +++ b/src/services/mcp/__tests__/McpHub.spec.ts @@ -7,6 +7,8 @@ import type { ClineProvider } from "../../../core/webview/ClineProvider" import type { McpHub as McpHubType, McpConnection, ConnectedMcpConnection, DisconnectedMcpConnection } from "../McpHub" import { ServerConfigSchema, McpHub } from "../McpHub" +import { OAUTH_FLOW_TIMEOUT_MS } from "../constants" +import { t } from "../../../i18n" // Mock fs/promises before importing anything that uses it vi.mock("fs/promises", () => ({ @@ -49,6 +51,8 @@ vi.mock("../../../utils/safeWriteJson", () => ({ }), })) +vi.mock("delay", () => ({ default: vi.fn().mockResolvedValue(undefined) })) + vi.mock("vscode", () => ({ workspace: { createFileSystemWatcher: vi.fn().mockReturnValue({ @@ -68,6 +72,22 @@ vi.mock("vscode", () => ({ createTextEditorDecorationType: vi.fn().mockReturnValue({ dispose: vi.fn(), }), + withProgress: vi.fn().mockImplementation((_options: any, task: any) => { + const progress = { report: vi.fn() } + const tokenListeners: Array<() => void> = [] + const cancellationToken = { + isCancellationRequested: false, + onCancellationRequested: vi.fn((cb: () => void) => { + tokenListeners.push(cb) + return { dispose: vi.fn() } + }), + _fire: () => tokenListeners.forEach((cb) => cb()), + } + return task(progress, cancellationToken) + }), + }, + ProgressLocation: { + Notification: 15, }, Disposable: { from: vi.fn(), @@ -2368,4 +2388,246 @@ describe("McpHub", () => { ) }) }) + + describe("_initiateOAuthFlow with persistent notification", () => { + const serverName = "oauth-server" + const serverUrl = "https://example.com/mcp" + const source = "global" as const + const config = { url: serverUrl } + + let mockAuthProvider: any + let mockTransport: any + let mockConnection: any + let mockSecretStorage: any + let vsc: any + + beforeEach(async () => { + vi.clearAllMocks() + vsc = await import("vscode") + + mockAuthProvider = { + openBrowser: vi.fn().mockResolvedValue(undefined), + waitForAuthCode: vi.fn().mockResolvedValue("auth-code-123"), + exchangeCodeForTokens: vi.fn().mockResolvedValue(undefined), + close: vi.fn().mockResolvedValue(undefined), + } + + mockTransport = {} + + mockConnection = { + server: { + status: "connecting", + config: JSON.stringify(config), + name: serverName, + }, + } + + mockSecretStorage = { + getOAuthData: vi.fn().mockResolvedValue(null), + onDidChange: vi.fn().mockReturnValue(vi.fn()), + } + ;(mcpHub as any).secretStorage = mockSecretStorage + + vi.spyOn(mcpHub as any, "deleteConnection").mockResolvedValue(undefined) + vi.spyOn(mcpHub as any, "connectToServer").mockResolvedValue(undefined) + vi.spyOn(mcpHub as any, "notifyWebviewOfServerChanges").mockResolvedValue(undefined) + vi.spyOn(mcpHub as any, "findConnection").mockReturnValue(mockConnection) + vi.spyOn(mcpHub as any, "validateServerConfig").mockReturnValue(config) + vi.spyOn(mcpHub as any, "appendErrorMessage").mockReturnValue(undefined) + }) + + afterEach(() => { + vi.useRealTimers() + }) + + it("should use withProgress for persistent notification", async () => { + vsc.window.showInformationMessage.mockResolvedValueOnce(t("mcp:oauth.flow.authenticateButton") as any) + vi.spyOn(mcpHub as any, "_completeOAuthFlow").mockResolvedValue(undefined) + + await (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + expect(vsc.window.withProgress).toHaveBeenCalledWith( + expect.objectContaining({ + location: vsc.ProgressLocation.Notification, + cancellable: true, + }), + expect.any(Function), + ) + }) + + it("should re-show toast when dismissed and complete on second attempt", async () => { + // First toast dismissed (undefined), second toast user clicks Authenticate + vsc.window.showInformationMessage + .mockResolvedValueOnce(undefined as any) + .mockResolvedValueOnce(t("mcp:oauth.flow.authenticateButton") as any) + + vi.spyOn(mcpHub as any, "_completeOAuthFlow").mockResolvedValue(undefined) + + await (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + expect(vsc.window.showInformationMessage).toHaveBeenCalledTimes(2) + }) + + it("should resolve when cross-window tokens arrive", async () => { + vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {})) + + mockSecretStorage.onDidChange.mockImplementation((_key: string, cb: () => void) => { + Promise.resolve().then(() => { + mockSecretStorage.getOAuthData.mockResolvedValue({ + expires_at: Date.now() + 10 * 60 * 1000, + }) + cb() + }) + return vi.fn() + }) + + await (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + expect(mockAuthProvider.close).toHaveBeenCalled() + expect((mcpHub as any).deleteConnection).toHaveBeenCalledWith(serverName, source) + expect((mcpHub as any).connectToServer).toHaveBeenCalled() + }) + + it("should skip flow when valid tokens already exist", async () => { + mockSecretStorage.getOAuthData.mockResolvedValue({ + expires_at: Date.now() + 10 * 60 * 1000, + }) + + await (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + expect(vsc.window.withProgress).not.toHaveBeenCalled() + expect(mockAuthProvider.close).toHaveBeenCalled() + expect((mcpHub as any).deleteConnection).toHaveBeenCalledWith(serverName, source) + expect((mcpHub as any).connectToServer).toHaveBeenCalled() + }) + + it("should disconnect and flag error when user cancels the OAuth flow", async () => { + vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {})) + + // Override withProgress for this test to capture the cancellation token + let capturedCancellationToken: any + vsc.window.withProgress.mockImplementationOnce((_options: any, task: any) => { + const progress = { report: vi.fn() } + const tokenListeners: Array<() => void> = [] + capturedCancellationToken = { + isCancellationRequested: false, + onCancellationRequested: vi.fn((cb: () => void) => { + tokenListeners.push(cb) + return { dispose: vi.fn() } + }), + _fire: () => tokenListeners.forEach((cb) => cb()), + } + return task(progress, capturedCancellationToken) + }) + + const flowPromise = (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + // _initiateOAuthFlow awaits getOAuthData() before calling withProgress. + // Flush that microtask so withProgress runs and capturedCancellationToken is set. + await Promise.resolve() + capturedCancellationToken._fire() + + await flowPromise + + expect(mockConnection.server.status).toBe("disconnected") + expect((mcpHub as any).appendErrorMessage).toHaveBeenCalled() + expect(mockAuthProvider.close).toHaveBeenCalled() + }) + + it("should disconnect and flag error when OAuth flow times out", async () => { + vi.useFakeTimers() + vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {})) + + const flowPromise = (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + await vi.advanceTimersByTimeAsync(OAUTH_FLOW_TIMEOUT_MS) + await flowPromise + + expect(mockConnection.server.status).toBe("disconnected") + expect((mcpHub as any).appendErrorMessage).toHaveBeenCalled() + expect(mockAuthProvider.close).toHaveBeenCalled() + }) + + 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) + vsc.window.showInformationMessage.mockResolvedValueOnce(t("mcp:oauth.flow.authenticateButton") as any) + mockSecretStorage.getOAuthData.mockResolvedValue({ + expires_at: Date.now() + 10 * 60 * 1000, + }) + + const completeOAuthSpy = vi.spyOn(mcpHub as any, "_completeOAuthFlow") + + await (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ) + + expect(completeOAuthSpy).not.toHaveBeenCalled() + expect(mockAuthProvider.close).toHaveBeenCalled() + expect((mcpHub as any).deleteConnection).toHaveBeenCalledWith(serverName, source) + expect((mcpHub as any).connectToServer).toHaveBeenCalled() + }) + + it("should resolve cleanly even when _completeOAuthFlow throws", async () => { + vsc.window.showInformationMessage.mockResolvedValueOnce(t("mcp:oauth.flow.authenticateButton") as any) + vi.spyOn(mcpHub as any, "_completeOAuthFlow").mockRejectedValue(new Error("network failure")) + + await expect( + (mcpHub as any)._initiateOAuthFlow( + serverName, + source, + config, + mockAuthProvider, + mockTransport, + mockConnection, + ), + ).resolves.toBeUndefined() + }) + }) }) diff --git a/src/services/mcp/constants.ts b/src/services/mcp/constants.ts index e49ed4e027..8107ee4a29 100644 --- a/src/services/mcp/constants.ts +++ b/src/services/mcp/constants.ts @@ -1 +1,2 @@ export const TOKEN_EXPIRY_BUFFER_MS = 5 * 60 * 1000 // 5 minutes +export const OAUTH_FLOW_TIMEOUT_MS = 5 * 60 * 1000 // 5 minutes diff --git a/src/services/mcp/utils/callbackServer.ts b/src/services/mcp/utils/callbackServer.ts index ac7a421058..e447222b0d 100644 --- a/src/services/mcp/utils/callbackServer.ts +++ b/src/services/mcp/utils/callbackServer.ts @@ -1,5 +1,6 @@ import * as http from "http" import { t } from "../../../i18n" +import { OAUTH_FLOW_TIMEOUT_MS } from "../constants" export interface CallbackResult { code?: string @@ -49,16 +50,13 @@ export function startCallbackServer( const resultPromise = new Promise((resolveResult, rejectResult) => { let resolved = false - const timeout = setTimeout( - () => { - if (!resolved) { - resolved = true - rejectResult(new Error("Callback timeout")) - server.close() - } - }, - 5 * 60 * 1000, - ) // 5 minutes + const timeout = setTimeout(() => { + if (!resolved) { + resolved = true + rejectResult(new Error("Callback timeout")) + server.close() + } + }, OAUTH_FLOW_TIMEOUT_MS) server.on("request", (req: any, res: any) => { if (resolved) return