feat(McpHub): handling concurrent window auth

This commit is contained in:
Elliott de Launay 2026-04-23 22:13:50 -04:00
parent 8eb8be93d3
commit f10d2541f6
No known key found for this signature in database
GPG key ID: BB899BED766D1806
5 changed files with 139 additions and 12 deletions

View file

@ -1111,20 +1111,32 @@ export class McpHub {
return new Promise<void>((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()

View file

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

View file

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

View file

@ -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(`
<!DOCTYPE html>
<html>
@ -90,11 +93,16 @@ export function startCallbackServer(
</html>
`)
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(`
<!DOCTYPE html>
<html>
@ -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 })}\`;

View file

@ -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<OAuthDiscoveryResult | null> {
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<never>((_, 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<O
const response = await fetch(discoveryUrl, {
headers: { Accept: "application/json" },
signal: AbortSignal.timeout(DISCOVERY_TIMEOUT_MS),
})
if (!response.ok) return null
const authServerMeta = (await response.json()) as Record<string, any>