refactor(McpHub): cleanup crosswindow auths

This commit is contained in:
Elliott de Launay 2026-04-24 09:47:43 -04:00
parent ba5345c084
commit a127704207
No known key found for this signature in database
GPG key ID: BB899BED766D1806
3 changed files with 100 additions and 21 deletions

View file

@ -171,8 +171,11 @@ export class McpHub {
private reauthPromises: Map<string, Promise<void>> = new Map()
private _oauthWatchers: Map<string, { unsubscribe: () => void; abortHandle: NodeJS.Timeout }> = new Map()
constructor(provider: ClineProvider) {
constructor(provider: ClineProvider, secretStorage?: SecretStorageService) {
this.providerRef = new WeakRef(provider)
if (secretStorage) {
this.secretStorage = secretStorage
}
this.watchMcpSettingsFile()
this.watchProjectMcpFile().catch(console.error)
this.setupWorkspaceFoldersWatcher()
@ -190,16 +193,12 @@ export class McpHub {
await this.initializationPromise
}
public setSecretStorage(secretStorage: SecretStorageService): void {
this.secretStorage = secretStorage
}
/**
* Registers a client (e.g., ClineProvider) using this hub.
* Increments the reference count.
*/
public registerClient(): void {
this.refCount++
// console.log(`McpHub: Client registered. Ref count: ${this.refCount}`)
}
/**
@ -209,8 +208,6 @@ export class McpHub {
public async unregisterClient(): Promise<void> {
this.refCount--
// console.log(`McpHub: Client unregistered. Ref count: ${this.refCount}`)
if (this.refCount <= 0) {
console.log("McpHub: Last client unregistered. Disposing hub.")
await this.dispose()
@ -1054,9 +1051,31 @@ export class McpHub {
return
}
// Check if another window already saved valid tokens
// Register the cross-window token watcher BEFORE the initial token read so we
// don't miss a write that lands in the gap between the read and the subscription.
// After the read we immediately unsubscribe; the watcher inside _runOAuthFlow
// will set up its own long-lived subscription for the duration of the flow.
// Note: onDidChange fires on both writes AND deletes — the re-read below is
// authoritative; we only reconnect if it returns a valid (non-expired) token.
let tokenChangedDuringRead = false
const unsubscribeCrossWindow = this.secretStorage.onDidChange(serverUrl, () => {
tokenChangedDuringRead = true
})
// Check if another window already saved valid tokens.
// Re-read after registering the watcher in case a write landed in the gap.
const existing = await this.secretStorage.getOAuthData(serverUrl)
if (existing && Date.now() < existing.expires_at - TOKEN_EXPIRY_BUFFER_MS) {
if (this.isDisposed) {
unsubscribeCrossWindow()
return
}
// If the first read missed but the watcher fired during it, re-read — a write
// may have landed between subscription and read completion.
const tokenToUse =
existing ?? (tokenChangedDuringRead ? await this.secretStorage.getOAuthData(serverUrl) : undefined)
unsubscribeCrossWindow()
if (!this.isDisposed && tokenToUse && Date.now() < tokenToUse.expires_at - TOKEN_EXPIRY_BUFFER_MS) {
await authProvider.close()
await this.deleteConnection(name, source)
await this.connectToServer(name, config, source)
@ -1185,8 +1204,8 @@ export class McpHub {
// Register in _oauthWatchers so deleteConnection() and dispose() can clean up
this._oauthWatchers.set(watcherKey, { unsubscribe, abortHandle: timeoutHandle })
// Non-modal toast — fires and forgets. When it auto-dismisses, update the
// persistent progress bar to tell the user how to re-trigger the flow.
// Non-modal toast — fires once. If dismissed without clicking Authenticate,
// the flow stays alive via the persistent progress bar (Cancel to abort).
const authenticateLabel = t("mcp:oauth.flow.authenticateButton")
void (async () => {
const choice = await vscode.window.showInformationMessage(
@ -1197,8 +1216,6 @@ export class McpHub {
if (disposed) return
if (choice !== authenticateLabel) {
// Toast was dismissed (timed out or closed) without clicking Authenticate.
// Update the progress bar so the user knows how to re-trigger.
progress.report({ message: t("mcp:oauth.flow.dismissedHint") })
return
}

View file

@ -37,10 +37,8 @@ export class McpServerManager {
try {
// Double-check instance in case it was created while we were waiting
if (!this.instance) {
const hub = new McpHub(provider)
// Set the secret storage service for OAuth operations
const secretStorage = new SecretStorageService(context)
hub.setSecretStorage(secretStorage)
const hub = new McpHub(provider, secretStorage)
// Wait for all MCP servers to finish connecting (or timing out)
await hub.waitUntilReady()
this.instance = hub

View file

@ -2463,8 +2463,6 @@ describe("McpHub", () => {
})
it("should update progress bar hint when toast is dismissed without clicking Authenticate", async () => {
// Toast dismissed (undefined) without clicking Authenticate — flow stays alive
// via the progress bar. The progress bar message should update to the dismissedHint.
let capturedProgress: any
vsc.window.withProgress.mockImplementationOnce((_options: any, task: any) => {
capturedProgress = { report: vi.fn() }
@ -2476,7 +2474,6 @@ describe("McpHub", () => {
return task(capturedProgress, cancellationToken)
})
// Toast dismissed (no button clicked), then flow stays open forever (cancel via timeout)
vsc.window.showInformationMessage.mockResolvedValueOnce(undefined as any)
vi.useFakeTimers()
@ -2546,8 +2543,6 @@ describe("McpHub", () => {
})
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) => {
@ -2738,6 +2733,75 @@ describe("McpHub", () => {
expect((mcpHub as any).connectToServer).not.toHaveBeenCalled()
})
it("should not reconnect when cross-window watcher fires but token is missing or expired", async () => {
vi.useFakeTimers()
vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {}))
// onDidChange fires immediately (simulates a storage event during the read gap)
// but getOAuthData returns undefined — no valid token written.
mockSecretStorage.onDidChange.mockImplementation((_key: string, cb: () => void) => {
cb() // fires synchronously — sets crossWindowTokenWritten flag
return vi.fn()
})
mockSecretStorage.getOAuthData.mockResolvedValue(undefined)
const flowPromise = (mcpHub as any)._initiateOAuthFlow(
serverName,
source,
config,
mockAuthProvider,
mockTransport,
mockConnection,
)
// Advance past the timeout so the flow can settle
await vi.advanceTimersByTimeAsync(OAUTH_FLOW_TIMEOUT_MS)
await flowPromise
// Should NOT have tried to reconnect — the watcher fired but no valid token exists
expect((mcpHub as any).connectToServer).not.toHaveBeenCalledWith(serverName, expect.anything(), source)
})
it("should reconnect when cross-window token is written while getOAuthData is in-flight", async () => {
// This covers the race: onDidChange fires BEFORE getOAuthData resolves.
// The flag causes a second read which finds the now-valid token.
let resolveGetOAuthData!: (value: any) => void
mockSecretStorage.getOAuthData
// First call (during the race window) — delayed, returns undefined
.mockImplementationOnce(
() =>
new Promise((r) => {
resolveGetOAuthData = r
}),
)
// Second call (after watcher fires) — valid token available
.mockResolvedValueOnce({ expires_at: Date.now() + 10 * 60 * 1000 })
// Watcher fires synchronously before getOAuthData resolves
mockSecretStorage.onDidChange.mockImplementation((_key: string, cb: () => void) => {
cb()
return vi.fn()
})
const flowPromise = (mcpHub as any)._initiateOAuthFlow(
serverName,
source,
config,
mockAuthProvider,
mockTransport,
mockConnection,
)
// Now let the first getOAuthData resolve with undefined
resolveGetOAuthData(undefined)
await flowPromise
expect(mockAuthProvider.close).toHaveBeenCalled()
expect((mcpHub as any).deleteConnection).toHaveBeenCalledWith(serverName, source)
expect((mcpHub as any).connectToServer).toHaveBeenCalled()
expect(vsc.window.withProgress).not.toHaveBeenCalled()
})
it("should cancel the previous watcher when called again for the same server", async () => {
vi.useFakeTimers()
vsc.window.showInformationMessage.mockImplementation(() => new Promise(() => {}))