diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index 20b777eea..ae04587ca 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -336,6 +336,7 @@ export interface AuthCredentialStore { tryAcquireCredentialRefreshLease?(credentialId: number, owner: string, expiresAtMs: number): boolean; getCredentialRefreshLeaseExpiresAt?(credentialId: number): number | undefined; releaseCredentialRefreshLease?(credentialId: number, owner: string): void; + renewCredentialRefreshLease?(credentialId: number, owner: string, expiresAtMs: number): boolean; /** * Append usage-limit snapshots for trend history. Optional: stores without * durable storage (e.g. the broker remote store) omit it and recording is @@ -599,6 +600,7 @@ const DEFAULT_OAUTH_REFRESH_TIMEOUT_MS = 10_000; const OAUTH_REFRESH_SKEW_MS = 60_000; const OAUTH_REFRESH_LEASE_TTL_MS = 15_000; const OAUTH_REFRESH_LEASE_POLL_MS = 50; +const OAUTH_REFRESH_LEASE_RENEW_MS = 5_000; /** * Cap on the buffered credential_disabled backlog held while no handler is attached. * In practice the backlog is 0–N where N ≈ active providers (≤ ~20). The cap exists so @@ -1873,7 +1875,8 @@ export class AuthStorage { const hasDurableLease = !!this.#store.tryAcquireCredentialRefreshLease && !!this.#store.getCredentialRefreshLeaseExpiresAt && - !!this.#store.releaseCredentialRefreshLease; + !!this.#store.releaseCredentialRefreshLease && + !!this.#store.renewCredentialRefreshLease; const owner = crypto.randomUUID(); let leasedCredentialId: number | undefined; @@ -1943,42 +1946,76 @@ export class AuthStorage { const serialized = serializeCredential(provider, current); if (!serialized) return { credential: current, refreshed: false, removed: false }; + let stopLeaseRenewal = false; + let leaseRenewalError: unknown; + const leaseRenewalStopped = Promise.withResolvers(); + const leaseRenewal = + leasedCredentialId !== undefined + ? (async () => { + while (!stopLeaseRenewal) { + await Promise.race([Bun.sleep(OAUTH_REFRESH_LEASE_RENEW_MS), leaseRenewalStopped.promise]); + if (stopLeaseRenewal) return; + const renewed = this.#store.renewCredentialRefreshLease?.( + leasedCredentialId, + owner, + Date.now() + OAUTH_REFRESH_LEASE_TTL_MS, + ); + if (!renewed) { + throw new AIError.ConfigurationError("OAuth refresh ownership was lost before persistence"); + } + } + })().catch(error => { + leaseRenewalError = error; + }) + : undefined; + let refreshed: OAuthCredentials; try { - refreshed = await options.refresh(current); - } catch (error) { - if (options.isDefinitiveFailure?.(error)) { - const disabledCause = options.disabledCause?.(error) ?? `oauth refresh failed: ${String(error)}`; - const disabled = this.#store.tryDisableAuthCredentialIfMatches(row.id, serialized.data, disabledCause); - if (disabled) { - this.#setStoredCredentials( - provider, - rows - .filter(entry => entry.id !== row.id) - .map(entry => ({ id: entry.id, credential: entry.credential })), + try { + refreshed = await options.refresh(current); + } catch (error) { + if (options.isDefinitiveFailure?.(error)) { + const disabledCause = options.disabledCause?.(error) ?? `oauth refresh failed: ${String(error)}`; + const disabled = this.#store.tryDisableAuthCredentialIfMatches( + row.id, + serialized.data, + disabledCause, ); - this.#resetProviderAssignments(provider); - this.#emitCredentialDisabled({ provider, disabledCause }); - return { credential: undefined, refreshed: false, removed: true }; + if (disabled) { + this.#setStoredCredentials( + provider, + rows + .filter(entry => entry.id !== row.id) + .map(entry => ({ id: entry.id, credential: entry.credential })), + ); + this.#resetProviderAssignments(provider); + this.#emitCredentialDisabled({ provider, disabledCause }); + return { credential: undefined, refreshed: false, removed: true }; + } + await this.reload(); + const latest = this.get(provider); + return { + credential: latest?.type === "oauth" ? options.credentialFromRow(latest) : undefined, + refreshed: false, + removed: false, + }; } - await this.reload(); - const latest = this.get(provider); - return { - credential: latest?.type === "oauth" ? options.credentialFromRow(latest) : undefined, - refreshed: false, - removed: false, - }; + options.onRefreshFailure?.(error); + const keepCredential = + typeof options.keepCredentialOnRefreshFailure === "function" + ? options.keepCredentialOnRefreshFailure(error) + : options.keepCredentialOnRefreshFailure === true; + if (keepCredential) { + return { credential: current, refreshed: false, removed: false }; + } + throw error; } - options.onRefreshFailure?.(error); - const keepCredential = - typeof options.keepCredentialOnRefreshFailure === "function" - ? options.keepCredentialOnRefreshFailure(error) - : options.keepCredentialOnRefreshFailure === true; - if (keepCredential) { - return { credential: current, refreshed: false, removed: false }; - } - throw error; + } finally { + stopLeaseRenewal = true; + leaseRenewalStopped.resolve(); + await leaseRenewal; } + if (leaseRenewalError) throw leaseRenewalError; const merged: T = options.mergeRefreshedCredential ? options.mergeRefreshedCredential(current, refreshed) @@ -5413,6 +5450,7 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { #deleteExpiredCredentialBlocksStmt: Statement; #acquireCredentialRefreshLeaseStmt: Statement; #getCredentialRefreshLeaseStmt: Statement; + #renewCredentialRefreshLeaseStmt: Statement; #releaseCredentialRefreshLeaseStmt: Statement; #credentialBlockReconcileAfter: Map = new Map(); #insertUsageHistoryStmt: Statement; @@ -5492,6 +5530,9 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { this.#getCredentialRefreshLeaseStmt = this.#db.prepare( "SELECT expires_at_ms FROM auth_credential_refresh_leases WHERE credential_id = ?", ); + this.#renewCredentialRefreshLeaseStmt = this.#db.prepare( + `UPDATE auth_credential_refresh_leases SET expires_at_ms = ?, updated_at = ${SQLITE_NOW_EPOCH} WHERE credential_id = ? AND owner = ?`, + ); this.#releaseCredentialRefreshLeaseStmt = this.#db.prepare( "DELETE FROM auth_credential_refresh_leases WHERE credential_id = ? AND owner = ?", ); @@ -6058,32 +6099,28 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { } tryUpdateAuthCredentialIfMatches(id: number, expectedData: string, credential: AuthCredential): boolean { + const providerStmt = this.#db.prepare("SELECT provider FROM auth_credentials WHERE id = ?"); + let providerRow: { provider?: string } | undefined; try { - const providerStmt = this.#db.prepare("SELECT provider FROM auth_credentials WHERE id = ?"); - let providerRow: { provider?: string } | undefined; - try { - providerRow = providerStmt.get(id) as { provider?: string } | undefined; - } finally { - providerStmt.finalize(); - } - const provider = providerRow?.provider ?? ""; - const serialized = serializeCredential(provider, credential); - if (!serialized) return false; - const result = this.#updateIfMatchesStmt.run( - serialized.credentialType, - serialized.data, - serialized.identityKey, - id, - expectedData, - ) as { changes: number }; - if (result.changes !== 1) return false; - if (provider) { - this.#purgeSupersededDisabledRows(provider, this.listAuthCredentials(provider)); - } - return true; - } catch { - return false; + providerRow = providerStmt.get(id) as { provider?: string } | undefined; + } finally { + providerStmt.finalize(); } + const provider = providerRow?.provider ?? ""; + const serialized = serializeCredential(provider, credential); + if (!serialized) return false; + const result = this.#updateIfMatchesStmt.run( + serialized.credentialType, + serialized.data, + serialized.identityKey, + id, + expectedData, + ) as { changes: number }; + if (result.changes !== 1) return false; + if (provider) { + this.#purgeSupersededDisabledRows(provider, this.listAuthCredentials(provider)); + } + return true; } deleteAuthCredential(id: number, disabledCause: string): void { @@ -6101,16 +6138,11 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { * row between our pre-check and the disable. */ tryDisableAuthCredentialIfMatches(id: number, expectedData: string, disabledCause: string): boolean { - try { - const result = this.#deleteIfMatchesStmt.run(normalizeDisabledCause(disabledCause), id, expectedData) as { - changes: number; - }; - return result.changes === 1; - } catch { - return false; - } + const result = this.#deleteIfMatchesStmt.run(normalizeDisabledCause(disabledCause), id, expectedData) as { + changes: number; + }; + return result.changes === 1; } - deleteAuthCredentialsForProvider(provider: string, disabledCause: string): void { try { this.#deleteByProviderStmt.run(normalizeDisabledCause(disabledCause), provider); @@ -6233,6 +6265,13 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { return row.expires_at_ms; } + renewCredentialRefreshLease(credentialId: number, owner: string, expiresAtMs: number): boolean { + const result = this.#renewCredentialRefreshLeaseStmt.run(expiresAtMs, credentialId, owner) as { + changes: number; + }; + return result.changes === 1; + } + releaseCredentialRefreshLease(credentialId: number, owner: string): void { try { this.#releaseCredentialRefreshLeaseStmt.run(credentialId, owner); diff --git a/packages/ai/test/auth-storage-oauth-refresh-race.test.ts b/packages/ai/test/auth-storage-oauth-refresh-race.test.ts index 6b9323958..b5179e450 100644 --- a/packages/ai/test/auth-storage-oauth-refresh-race.test.ts +++ b/packages/ai/test/auth-storage-oauth-refresh-race.test.ts @@ -432,4 +432,82 @@ describe("AuthStorage OAuth refresh race", () => { expect(cRow?.credential.type).toBe("oauth"); if (cRow?.credential.type === "oauth") expect(cRow.credential.refresh).toBe("c-ref"); }); + + test("propagates CAS update storage errors instead of treating them as peer refresh wins", async () => { + if (!authStorage || !store) throw new Error("test setup failed"); + + await authStorage.set("unit-oauth-cas-update-error", [ + { + type: "oauth", + access: "access-old", + refresh: "refresh-old", + expires: Date.now() - 60_000, + }, + ]); + + const failure = new Error("sqlite update failed"); + vi.spyOn(store, "tryUpdateAuthCredentialIfMatches").mockImplementation(() => { + throw failure; + }); + + await expect( + authStorage.refreshStoredOAuthCredential("unit-oauth-cas-update-error", { + credentialFromRow: row => row, + forceRefresh: true, + refresh: async credential => ({ + ...credential, + access: "access-fresh", + refresh: "refresh-fresh", + expires: Date.now() + 60 * 60_000, + }), + }), + ).rejects.toThrow("sqlite update failed"); + + const stored = store.listAuthCredentials("unit-oauth-cas-update-error"); + expect(stored).toHaveLength(1); + expect(stored[0]?.credential).toMatchObject({ + type: "oauth", + access: "access-old", + refresh: "refresh-old", + }); + }); + + test("propagates CAS disable storage errors instead of treating them as peer rotations", async () => { + if (!authStorage || !store) throw new Error("test setup failed"); + + await authStorage.set("unit-oauth-cas-disable-error", [ + { + type: "oauth", + access: "access-old", + refresh: "refresh-old", + expires: Date.now() - 60_000, + }, + ]); + + const failure = new Error("sqlite disable failed"); + vi.spyOn(store, "tryDisableAuthCredentialIfMatches").mockImplementation(() => { + throw failure; + }); + + await expect( + authStorage.refreshStoredOAuthCredential("unit-oauth-cas-disable-error", { + credentialFromRow: row => row, + forceRefresh: true, + refresh: async () => { + throw new Error('HTTP 400 invalid_grant {"error":"invalid_grant"}'); + }, + isDefinitiveFailure: error => error instanceof Error && error.message.includes("invalid_grant"), + disabledCause: error => `oauth refresh failed: ${error instanceof Error ? error.message : String(error)}`, + }), + ).rejects.toThrow("sqlite disable failed"); + + expect(events).toHaveLength(0); + const stored = store.listAuthCredentials("unit-oauth-cas-disable-error"); + expect(stored).toHaveLength(1); + expect(stored[0]?.credential).toMatchObject({ + type: "oauth", + access: "access-old", + refresh: "refresh-old", + }); + }); }); diff --git a/packages/coding-agent/test/mcp-manager-oauth-refresh.test.ts b/packages/coding-agent/test/mcp-manager-oauth-refresh.test.ts index 7c5f1521e..8ad061a0f 100644 --- a/packages/coding-agent/test/mcp-manager-oauth-refresh.test.ts +++ b/packages/coding-agent/test/mcp-manager-oauth-refresh.test.ts @@ -10,7 +10,7 @@ * Bearer injection, so the next request surfaces a clean auth error instead. */ import { Database } from "bun:sqlite"; -import { afterEach, beforeEach, describe, expect, test, vi } from "bun:test"; +import { afterEach, beforeEach, describe, expect, setSystemTime, test, vi } from "bun:test"; import * as fs from "node:fs/promises"; import * as os from "node:os"; import * as path from "node:path"; @@ -37,6 +37,54 @@ function getAuthorizationHeader(config: MCPServerConfig): string | undefined { return config.headers?.Authorization; } +type ControlledSleep = { + ms: number; + resolved: boolean; + resolve: () => void; +}; + +function installControlledBunSleep(): ControlledSleep[] { + const calls: ControlledSleep[] = []; + vi.spyOn(Bun, "sleep").mockImplementation((ms: number | Date) => { + const { promise, resolve } = Promise.withResolvers(); + let call: ControlledSleep; + const delayMs = typeof ms === "number" ? ms : Math.max(0, ms.getTime() - Date.now()); + call = { + ms: delayMs, + resolved: false, + resolve: () => { + if (call.resolved) return; + call.resolved = true; + resolve(); + }, + }; + calls.push(call); + return promise; + }); + return calls; +} + +async function drainMicrotasks(count = 10): Promise { + for (let attempt = 0; attempt < count; attempt++) { + await Promise.resolve(); + } +} + +async function waitForControlledSleep(calls: ControlledSleep[], ms: number): Promise { + for (let attempt = 0; attempt < 20; attempt++) { + const call = calls.find(candidate => !candidate.resolved && candidate.ms === ms); + if (call) return call; + await Promise.resolve(); + } + throw new Error(`Timed out waiting for Bun.sleep(${ms})`); +} + +function resolvePendingControlledSleeps(calls: ControlledSleep[]): void { + for (const call of calls) { + if (!call.resolved) call.resolve(); + } +} + async function withSharedSQLiteAuth( fn: (context: { authA: AuthStorage; @@ -176,6 +224,98 @@ describe("MCPManager OAuth refresh failure", () => { }); describe("MCPManager shared SQLite OAuth refresh", () => { + afterEach(() => { + setSystemTime(); + vi.restoreAllMocks(); + }); + + test("renews refresh ownership while the token endpoint is blocked", async () => { + await withSharedSQLiteAuth(async ({ authA, authB, storeA }) => { + const startMs = Date.parse("2026-07-10T12:00:00.000Z"); + setSystemTime(new Date(startMs)); + const sleeps = installControlledBunSleep(); + const renewSpy = vi.spyOn(storeA, "renewCredentialRefreshLease"); + + await authA.set(SHARED_CREDENTIAL_ID, { + type: "oauth", + access: SHARED_STALE_ACCESS, + refresh: SHARED_STALE_REFRESH, + expires: startMs - 60_000, + }); + await authB.reload(); + + const refreshStarted = Promise.withResolvers(); + const allowRefreshResponse = Promise.withResolvers(); + const refreshTokens: string[] = []; + let refreshRequests = 0; + const tokenServer = Bun.serve({ + hostname: "127.0.0.1", + port: 0, + async fetch(req) { + if (req.method !== "POST" || new URL(req.url).pathname !== "/token") { + return new Response("not found", { status: 404 }); + } + refreshRequests += 1; + const body = new URLSearchParams(await req.text()); + refreshTokens.push(body.get("refresh_token") ?? ""); + if (refreshRequests === 1) { + refreshStarted.resolve(); + await allowRefreshResponse.promise; + return Response.json({ + access_token: SHARED_FRESH_ACCESS, + refresh_token: SHARED_FRESH_REFRESH, + expires_in: 3600, + }); + } + return Response.json({ error: "invalid_grant" }, { status: 400 }); + }, + }); + try { + const managerA = new MCPManager(process.cwd()); + managerA.setAuthStorage(authA); + const managerB = new MCPManager(process.cwd()); + managerB.setAuthStorage(authB); + const config: MCPServerConfig = { + type: "http", + url: "https://logfire.example.com/mcp", + auth: { + type: "oauth", + credentialId: SHARED_CREDENTIAL_ID, + tokenUrl: `http://127.0.0.1:${tokenServer.port}/token`, + }, + }; + + const preparedA = managerA.prepareConfig(config); + await refreshStarted.promise; + const renewalSleep = await waitForControlledSleep(sleeps, 5_000); + + setSystemTime(new Date(startMs + 5_000)); + renewalSleep.resolve(); + await drainMicrotasks(); + expect(renewSpy).toHaveBeenCalledTimes(1); + + setSystemTime(new Date(startMs + 16_000)); + const preparedB = managerB.prepareConfig(config); + const peerLeaseWait = await waitForControlledSleep(sleeps, 250); + + expect(refreshRequests).toBe(1); + expect(refreshTokens).toEqual([SHARED_STALE_REFRESH]); + + allowRefreshResponse.resolve(); + const resolvedA = await preparedA; + peerLeaseWait.resolve(); + const resolvedB = await preparedB; + + expect(getAuthorizationHeader(resolvedA)).toBe(`Bearer ${SHARED_FRESH_ACCESS}`); + expect(getAuthorizationHeader(resolvedB)).toBe(`Bearer ${SHARED_FRESH_ACCESS}`); + expect(refreshRequests).toBe(1); + expect(refreshTokens).toEqual([SHARED_STALE_REFRESH]); + } finally { + resolvePendingControlledSleeps(sleeps); + tokenServer.stop(true); + } + }); + }); test("shares refresh ownership so peer managers do not replay a rotating refresh token", async () => { await withSharedSQLiteAuth(async ({ authA, authB }) => { await authA.set(SHARED_CREDENTIAL_ID, {