diff --git a/shared/glean/mcp/src/auth-provider.ts b/shared/glean/mcp/src/auth-provider.ts index 10ba626..4143cb4 100644 --- a/shared/glean/mcp/src/auth-provider.ts +++ b/shared/glean/mcp/src/auth-provider.ts @@ -6,9 +6,14 @@ import type { } from "@modelcontextprotocol/client"; import { execFile, spawn } from "node:child_process"; import { randomUUID } from "node:crypto"; +import { setTimeout as sleep } from "node:timers/promises"; import { platform } from "node:os"; import { getCallbackUrl, setExpectedState } from "./auth-callback-server.js"; -import { clearCredentials, loadCredentials, saveCredentials } from "./token-store.js"; +import { + clearCredentials, + loadCredentials, + saveCredentials, +} from "./token-store.js"; export type InvalidationScope = | "all" @@ -17,6 +22,10 @@ export type InvalidationScope = | "verifier" | "discovery"; +// Grace window for a sibling's in-flight refresh to land on disk. +const ROTATION_GRACE_MS = 2000; +const ROTATION_POLL_MS = 500; + /** * Open `url` in the user's default browser. Used for the self-open sign-in * path when the client does not support URL-mode elicitation (where the client @@ -53,7 +62,6 @@ export class GleanOAuthClientProvider implements OAuthClientProvider { // explicitly invalidating. Used to detect when a previous auth URL didn't // complete — likely because the server rejected the (stale) client_id. private _authUrlPending = false; - authorizationUrl: string | undefined; /** @@ -72,6 +80,33 @@ export class GleanOAuthClientProvider implements OAuthClientProvider { } } + // Re-read the shared store on every token access so a sibling's rotated + // grant is used instead of a stale in-memory copy. + private syncTokensFromDisk(): void { + const stored = loadCredentials(); + if (!stored) return; + if (stored.tokens) { + this._tokens = stored.tokens; + } + if (stored.clientInfo) { + this._clientInfo = stored.clientInfo; + } + } + + // Wait for a sibling's refresh to land on disk. Returns true once a + // different access token is available for adoption/retry. + async waitForSiblingRefresh( + previousAccessToken: string | undefined, + ): Promise { + const deadline = Date.now() + ROTATION_GRACE_MS; + for (;;) { + const current = this.tokens()?.access_token; + if (current && current !== previousAccessToken) return true; + if (Date.now() >= deadline) return false; + await sleep(ROTATION_POLL_MS); + } + } + get redirectUrl(): string { return getCallbackUrl(); } @@ -93,6 +128,7 @@ export class GleanOAuthClientProvider implements OAuthClientProvider { } tokens(): StoredOAuthTokens | undefined { + this.syncTokensFromDisk(); return this._tokens; } @@ -118,10 +154,23 @@ export class GleanOAuthClientProvider implements OAuthClientProvider { this._clientInfo = undefined; saveCredentials(this._tokens, undefined); break; - case "tokens": + case "tokens": { + // SDK auth() invalidates tokens after invalid_grant, which can mean a + // sibling already rotated our refresh token. Client errors invalidate + // "client" before "tokens" instead: without a retained client, clear + // immediately. A newer token must not cancel a client or full reset. + const previousAccessToken = this._tokens?.access_token; + if ( + this._clientInfo && + this._tokens?.refresh_token && + (await this.waitForSiblingRefresh(previousAccessToken)) + ) { + return; + } this._tokens = undefined; saveCredentials(undefined, this._clientInfo); break; + } case "verifier": this._codeVerifier = ""; break; diff --git a/shared/glean/mcp/src/remote-client.ts b/shared/glean/mcp/src/remote-client.ts index db991d8..cb19530 100644 --- a/shared/glean/mcp/src/remote-client.ts +++ b/shared/glean/mcp/src/remote-client.ts @@ -2,6 +2,8 @@ import { Client, StreamableHTTPClientTransport, UnauthorizedError, + OAuthError, + OAuthErrorCode, type ElicitRequest, type ElicitResult, } from "@modelcontextprotocol/client"; @@ -177,6 +179,7 @@ export async function createRemoteClient( serverUrl: string, opts: RemoteClientOptions, chatSessionId?: string, + authRetry = false, ): Promise { const authProvider = opts.authProvider; @@ -245,14 +248,48 @@ export async function createRemoteClient( }); } + // Snapshot to detect a sibling's refresh between connect and failure. + const accessTokenAtConnect = authProvider?.tokens()?.access_token; + const transport = buildTransport(serverUrl, opts, chatSessionId); try { await withConnectLock(() => client.connect(transport)); } catch (error) { - if (error instanceof UnauthorizedError && authProvider?.authorizationUrl) { - pendingTransport = transport; - throw new AuthRequiredError(authProvider.authorizationUrl); + if (!authProvider) { + throw error; + } + + if (error instanceof UnauthorizedError) { + const refreshedAccessToken = authProvider.tokens()?.access_token; + if ( + !authRetry && + refreshedAccessToken && + refreshedAccessToken !== accessTokenAtConnect + ) { + console.error( + "[auth] Auth failed but a newer token is on disk " + + "(sibling refresh) — retrying once", + ); + return createRemoteClient(serverUrl, opts, chatSessionId, true); + } + if (authProvider.authorizationUrl) { + pendingTransport = transport; + throw new AuthRequiredError(authProvider.authorizationUrl); + } + } + // Concurrent-refresh losers are reported with structured OAuth errors + // (typically invalid_request); retry once if a sibling's grant lands in the + // grace window. + if ( + !authRetry && + isRefreshOAuthError(error) && + (await authProvider.waitForSiblingRefresh(accessTokenAtConnect)) + ) { + console.error( + "[auth] Refresh failed but a sibling refreshed — retrying with its token", + ); + return createRemoteClient(serverUrl, opts, chatSessionId, true); } throw error; } @@ -260,6 +297,16 @@ export async function createRemoteClient( return client; } +// Restrict recovery to OAuth errors that can indicate a refresh race. The SDK +// preserves the response's machine-readable error code. +function isRefreshOAuthError(error: unknown): boolean { + return ( + error instanceof OAuthError && + (error.code === OAuthErrorCode.InvalidRequest || + error.code === OAuthErrorCode.InvalidGrant) + ); +} + export async function callRemoteTool( client: Client, name: string, diff --git a/shared/glean/mcp/src/token-store.ts b/shared/glean/mcp/src/token-store.ts index 1331dc2..804ee3f 100644 --- a/shared/glean/mcp/src/token-store.ts +++ b/shared/glean/mcp/src/token-store.ts @@ -5,6 +5,7 @@ import type { import fs from "node:fs"; import path from "node:path"; import { serverDataDir } from "./data-dir.js"; +import { writeFileAtomicSync } from "./atomic-write.js"; const CREDENTIALS_FILENAME = "mcp-credentials.json"; const DIR_MODE = 0o700; @@ -38,11 +39,7 @@ export function saveCredentials( fs.mkdirSync(dir, { recursive: true, mode: DIR_MODE }); fs.chmodSync(dir, DIR_MODE); const data: StoredCredentials = { tokens, clientInfo }; - fs.writeFileSync(filePath, JSON.stringify(data, null, 2), { - encoding: "utf-8", - mode: FILE_MODE, - }); - fs.chmodSync(filePath, FILE_MODE); + writeFileAtomicSync(filePath, JSON.stringify(data, null, 2), FILE_MODE); } catch (err) { const msg = err instanceof Error ? err.message : String(err); console.error(`[auth] Failed to persist credentials: ${msg}`); diff --git a/shared/glean/mcp/tests/auth-provider-sdk.test.ts b/shared/glean/mcp/tests/auth-provider-sdk.test.ts new file mode 100644 index 0000000..f0d9d5e --- /dev/null +++ b/shared/glean/mcp/tests/auth-provider-sdk.test.ts @@ -0,0 +1,342 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { + auth, + type StoredOAuthClientInformation, + type StoredOAuthTokens, +} from "@modelcontextprotocol/client"; +import { createHash } from "node:crypto"; +import fs from "node:fs"; +import os from "node:os"; +import path from "node:path"; + +// Keep the provider's real polling loop, but let Vitest control its sleeps. +vi.mock("node:timers/promises", () => ({ + setTimeout: (delay: number) => + new Promise((resolve) => globalThis.setTimeout(resolve, delay)), +})); + +vi.mock("../src/auth-callback-server.js", () => ({ + getCallbackUrl: () => "http://127.0.0.1:29107/glean-cli-callback", + setExpectedState: vi.fn(), +})); + +import { GleanOAuthClientProvider } from "../src/auth-provider.js"; +import { setExpectedState } from "../src/auth-callback-server.js"; +import { loadCredentials, saveCredentials } from "../src/token-store.js"; + +const serverUrl = "https://mcp.example.test/mcp"; +const issuer = "https://auth.example.test"; +const resourceMetadataUrl = new URL( + "https://mcp.example.test/.well-known/oauth-protected-resource/mcp", +); +const authMetadataUrl = `${issuer}/.well-known/oauth-authorization-server`; +const authorizationEndpoint = `${issuer}/authorize`; +const tokenEndpoint = `${issuer}/token`; +const registrationEndpoint = `${issuer}/register`; +const callbackUrl = "http://127.0.0.1:29107/glean-cli-callback"; + +// Stamp both clients and tokens up front. SDK v2 otherwise writes an issuer +// migration before refreshing, which can overwrite the sibling's disk state. +const originalClient: StoredOAuthClientInformation = { + client_id: "client-0", + issuer, +}; +const originalTokens: StoredOAuthTokens = { + access_token: "T0", + refresh_token: "R0", + token_type: "Bearer", + issuer, +}; +const siblingTokens: StoredOAuthTokens = { + access_token: "T1", + refresh_token: "R1", + token_type: "Bearer", + issuer, +}; +const refreshedTokenResponse = { + access_token: "T2", + refresh_token: "R2", + token_type: "Bearer", +}; +const refreshedTokens: StoredOAuthTokens = { ...refreshedTokenResponse, issuer }; +const registeredClientResponse = { + client_id: "client-1", + redirect_uris: [callbackUrl], + token_endpoint_auth_method: "none", +}; +const registeredClient = { ...registeredClientResponse, issuer }; + +const discoveryRequests = [ + `GET ${resourceMetadataUrl.href}`, + `GET ${authMetadataUrl}`, +]; + +type RefreshError = "invalid_grant" | "invalid_client" | "unauthorized_client"; + +// Only HTTP is faked. Real SDK discovery, response parsing, error dispatch, +// registration, PKCE, and auth retries run against the real provider/store. +function makeOAuthServer(error: RefreshError) { + const requests: string[] = []; + const refreshRequests: Record[] = []; + const registrations: unknown[] = []; + const credentialsAtRegistration: ReturnType[] = []; + const fetchFn = vi.fn(async (url: string | URL, init?: RequestInit) => { + const request = `${init?.method ?? "GET"} ${String(url)}`; + requests.push(request); + switch (request) { + case `GET ${resourceMetadataUrl.href}`: + return Response.json({ + resource: serverUrl, + authorization_servers: [issuer], + scopes_supported: ["mcp"], + }); + case `GET ${authMetadataUrl}`: + return Response.json({ + issuer, + authorization_endpoint: authorizationEndpoint, + token_endpoint: tokenEndpoint, + registration_endpoint: registrationEndpoint, + response_types_supported: ["code"], + grant_types_supported: ["authorization_code", "refresh_token"], + token_endpoint_auth_methods_supported: ["none"], + code_challenge_methods_supported: ["S256"], + }); + case `POST ${tokenEndpoint}`: { + const params = new URLSearchParams(String(init?.body)); + refreshRequests.push(Object.fromEntries(params)); + if (params.get("refresh_token") === "R0") { + return Response.json( + { error, error_description: "Synthetic refresh rejection" }, + { status: error === "invalid_client" ? 401 : 400 }, + ); + } + if (params.get("refresh_token") === "R1") { + return Response.json(refreshedTokenResponse); + } + throw new Error("Unexpected refresh token in test HTTP request"); + } + case `POST ${registrationEndpoint}`: + registrations.push(JSON.parse(String(init?.body))); + credentialsAtRegistration.push(loadCredentials()); + return Response.json(registeredClientResponse, { status: 201 }); + default: + throw new Error(`Unexpected test HTTP request: ${request}`); + } + }); + return { + fetchFn, + requests, + refreshRequests, + registrations, + credentialsAtRegistration, + }; +} + +function refreshRequest(refreshToken: string) { + return { + grant_type: "refresh_token", + refresh_token: refreshToken, + client_id: originalClient.client_id, + resource: serverUrl, + }; +} + +function observeProvider(provider: GleanOAuthClientProvider) { + provider.onTokensChanged = vi.fn(); + return { + invalidate: vi.spyOn(provider, "invalidateCredentials"), + wait: vi.spyOn(provider, "waitForSiblingRefresh"), + saveClient: vi.spyOn(provider, "saveClientInformation"), + saveTokens: vi.spyOn(provider, "saveTokens"), + redirect: vi.spyOn(provider, "redirectToAuthorization"), + }; +} + +function expectAuthorizationRedirect( + provider: GleanOAuthClientProvider, + clientId: string, +) { + expect(provider.authorizationUrl).toBeDefined(); + const url = new URL(provider.authorizationUrl!); + expect(`${url.origin}${url.pathname}`).toBe(authorizationEndpoint); + expect(Object.fromEntries(url.searchParams)).toMatchObject({ + response_type: "code", + client_id: clientId, + redirect_uri: callbackUrl, + resource: serverUrl, + scope: "mcp", + code_challenge_method: "S256", + code_challenge: createHash("sha256") + .update(provider.codeVerifier()) + .digest("base64url"), + }); + expect(provider.codeVerifier()).toMatch(/^[A-Za-z0-9._~-]{43,128}$/); + expect(url.searchParams.get("state")).toBeTruthy(); + expect(setExpectedState).toHaveBeenCalledExactlyOnceWith( + url.searchParams.get("state"), + ); +} + +describe("GleanOAuthClientProvider with real SDK auth()", () => { + let dataDir: string; + let provider: GleanOAuthClientProvider; + + beforeEach(() => { + dataDir = fs.mkdtempSync(path.join(os.tmpdir(), "auth-provider-sdk-test-")); + vi.stubEnv("PLUGIN_DATA_DIR", dataDir); + // Fail closed if any SDK path forgets the explicit fake HTTP boundary. + vi.stubGlobal("fetch", vi.fn(() => { + throw new Error("Network access is forbidden in OAuth SDK tests"); + })); + vi.clearAllMocks(); + vi.useFakeTimers(); + vi.setSystemTime(new Date("2026-01-01T00:00:00Z")); + saveCredentials(originalTokens, originalClient); + provider = new GleanOAuthClientProvider(); + }); + + afterEach(() => { + vi.clearAllTimers(); + vi.useRealTimers(); + vi.restoreAllMocks(); + vi.unstubAllGlobals(); + vi.unstubAllEnvs(); + fs.rmSync(dataDir, { recursive: true, force: true }); + }); + + it("retries invalid_grant with a sibling's R1 and persists T2/R2 without resetting the client", async () => { + const sibling = new GleanOAuthClientProvider(); + const observed = observeProvider(provider); + const server = makeOAuthServer("invalid_grant"); + const result = auth(provider, { serverUrl, resourceMetadataUrl, fetchFn: server.fetchFn }); + + await vi.advanceTimersByTimeAsync(0); + expect(server.refreshRequests).toEqual([refreshRequest("R0")]); + expect(observed.invalidate.mock.calls).toEqual([["tokens"]]); + expect(observed.wait).toHaveBeenCalledExactlyOnceWith("T0"); + + // The other process finishes after invalid_grant, while our real provider + // is asleep. It must adopt the atomic disk write at the next 500 ms poll. + await vi.advanceTimersByTimeAsync(250); + sibling.saveTokens(siblingTokens); + expect(loadCredentials()).toEqual({ tokens: siblingTokens, clientInfo: originalClient }); + await vi.advanceTimersByTimeAsync(249); + expect(server.refreshRequests).toHaveLength(1); + expect(observed.saveTokens).not.toHaveBeenCalled(); + await vi.advanceTimersByTimeAsync(1); + + await expect(result).resolves.toBe("AUTHORIZED"); + expect(server.refreshRequests).toEqual([refreshRequest("R0"), refreshRequest("R1")]); + expect(server.requests).toEqual([ + ...discoveryRequests, `POST ${tokenEndpoint}`, + ...discoveryRequests, `POST ${tokenEndpoint}`, + ]); + expect(observed.invalidate.mock.calls).toEqual([["tokens"]]); + expect(observed.saveTokens).toHaveBeenCalledExactlyOnceWith(refreshedTokens, { issuer }); + expect(provider.onTokensChanged).toHaveBeenCalledExactlyOnceWith(refreshedTokens); + expect(provider.tokens()).toEqual(refreshedTokens); + expect(provider.clientInformation()).toEqual(originalClient); + expect(loadCredentials()).toEqual({ tokens: refreshedTokens, clientInfo: originalClient }); + expect(observed.saveClient).not.toHaveBeenCalled(); + expect(server.registrations).toEqual([]); + expect(observed.redirect).not.toHaveBeenCalled(); + expect(provider.authorizationUrl).toBeUndefined(); + expect(vi.getTimerCount()).toBe(0); + }); + + it("waits the full 2 seconds on genuine invalid_grant, then clears tokens and redirects with the same client", async () => { + const observed = observeProvider(provider); + const server = makeOAuthServer("invalid_grant"); + const startedAt = Date.now(); + const result = auth(provider, { serverUrl, resourceMetadataUrl, fetchFn: server.fetchFn }); + + await vi.advanceTimersByTimeAsync(0); + expect(observed.invalidate.mock.calls).toEqual([["tokens"]]); + expect(observed.wait).toHaveBeenCalledExactlyOnceWith("T0"); + for (const elapsed of [500, 500, 500, 499]) { + await vi.advanceTimersByTimeAsync(elapsed); + expect(loadCredentials()).toEqual({ tokens: originalTokens, clientInfo: originalClient }); + expect(server.refreshRequests).toEqual([refreshRequest("R0")]); + expect(observed.redirect).not.toHaveBeenCalled(); + expect(provider.onTokensChanged).not.toHaveBeenCalled(); + } + expect(Date.now() - startedAt).toBe(1999); + await vi.advanceTimersByTimeAsync(1); + + await expect(result).resolves.toBe("REDIRECT"); + expect(Date.now() - startedAt).toBe(2000); + expect(server.requests).toEqual([ + ...discoveryRequests, `POST ${tokenEndpoint}`, ...discoveryRequests, + ]); + expect(observed.invalidate.mock.calls).toEqual([["tokens"]]); + expect(provider.tokens()).toBeUndefined(); + expect(provider.clientInformation()).toEqual(originalClient); + expect(loadCredentials()).toEqual({ clientInfo: originalClient }); + expect(provider.onTokensChanged).toHaveBeenCalledExactlyOnceWith(undefined); + expect(observed.saveClient).not.toHaveBeenCalled(); + expect(observed.saveTokens).not.toHaveBeenCalled(); + expect(server.registrations).toEqual([]); + expect(observed.redirect).toHaveBeenCalledTimes(1); + expectAuthorizationRedirect(provider, originalClient.client_id); + expect(vi.getTimerCount()).toBe(0); + }); + + describe.each(["invalid_client", "unauthorized_client"] as const)("%s", (error) => { + it.each([false, true])( + "resets client then tokens without grace and re-registers (sibling write after client reset: %s)", + async (writeSibling) => { + const sibling = new GleanOAuthClientProvider(); + const observed = observeProvider(provider); + const server = makeOAuthServer(error); + const startedAt = Date.now(); + const result = auth(provider, { serverUrl, resourceMetadataUrl, fetchFn: server.fetchFn }); + + if (writeSibling) { + // Observe the SDK's await between its two real invalidations. A + // bounded microtask loop permits the sibling write at that boundary + // without replacing authInternal() or invalidateCredentials(). + for (let hop = 0; hop < 100 && provider.clientInformation(); hop++) { + await Promise.resolve(); + } + expect(observed.invalidate.mock.calls).toEqual([["client"]]); + expect(provider.clientInformation()).toBeUndefined(); + expect(loadCredentials()).toEqual({ tokens: originalTokens }); + sibling.saveTokens(siblingTokens); + expect(loadCredentials()).toEqual({ tokens: siblingTokens, clientInfo: originalClient }); + } + + await vi.advanceTimersByTimeAsync(0); + expect(observed.invalidate.mock.calls).toEqual([["client"], ["tokens"]]); + expect(observed.wait).not.toHaveBeenCalled(); + expect(vi.getTimerCount()).toBe(0); + await expect(result).resolves.toBe("REDIRECT"); + + expect(Date.now() - startedAt).toBe(0); + expect(server.refreshRequests).toEqual([refreshRequest("R0")]); + expect(server.requests).toEqual([ + ...discoveryRequests, `POST ${tokenEndpoint}`, + ...discoveryRequests, `POST ${registrationEndpoint}`, + ]); + expect(server.credentialsAtRegistration).toEqual([{}]); + expect(server.registrations).toEqual([{ + redirect_uris: [callbackUrl], + client_name: "Glean Claude Code Plugin", + application_type: "native", + grant_types: ["authorization_code", "refresh_token"], + scope: "mcp", + }]); + expect(observed.saveClient).toHaveBeenCalledExactlyOnceWith(registeredClient, { issuer }); + expect(observed.invalidate.mock.invocationCallOrder[1]).toBeLessThan( + observed.saveClient.mock.invocationCallOrder[0], + ); + expect(observed.saveTokens).not.toHaveBeenCalled(); + expect(provider.tokens()).toBeUndefined(); + expect(provider.clientInformation()).toEqual(registeredClient); + expect(loadCredentials()).toEqual({ clientInfo: registeredClient }); + expect(provider.onTokensChanged).toHaveBeenCalledExactlyOnceWith(undefined); + expect(observed.redirect).toHaveBeenCalledTimes(1); + expectAuthorizationRedirect(provider, registeredClient.client_id); + }, + ); + }); +}); diff --git a/shared/glean/mcp/tests/auth-provider.test.ts b/shared/glean/mcp/tests/auth-provider.test.ts index ba181f0..7817958 100644 --- a/shared/glean/mcp/tests/auth-provider.test.ts +++ b/shared/glean/mcp/tests/auth-provider.test.ts @@ -9,6 +9,10 @@ vi.mock("node:os", async () => { return { ...actual, homedir: () => tmpDir }; }); +vi.mock("node:timers/promises", () => ({ + setTimeout: (delay: number) => new Promise((resolve) => setTimeout(resolve, delay)), +})); + vi.mock("node:child_process", () => ({ exec: vi.fn(), execFile: vi.fn(), @@ -27,12 +31,15 @@ describe("GleanOAuthClientProvider", () => { const gleanDir = path.join(tmpDir, ".glean"); beforeEach(() => { - delete process.env.PLUGIN_DATA_DIR; + vi.stubEnv("PLUGIN_DATA_DIR", gleanDir); fs.rmSync(gleanDir, { recursive: true, force: true }); vi.clearAllMocks(); }); afterEach(() => { + vi.useRealTimers(); + vi.restoreAllMocks(); + vi.unstubAllEnvs(); fs.rmSync(gleanDir, { recursive: true, force: true }); }); @@ -74,6 +81,161 @@ describe("GleanOAuthClientProvider", () => { expect(raw.tokens.access_token).toBe("new_tok"); }); + // --- Cross-process sync: tokens() must pick up a sibling's rewrite. --- + + const credFile = path.join(gleanDir, "mcp-credentials.json"); + + function writeCredFile(tokens: unknown, clientInfo?: unknown): void { + fs.mkdirSync(gleanDir, { recursive: true }); + fs.writeFileSync(credFile, JSON.stringify({ tokens, clientInfo })); + } + + it("tokens() adopts a token written by another process", () => { + fs.mkdirSync(gleanDir, { recursive: true }); + fs.writeFileSync( + credFile, + JSON.stringify({ + tokens: { access_token: "T0", refresh_token: "R0" }, + clientInfo: { client_id: "cid" }, + }), + ); + const provider = new GleanOAuthClientProvider(); + expect(provider.tokens()?.access_token).toBe("T0"); + const originalMtime = fs.statSync(credFile).mtime; + + // Sibling refreshes: new access + rotated refresh token on disk. + writeCredFile( + { access_token: "T1", refresh_token: "R1" }, + { client_id: "cid" }, + ); + // The provider must not rely on mtime to observe this rewrite. + fs.utimesSync(credFile, originalMtime, originalMtime); + + expect(provider.tokens()?.access_token).toBe("T1"); + expect(provider.tokens()?.refresh_token).toBe("R1"); + }); + + it("tokens() keeps the in-memory token when the file is deleted", () => { + fs.mkdirSync(gleanDir, { recursive: true }); + fs.writeFileSync( + credFile, + JSON.stringify({ tokens: { access_token: "T0" }, clientInfo: {} }), + ); + const provider = new GleanOAuthClientProvider(); + expect(provider.tokens()?.access_token).toBe("T0"); + + // Transient disappearance / another process mid-write — don't self-evict. + fs.rmSync(credFile, { force: true }); + expect(provider.tokens()?.access_token).toBe("T0"); + }); + + it("tokens() does not adopt a rewrite that carries no tokens", () => { + fs.mkdirSync(gleanDir, { recursive: true }); + fs.writeFileSync( + credFile, + JSON.stringify({ tokens: { access_token: "T0" }, clientInfo: {} }), + ); + const provider = new GleanOAuthClientProvider(); + expect(provider.tokens()?.access_token).toBe("T0"); + + // A client-only rewrite (tokens dropped) must not log us out in-memory. + writeCredFile(undefined, { client_id: "cid" }); + expect(provider.tokens()?.access_token).toBe("T0"); + }); + + it("invalidateCredentials('tokens') adopts a sibling's token instead of wiping the store", async () => { + fs.mkdirSync(gleanDir, { recursive: true }); + fs.writeFileSync( + credFile, + JSON.stringify({ + tokens: { access_token: "T0", refresh_token: "R0" }, + clientInfo: { client_id: "cid" }, + }), + ); + const provider = new GleanOAuthClientProvider(); + expect(provider.tokens()?.access_token).toBe("T0"); + + // A sibling refreshed + rotated: fresh grant is now on disk. + writeCredFile( + { access_token: "T1", refresh_token: "R1" }, + { client_id: "cid" }, + ); + + // The SDK calls this on invalid_grant. It must NOT clear — the failure was + // just our stale token; adopt the sibling's fresh one and leave it on disk. + await provider.invalidateCredentials("tokens"); + + expect(provider.tokens()?.access_token).toBe("T1"); + expect(provider.tokens()?.refresh_token).toBe("R1"); + const raw = JSON.parse(fs.readFileSync(credFile, "utf-8")); + expect(raw.tokens.access_token).toBe("T1"); // not clobbered with undefined + }); + + it("invalidateCredentials('tokens') clears when there is no newer token on disk", async () => { + const provider = new GleanOAuthClientProvider(); + provider.saveClientInformation({ client_id: "cid" }); + provider.saveTokens({ access_token: "T0", refresh_token: "R0" } as any); + expect(provider.tokens()?.access_token).toBe("T0"); + + // No sibling write since our snapshot → a genuine invalidation → clear. + vi.useFakeTimers(); + const invalidation = provider.invalidateCredentials("tokens"); + await vi.advanceTimersByTimeAsync(2000); + await invalidation; + + expect(provider.tokens()).toBeUndefined(); + const raw = JSON.parse(fs.readFileSync(credFile, "utf-8")); + expect(raw.tokens).toBeUndefined(); + }); + + it("adopts a sibling token at the next 500 ms poll", async () => { + vi.useFakeTimers(); + const provider = new GleanOAuthClientProvider(); + provider.saveClientInformation({ client_id: "cid" }); + provider.saveTokens({ access_token: "T0", refresh_token: "R0" } as any); + const completed = vi.fn(); + const invalidation = provider.invalidateCredentials("tokens").then(completed); + + await vi.advanceTimersByTimeAsync(150); + writeCredFile( + { access_token: "T1", refresh_token: "R1" }, + { client_id: "cid" }, + ); + await vi.advanceTimersByTimeAsync(349); + expect(completed).not.toHaveBeenCalled(); + await vi.advanceTimersByTimeAsync(1); + await invalidation; + + expect(provider.tokens()?.access_token).toBe("T1"); + const raw = JSON.parse(fs.readFileSync(credFile, "utf-8")); + expect(raw.tokens.access_token).toBe("T1"); + }); + + it("skips the grace window when no refresh token was held", async () => { + vi.useFakeTimers(); + const provider = new GleanOAuthClientProvider(); + provider.saveClientInformation({ client_id: "cid" }); + provider.saveTokens({ access_token: "T0" } as any); + + const start = Date.now(); + await provider.invalidateCredentials("tokens"); + + expect(Date.now()).toBe(start); + expect(vi.getTimerCount()).toBe(0); + expect(provider.tokens()).toBeUndefined(); + }); + + it("skips the grace window without a retained client", async () => { + vi.useFakeTimers(); + const provider = new GleanOAuthClientProvider(); + provider.saveTokens({ access_token: "T0", refresh_token: "R0" } as any); + + await provider.invalidateCredentials("tokens"); + + expect(vi.getTimerCount()).toBe(0); + expect(provider.tokens()).toBeUndefined(); + }); + it("saveClientInformation persists to disk", () => { const provider = new GleanOAuthClientProvider(); const info = { client_id: "cid", client_secret: "sec" } as any; @@ -133,15 +295,21 @@ describe("GleanOAuthClientProvider", () => { }); it("invalidateCredentials('all') clears all in-memory state and deletes file", async () => { + vi.useFakeTimers(); const provider = new GleanOAuthClientProvider(); - provider.saveTokens({ access_token: "tok", token_type: "Bearer" } as any); + provider.saveTokens({ access_token: "tok", refresh_token: "refresh", token_type: "Bearer" } as any); provider.saveClientInformation({ client_id: "cid" } as any); provider.saveCodeVerifier("verifier"); await provider.redirectToAuthorization(new URL("https://example.com/oauth/authorize?state=s1")); expect(fs.existsSync(path.join(gleanDir, "mcp-credentials.json"))).toBe(true); + writeCredFile( + { access_token: "sibling", refresh_token: "sibling-refresh" }, + { client_id: "cid" }, + ); await provider.invalidateCredentials("all"); + expect(vi.getTimerCount()).toBe(0); expect(provider.tokens()).toBeUndefined(); expect(provider.clientInformation()).toBeUndefined(); expect(provider.codeVerifier()).toBe(""); @@ -149,12 +317,17 @@ describe("GleanOAuthClientProvider", () => { expect(fs.existsSync(path.join(gleanDir, "mcp-credentials.json"))).toBe(false); }); - it("invalidateCredentials('client') drops client but keeps tokens", async () => { + it("invalidateCredentials('client') drops client without waiting but keeps tokens", async () => { + vi.useFakeTimers(); const provider = new GleanOAuthClientProvider(); - provider.saveTokens({ access_token: "tok" } as any); - provider.saveClientInformation({ client_id: "cid" } as any); + const tokens = { access_token: "tok", refresh_token: "refresh", token_type: "Bearer" }; + provider.saveTokens(tokens); + provider.saveClientInformation({ client_id: "cid" }); + await provider.invalidateCredentials("client"); - expect(provider.tokens()).toEqual({ access_token: "tok" }); + + expect(vi.getTimerCount()).toBe(0); + expect(provider.tokens()).toEqual(tokens); expect(provider.clientInformation()).toBeUndefined(); }); diff --git a/shared/glean/mcp/tests/remote-client-auth-retry.test.ts b/shared/glean/mcp/tests/remote-client-auth-retry.test.ts new file mode 100644 index 0000000..3661a52 --- /dev/null +++ b/shared/glean/mcp/tests/remote-client-auth-retry.test.ts @@ -0,0 +1,216 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { + UnauthorizedError, + OAuthError, + OAuthErrorCode, +} from "@modelcontextprotocol/client"; + +// Control client.connect() across (re)tries while preserving the real SDK errors. +const { connectMock } = vi.hoisted(() => ({ connectMock: vi.fn() })); + +vi.mock("@modelcontextprotocol/client", async (importOriginal) => ({ + ...(await importOriginal()), + Client: class { + async connect(...args: unknown[]) { + return connectMock(...args); + } + }, + StreamableHTTPClientTransport: class { + constructor() {} + async close() {} + }, +})); + +const { createRemoteClient, AuthRequiredError } = await import( + "../src/remote-client.js" +); + +const serverUrl = "https://acme-be.glean.com/mcp/gateway/proxy"; + +/** + * Minimal OAuthClientProvider stand-in. tokens() returns the next value in + * `seq` on each call, mirroring how the real provider re-reads disk: the + * pre-connect snapshot, then the value after a sibling may have rewritten it. + */ +function makeProvider(seq: Array<{ access_token?: string } | undefined>) { + let i = 0; + return { + tokens() { + const t = seq[Math.min(i, seq.length - 1)]; + i += 1; + return t; + }, + authorizationUrl: "https://example.com/oauth/authorize?state=s1", + pendingAuthCode: undefined, + needsFreshClient: () => false, + } as any; +} + +beforeEach(() => { + connectMock.mockReset(); +}); + +describe("createRemoteClient sibling-refresh retry", () => { + it("retries once and succeeds when a newer token appears on disk", async () => { + connectMock + .mockRejectedValueOnce(new UnauthorizedError("401")) + .mockResolvedValueOnce(undefined); + + // Pre-connect T0, post-failure T1, retry snapshot T1. + const provider = makeProvider([ + { access_token: "T0" }, + { access_token: "T1" }, + { access_token: "T1" }, + ]); + + const client = await createRemoteClient( + serverUrl, + { authProvider: provider }, + "sess-1", + ); + + expect(client).toBeTruthy(); + expect(connectMock).toHaveBeenCalledTimes(2); + }); + + it("retries when a sibling supplies the first available token", async () => { + connectMock + .mockRejectedValueOnce(new UnauthorizedError("401")) + .mockResolvedValueOnce(undefined); + const provider = makeProvider([undefined, { access_token: "T1" }]); + + const client = await createRemoteClient(serverUrl, { authProvider: provider }); + + expect(client).toBeTruthy(); + expect(connectMock).toHaveBeenCalledTimes(2); + }); + + it("does not retry when the on-disk token is unchanged", async () => { + connectMock.mockRejectedValue(new UnauthorizedError("401")); + const provider = makeProvider([{ access_token: "T0" }]); + + await expect( + createRemoteClient(serverUrl, { authProvider: provider }), + ).rejects.toBeInstanceOf(AuthRequiredError); + + expect(connectMock).toHaveBeenCalledTimes(1); + }); + + it("does not retry twice even if another token appears after the retry fails", async () => { + connectMock.mockRejectedValue(new UnauthorizedError("401")); + const provider = makeProvider([ + { access_token: "T0" }, + { access_token: "T1" }, + { access_token: "T1" }, + { access_token: "T2" }, + ]); + + await expect( + createRemoteClient(serverUrl, { authProvider: provider }), + ).rejects.toBeInstanceOf(AuthRequiredError); + + expect(connectMock).toHaveBeenCalledTimes(2); + }); +}); + +function makeCollisionProvider(siblingRefreshed: boolean) { + let accessToken = "T0"; + return { + tokens: () => ({ access_token: accessToken, refresh_token: "R0" }), + authorizationUrl: undefined, + pendingAuthCode: undefined, + needsFreshClient: () => false, + waitForSiblingRefresh: vi.fn(async () => { + if (siblingRefreshed) accessToken = "T1"; + return siblingRefreshed; + }), + invalidateCredentials: vi.fn(), + } as any; +} + +describe.each([ + OAuthErrorCode.InvalidRequest, + OAuthErrorCode.InvalidGrant, +])("createRemoteClient %s refresh-collision retry", (code) => { + const collisionError = new OAuthError(code, "Refresh request rejected"); + + it("retries once when a sibling's refresh lands during the grace wait", async () => { + connectMock + .mockRejectedValueOnce(collisionError) + .mockResolvedValueOnce(undefined); + const provider = makeCollisionProvider(true); + + const client = await createRemoteClient(serverUrl, { authProvider: provider }); + + expect(client).toBeTruthy(); + expect(connectMock).toHaveBeenCalledTimes(2); + expect(provider.waitForSiblingRefresh).toHaveBeenCalledExactlyOnceWith("T0"); + }); + + it("rethrows when no sibling token appears within the grace window", async () => { + connectMock.mockRejectedValue(collisionError); + const provider = makeCollisionProvider(false); + + await expect( + createRemoteClient(serverUrl, { authProvider: provider }), + ).rejects.toBe(collisionError); + + expect(connectMock).toHaveBeenCalledTimes(1); + }); + + it("rethrows a failed retry without waiting or connecting again", async () => { + connectMock.mockRejectedValue(collisionError); + const provider = makeCollisionProvider(true); + + await expect( + createRemoteClient(serverUrl, { authProvider: provider }), + ).rejects.toBe(collisionError); + + expect(connectMock).toHaveBeenCalledTimes(2); + expect(provider.waitForSiblingRefresh).toHaveBeenCalledTimes(1); + }); +}); + +describe("createRemoteClient non-recoverable errors", () => { + it("rethrows unauthorized errors without a newer token or a sign-in URL", async () => { + const error = new UnauthorizedError("401"); + connectMock.mockRejectedValue(error); + const provider = makeCollisionProvider(false); + + await expect( + createRemoteClient(serverUrl, { authProvider: provider }), + ).rejects.toBe(error); + + expect(connectMock).toHaveBeenCalledTimes(1); + expect(provider.waitForSiblingRefresh).not.toHaveBeenCalled(); + }); + + it.each([ + new Error("Failed to refresh token"), + new OAuthError(OAuthErrorCode.InvalidClient, "Invalid client"), + new OAuthError(OAuthErrorCode.InvalidScope, "Invalid scope"), + { code: OAuthErrorCode.InvalidGrant }, + ])("does not retry an unrelated or untyped error: %s", async (error) => { + connectMock.mockRejectedValue(error); + const provider = makeCollisionProvider(true); + + await expect( + createRemoteClient(serverUrl, { authProvider: provider }), + ).rejects.toBe(error); + + expect(connectMock).toHaveBeenCalledTimes(1); + expect(provider.waitForSiblingRefresh).not.toHaveBeenCalled(); + }); + + it.each([ + new UnauthorizedError("401"), + new OAuthError(OAuthErrorCode.InvalidRequest, "Invalid request"), + new OAuthError(OAuthErrorCode.InvalidGrant, "Invalid grant"), + ])("rethrows without an auth provider and does not retry: %s", async (error) => { + connectMock.mockRejectedValue(error); + + await expect(createRemoteClient(serverUrl, {})).rejects.toBe(error); + + expect(connectMock).toHaveBeenCalledTimes(1); + }); +}); diff --git a/shared/glean/mcp/tests/token-store.test.ts b/shared/glean/mcp/tests/token-store.test.ts index 14fd6cc..161ad8d 100644 --- a/shared/glean/mcp/tests/token-store.test.ts +++ b/shared/glean/mcp/tests/token-store.test.ts @@ -10,19 +10,21 @@ vi.mock("node:os", async () => { return { ...actual, homedir: () => tmpDir }; }); -const { clearCredentials, loadCredentials, saveCredentials } = await import( - "../src/token-store.js" -); +const { clearCredentials, loadCredentials, saveCredentials } = + await import("../src/token-store.js"); describe("token-store", () => { const gleanDir = path.join(tmpDir, ".glean"); const credFile = path.join(gleanDir, "mcp-credentials.json"); beforeEach(() => { + vi.stubEnv("PLUGIN_DATA_DIR", gleanDir); fs.rmSync(gleanDir, { recursive: true, force: true }); }); afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllEnvs(); fs.rmSync(gleanDir, { recursive: true, force: true }); }); @@ -56,6 +58,41 @@ describe("token-store", () => { expect(mode).toBe(0o600); }); + it("sets the credentials directory to mode 0700", () => { + fs.mkdirSync(gleanDir, { recursive: true, mode: 0o755 }); + + saveCredentials({ access_token: "x" }, undefined); + + expect(fs.statSync(gleanDir).mode & 0o777).toBe(0o700); + }); + + it("tightens a leftover temp file before replacing credentials", () => { + fs.mkdirSync(gleanDir, { recursive: true }); + const tmpPath = path.join(gleanDir, `.mcp-credentials.json.${process.pid}.tmp`); + fs.writeFileSync(tmpPath, "stale", { mode: 0o644 }); + + saveCredentials({ access_token: "new" }, undefined); + + expect(loadCredentials()?.tokens?.access_token).toBe("new"); + expect(fs.statSync(credFile).mode & 0o777).toBe(0o600); + expect(fs.readdirSync(gleanDir)).toEqual(["mcp-credentials.json"]); + }); + + it("preserves credentials and removes the temp file when rename fails", () => { + const original = { access_token: "old" }; + saveCredentials(original, { client_id: "cid" }); + vi.spyOn(fs, "renameSync").mockImplementationOnce(() => { + throw new Error("rename blocked"); + }); + const log = vi.spyOn(console, "error").mockImplementation(() => {}); + + expect(() => saveCredentials({ access_token: "new" }, undefined)).not.toThrow(); + + expect(loadCredentials()).toEqual({ tokens: original, clientInfo: { client_id: "cid" } }); + expect(fs.readdirSync(gleanDir)).toEqual(["mcp-credentials.json"]); + expect(log).toHaveBeenCalledWith("[auth] Failed to persist credentials: rename blocked"); + }); + it("returns undefined for corrupted JSON", () => { fs.mkdirSync(gleanDir, { recursive: true }); fs.writeFileSync(credFile, "not-json{{{", "utf-8"); @@ -88,4 +125,5 @@ describe("token-store", () => { expect(fs.existsSync(credFile)).toBe(false); expect(() => clearCredentials()).not.toThrow(); }); + });