mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-03 02:22:24 +00:00
feat(McpHub): handling concurrent window auth
This commit is contained in:
parent
8eb8be93d3
commit
f10d2541f6
5 changed files with 139 additions and 12 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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 })}\`;
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue