fix(auth): fenced oauth refresh writes

- Fenced final OAuth refresh update and terminal-disable CAS statements by row id, serialized credential data, active lease owner, and unexpired lease time.
- Passed an AbortSignal through MCP OAuth token refresh and bounded owned refresh operations below the lease TTL while awaiting the aborted fetch to settle.
- Added regressions for stolen-lease update/disable attempts and timed-out MCP token fetch abort behavior.

Fixes #5081
This commit is contained in:
roboomp
2026-07-10 21:49:18 +00:00
parent cb32853130
commit b60dc669ea
5 changed files with 272 additions and 21 deletions
+102 -17
View File
@@ -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<T extends OAuthCredential = OAuthCred
signal?: AbortSignal;
keepCredentialOnRefreshFailure?: boolean | ((error: unknown) => boolean);
onRefreshFailure?: (error: unknown) => void;
refresh: (credential: T) => Promise<OAuthCredentials>;
refreshTimeoutMs?: number;
refresh: (credential: T, signal?: AbortSignal) => Promise<OAuthCredentials>;
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 {
@@ -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",
});
});
});
+2 -1
View File
@@ -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) => {
@@ -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) {
@@ -115,10 +115,11 @@ async function withSharedSQLiteAuth<T>(
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<void>();
const abortObserved = Promise.withResolvers<void>();
const allowFetchReject = Promise.withResolvers<void>();
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<Response> => {
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", () => {