mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-05 02:41:26 +00:00
refactor(McpHub): making prompt long running or dismissable
This commit is contained in:
parent
f2a3cfba07
commit
8eb8be93d3
5 changed files with 435 additions and 116 deletions
|
|
@ -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."
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 =
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue