diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index ae04587ca..60e535093 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -308,13 +308,28 @@ export interface AuthCredentialSnapshot { * a remote broker; mutating methods (`replace*`, `upsert*`, `delete*ForProvider`) * throw because login flows route through the broker, not the client. */ +export interface CredentialRefreshLeaseFence { + owner: string; + nowMs: number; +} + export interface AuthCredentialStore { close(): void; listAuthCredentials(provider?: string): StoredAuthCredential[]; updateAuthCredential(id: number, credential: AuthCredential): void; deleteAuthCredential(id: number, disabledCause: string): void; - tryDisableAuthCredentialIfMatches(id: number, expectedData: string, disabledCause: string): boolean; - tryUpdateAuthCredentialIfMatches?(id: number, expectedData: string, credential: AuthCredential): boolean; + tryDisableAuthCredentialIfMatches( + id: number, + expectedData: string, + disabledCause: string, + lease?: CredentialRefreshLeaseFence, + ): boolean; + tryUpdateAuthCredentialIfMatches?( + id: number, + expectedData: string, + credential: AuthCredential, + lease?: CredentialRefreshLeaseFence, + ): boolean; replaceAuthCredentialsForProvider(provider: string, credentials: AuthCredential[]): StoredAuthCredential[]; upsertAuthCredentialForProvider(provider: string, credential: AuthCredential): StoredAuthCredential[]; deleteAuthCredentialsForProvider(provider: string, disabledCause: string): void; @@ -601,6 +616,7 @@ 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; +const OAUTH_REFRESH_OPERATION_TIMEOUT_MS = 10_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 @@ -736,7 +752,8 @@ export interface StoredOAuthRefreshOptions boolean); onRefreshFailure?: (error: unknown) => void; - refresh: (credential: T) => Promise; + refreshTimeoutMs?: number; + refresh: (credential: T, signal?: AbortSignal) => Promise; mergeRefreshedCredential?: (credential: T, refreshed: OAuthCredentials) => T; isDefinitiveFailure?: (error: unknown) => boolean; disabledCause?: (error: unknown) => string; @@ -1968,11 +1985,20 @@ export class AuthStorage { leaseRenewalError = error; }) : undefined; + const refreshAbort = new AbortController(); + const refreshTimeout = setTimeout(() => { + refreshAbort.abort( + new AIError.OAuthError(`OAuth token refresh timed out for provider: ${provider}`, { + kind: "timeout", + provider, + }), + ); + }, options.refreshTimeoutMs ?? OAUTH_REFRESH_OPERATION_TIMEOUT_MS); let refreshed: OAuthCredentials; try { try { - refreshed = await options.refresh(current); + refreshed = await options.refresh(current, refreshAbort.signal); } catch (error) { if (options.isDefinitiveFailure?.(error)) { const disabledCause = options.disabledCause?.(error) ?? `oauth refresh failed: ${String(error)}`; @@ -1980,6 +2006,7 @@ export class AuthStorage { row.id, serialized.data, disabledCause, + leasedCredentialId !== undefined ? { owner, nowMs: Date.now() } : undefined, ); if (disabled) { this.#setStoredCredentials( @@ -2014,6 +2041,7 @@ export class AuthStorage { stopLeaseRenewal = true; leaseRenewalStopped.resolve(); await leaseRenewal; + clearTimeout(refreshTimeout); } if (leaseRenewalError) throw leaseRenewalError; @@ -2031,7 +2059,14 @@ export class AuthStorage { apiEndpoint: refreshed.apiEndpoint ?? current.apiEndpoint, }; if (this.#store.tryUpdateAuthCredentialIfMatches) { - if (!this.#store.tryUpdateAuthCredentialIfMatches(row.id, serialized.data, merged)) { + if ( + !this.#store.tryUpdateAuthCredentialIfMatches( + row.id, + serialized.data, + merged, + leasedCredentialId !== undefined ? { owner, nowMs: Date.now() } : undefined, + ) + ) { await this.reload(); const latest = this.get(provider); return { @@ -5443,6 +5478,8 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { #getCacheIncludingExpiredStmt: Statement; #upsertCacheStmt: Statement; #deleteExpiredCacheStmt: Statement; + #updateIfMatchesWithLeaseStmt: Statement; + #deleteIfMatchesWithLeaseStmt: Statement; #getCredentialBlockStmt: Statement; #listCredentialBlocksByCredentialStmt: Statement; #upsertCredentialBlockStmt: Statement; @@ -5483,12 +5520,30 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { this.#updateIfMatchesStmt = this.#db.prepare( `UPDATE auth_credentials SET credential_type = ?, data = ?, identity_key = ?, updated_at = ${SQLITE_NOW_EPOCH} WHERE id = ? AND data = ? AND disabled_cause IS NULL`, ); + this.#updateIfMatchesWithLeaseStmt = this.#db.prepare( + `UPDATE auth_credentials + SET credential_type = ?, data = ?, identity_key = ?, updated_at = ${SQLITE_NOW_EPOCH} + WHERE id = ? AND data = ? AND disabled_cause IS NULL + AND EXISTS ( + SELECT 1 FROM auth_credential_refresh_leases + WHERE credential_id = ? AND owner = ? AND expires_at_ms > ? + )`, + ); this.#deleteStmt = this.#db.prepare( `UPDATE auth_credentials SET disabled_cause = ?, updated_at = ${SQLITE_NOW_EPOCH} WHERE id = ?`, ); this.#deleteIfMatchesStmt = this.#db.prepare( `UPDATE auth_credentials SET disabled_cause = ?, updated_at = ${SQLITE_NOW_EPOCH} WHERE id = ? AND data = ? AND disabled_cause IS NULL`, ); + this.#deleteIfMatchesWithLeaseStmt = this.#db.prepare( + `UPDATE auth_credentials + SET disabled_cause = ?, updated_at = ${SQLITE_NOW_EPOCH} + WHERE id = ? AND data = ? AND disabled_cause IS NULL + AND EXISTS ( + SELECT 1 FROM auth_credential_refresh_leases + WHERE credential_id = ? AND owner = ? AND expires_at_ms > ? + )`, + ); this.#deleteByProviderStmt = this.#db.prepare( `UPDATE auth_credentials SET disabled_cause = ?, updated_at = ${SQLITE_NOW_EPOCH} WHERE provider = ? AND disabled_cause IS NULL`, ); @@ -6098,7 +6153,12 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { } } - tryUpdateAuthCredentialIfMatches(id: number, expectedData: string, credential: AuthCredential): boolean { + tryUpdateAuthCredentialIfMatches( + id: number, + expectedData: string, + credential: AuthCredential, + lease?: CredentialRefreshLeaseFence, + ): boolean { const providerStmt = this.#db.prepare("SELECT provider FROM auth_credentials WHERE id = ?"); let providerRow: { provider?: string } | undefined; try { @@ -6109,13 +6169,24 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { 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 }; + const result = lease + ? (this.#updateIfMatchesWithLeaseStmt.run( + serialized.credentialType, + serialized.data, + serialized.identityKey, + id, + expectedData, + id, + lease.owner, + lease.nowMs, + ) as { changes: number }) + : (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)); @@ -6137,10 +6208,24 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { * the OAuth refresh-failure path to avoid clobbering a peer that rotated the * row between our pre-check and the disable. */ - tryDisableAuthCredentialIfMatches(id: number, expectedData: string, disabledCause: string): boolean { - const result = this.#deleteIfMatchesStmt.run(normalizeDisabledCause(disabledCause), id, expectedData) as { - changes: number; - }; + tryDisableAuthCredentialIfMatches( + id: number, + expectedData: string, + disabledCause: string, + lease?: CredentialRefreshLeaseFence, + ): boolean { + const result = lease + ? (this.#deleteIfMatchesWithLeaseStmt.run( + normalizeDisabledCause(disabledCause), + id, + expectedData, + id, + lease.owner, + lease.nowMs, + ) as { changes: number }) + : (this.#deleteIfMatchesStmt.run(normalizeDisabledCause(disabledCause), id, expectedData) as { + changes: number; + }); return result.changes === 1; } deleteAuthCredentialsForProvider(provider: string, disabledCause: string): void { 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 b5179e450..312a073dd 100644 --- a/packages/ai/test/auth-storage-oauth-refresh-race.test.ts +++ b/packages/ai/test/auth-storage-oauth-refresh-race.test.ts @@ -1,4 +1,4 @@ -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"; @@ -36,6 +36,7 @@ describe("AuthStorage OAuth refresh race", () => { afterEach(async () => { vi.restoreAllMocks(); + setSystemTime(); oauthUtils.unregisterOAuthProviders("auth-storage-oauth-refresh-race-test"); store?.close(); store = null; @@ -510,4 +511,110 @@ describe("AuthStorage OAuth refresh race", () => { refresh: "refresh-old", }); }); + + test("does not persist a refresh when durable lease ownership is lost before CAS update", async () => { + if (!authStorage || !store) throw new Error("test setup failed"); + + const now = Date.parse("2026-07-10T12:00:00.000Z"); + setSystemTime(new Date(now)); + await authStorage.set("unit-oauth-lease-update", [ + { + type: "oauth", + access: "access-old", + refresh: "refresh-old", + expires: now - 60_000, + }, + ]); + const storedBefore = store.listAuthCredentials("unit-oauth-lease-update"); + expect(storedBefore).toHaveLength(1); + const credentialId = storedBefore[0]!.id; + const stealLease = store.tryAcquireCredentialRefreshLease?.bind(store); + if (!stealLease) throw new Error("test store does not support refresh leases"); + const updateSpy = vi.spyOn(store, "tryUpdateAuthCredentialIfMatches"); + + const result = await authStorage.refreshStoredOAuthCredential("unit-oauth-lease-update", { + credentialFromRow: row => row, + forceRefresh: true, + refresh: async credential => { + // Keep the credential row bytes unchanged while expiring owner A's + // lease. A non-lease-fenced final CAS would still persist this token. + setSystemTime(new Date(now + 16_000)); + expect(stealLease(credentialId, "peer-owner", now + 31_000)).toBe(true); + return { + ...credential, + access: "access-from-lost-owner", + refresh: "refresh-from-lost-owner", + expires: now + 60 * 60_000, + }; + }, + }); + + expect(updateSpy).toHaveBeenCalled(); + expect(result).toMatchObject({ refreshed: false, removed: false }); + expect(result.credential).toMatchObject({ + type: "oauth", + access: "access-old", + refresh: "refresh-old", + }); + const stored = store.listAuthCredentials("unit-oauth-lease-update"); + expect(stored).toHaveLength(1); + expect(stored[0]?.id).toBe(credentialId); + expect(stored[0]?.credential).toMatchObject({ + type: "oauth", + access: "access-old", + refresh: "refresh-old", + }); + }); + + test("does not terminal-disable a credential when durable lease ownership is lost before CAS disable", async () => { + if (!authStorage || !store) throw new Error("test setup failed"); + + const now = Date.parse("2026-07-10T12:30:00.000Z"); + setSystemTime(new Date(now)); + await authStorage.set("unit-oauth-lease-disable", [ + { + type: "oauth", + access: "access-old", + refresh: "refresh-old", + expires: now - 60_000, + }, + ]); + const storedBefore = store.listAuthCredentials("unit-oauth-lease-disable"); + expect(storedBefore).toHaveLength(1); + const credentialId = storedBefore[0]!.id; + const stealLease = store.tryAcquireCredentialRefreshLease?.bind(store); + if (!stealLease) throw new Error("test store does not support refresh leases"); + const disableSpy = vi.spyOn(store, "tryDisableAuthCredentialIfMatches"); + + const result = await authStorage.refreshStoredOAuthCredential("unit-oauth-lease-disable", { + credentialFromRow: row => row, + forceRefresh: true, + refresh: async () => { + // The row still contains the same stale refresh token. Only the lease + // fence distinguishes stale owner A from the current row owner. + setSystemTime(new Date(now + 16_000)); + expect(stealLease(credentialId, "peer-owner", now + 31_000)).toBe(true); + 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)}`, + }); + + expect(disableSpy).toHaveBeenCalled(); + expect(result).toMatchObject({ refreshed: false, removed: false }); + expect(result.credential).toMatchObject({ + type: "oauth", + access: "access-old", + refresh: "refresh-old", + }); + expect(events).toHaveLength(0); + const stored = store.listAuthCredentials("unit-oauth-lease-disable"); + expect(stored).toHaveLength(1); + expect(stored[0]?.id).toBe(credentialId); + expect(stored[0]?.credential).toMatchObject({ + type: "oauth", + access: "access-old", + refresh: "refresh-old", + }); + }); }); diff --git a/packages/coding-agent/src/mcp/manager.ts b/packages/coding-agent/src/mcp/manager.ts index 304a0d245..1cb1adeb0 100644 --- a/packages/coding-agent/src/mcp/manager.ts +++ b/packages/coding-agent/src/mcp/manager.ts @@ -1238,7 +1238,7 @@ export class MCPManager { const material = selectMcpOAuthRefreshMaterial(current, auth); return Boolean(current.refresh && material?.tokenUrl); }, - refresh: current => { + refresh: (current, signal) => { if (current.refresh === REMOTE_REFRESH_SENTINEL) { throw new Error("MCP OAuth refresh token is broker-redacted; local refresh is unavailable"); } @@ -1257,6 +1257,7 @@ export class MCPManager { return refreshMCPOAuthToken(tokenUrl, current.refresh, clientId, clientSecret, resource, { authorizationUrl, stripSameOriginResource: resourceIsFallback, + signal, }); }, mergeRefreshedCredential: (current, refreshed) => { diff --git a/packages/coding-agent/src/mcp/oauth-flow.ts b/packages/coding-agent/src/mcp/oauth-flow.ts index 36e6f9bf7..c68916638 100644 --- a/packages/coding-agent/src/mcp/oauth-flow.ts +++ b/packages/coding-agent/src/mcp/oauth-flow.ts @@ -715,6 +715,7 @@ export class MCPOAuthFlow extends OAuthCallbackFlow { */ export interface RefreshMCPOAuthTokenOptions { fetch?: FetchImpl; + signal?: AbortSignal; /** * Authorization-server URL the original grant was minted against. Used to * filter same-origin resource indicators on refresh. Defaults to `tokenUrl`'s @@ -766,6 +767,7 @@ export async function refreshMCPOAuthToken( method: "POST", headers: { "Content-Type": "application/x-www-form-urlencoded" }, body: params.toString(), + signal: optsFromTrailing?.signal, }); if (!response.ok) { 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 8ad061a0f..997e4edad 100644 --- a/packages/coding-agent/test/mcp-manager-oauth-refresh.test.ts +++ b/packages/coding-agent/test/mcp-manager-oauth-refresh.test.ts @@ -115,10 +115,11 @@ async function withSharedSQLiteAuth( describe("MCPManager OAuth refresh failure", () => { let manager: MCPManager; let authStorage: AuthStorage; + let store: SqliteAuthCredentialStore; let serverConfig: MCPServerConfig; beforeEach(async () => { - const store = new SqliteAuthCredentialStore(new Database(":memory:")); + store = new SqliteAuthCredentialStore(new Database(":memory:")); authStorage = new AuthStorage(store); await authStorage.reload(); @@ -147,6 +148,7 @@ describe("MCPManager OAuth refresh failure", () => { }); afterEach(() => { + vi.useRealTimers(); authStorage.close(); vi.restoreAllMocks(); }); @@ -169,7 +171,7 @@ describe("MCPManager OAuth refresh failure", () => { undefined, undefined, "https://logfire.example.com/mcp", - { authorizationUrl: undefined, stripSameOriginResource: true }, + { authorizationUrl: undefined, stripSameOriginResource: true, signal: expect.any(AbortSignal) }, ); // The poisoned Bearer must not be re-injected — that is the loop the user // reported (#1908). @@ -221,6 +223,60 @@ describe("MCPManager OAuth refresh failure", () => { const remaining = authStorage.get(CREDENTIAL_ID); expect(remaining).toMatchObject({ type: "oauth", access: "fresh-access", refresh: "fresh-refresh" }); }); + + test("aborts a timed-out token fetch and waits for it before releasing refresh ownership", async () => { + vi.useFakeTimers(); + const fetchCalled = Promise.withResolvers(); + const abortObserved = Promise.withResolvers(); + const allowFetchReject = Promise.withResolvers(); + let capturedSignal: AbortSignal | undefined; + let preparedSettled = false; + const releaseSpy = vi.spyOn(store, "releaseCredentialRefreshLease"); + const fetchImpl = Object.assign( + async (_input: string | URL | Request, init?: RequestInit | BunFetchRequestInit): Promise => { + capturedSignal = init?.signal ?? undefined; + if (!capturedSignal) throw new Error("token refresh fetch did not receive an AbortSignal"); + fetchCalled.resolve(); + capturedSignal.addEventListener( + "abort", + () => { + abortObserved.resolve(); + }, + { once: true }, + ); + await allowFetchReject.promise; + throw capturedSignal.reason ?? new Error("fetch aborted"); + }, + { preconnect: globalThis.fetch.preconnect }, + ); + vi.spyOn(globalThis, "fetch").mockImplementation(fetchImpl); + + const prepared = manager.prepareConfig(serverConfig).finally(() => { + preparedSettled = true; + }); + await fetchCalled.promise; + expect(capturedSignal).toBeDefined(); + + vi.advanceTimersByTime(9_999); + await drainMicrotasks(); + expect(capturedSignal!.aborted).toBe(false); + expect(preparedSettled).toBe(false); + expect(releaseSpy).not.toHaveBeenCalled(); + + vi.advanceTimersByTime(1); + await abortObserved.promise; + expect(capturedSignal!.aborted).toBe(true); + await drainMicrotasks(); + expect(preparedSettled).toBe(false); + expect(releaseSpy).not.toHaveBeenCalled(); + + allowFetchReject.resolve(); + const preparedConfig = await prepared; + + expect(preparedSettled).toBe(true); + expect(releaseSpy).toHaveBeenCalledTimes(1); + expect(getAuthorizationHeader(preparedConfig)).toBe(`Bearer ${STALE_ACCESS}`); + }); }); describe("MCPManager shared SQLite OAuth refresh", () => {