refactor(McpHub): making prompt long running or dismissable

This commit is contained in:
Elliott de Launay 2026-04-02 23:45:07 -04:00
parent f2a3cfba07
commit 8eb8be93d3
No known key found for this signature in database
GPG key ID: BB899BED766D1806
5 changed files with 435 additions and 116 deletions

View file

@ -36,6 +36,15 @@
"tab_closing_in": "This tab will attempt to close in <span id=\"count\">{{count}}</span>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."
}
}
}

View file

@ -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<typeof ServerConfigSchema>,
serverUrl: string,
authProvider: McpOAuthClientProvider,
transport: StreamableHTTPClientTransport,
connection: ConnectedMcpConnection,
progress: vscode.Progress<{ message?: string; increment?: number }>,
cancellationToken: vscode.CancellationToken,
): Promise<void> {
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<void>((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<typeof ServerConfigSchema>,
): 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 =

View file

@ -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()
})
})
})

View file

@ -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

View file

@ -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<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
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