fix(auth): renewed oauth refresh leases
- Renewed durable OAuth refresh leases while token refreshes are in flight so slow endpoints cannot let a peer steal the row and replay a rotating refresh token. - Let CAS update and disable storage errors propagate instead of collapsing them into peer-win misses. - Added regressions for lease renewal and CAS storage failure propagation. Fixes #5081
This commit is contained in:
+103
-64
@@ -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<void>();
|
||||
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<string, number> = 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);
|
||||
|
||||
@@ -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",
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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<void>();
|
||||
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<void> {
|
||||
for (let attempt = 0; attempt < count; attempt++) {
|
||||
await Promise.resolve();
|
||||
}
|
||||
}
|
||||
|
||||
async function waitForControlledSleep(calls: ControlledSleep[], ms: number): Promise<ControlledSleep> {
|
||||
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<T>(
|
||||
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<void>();
|
||||
const allowRefreshResponse = Promise.withResolvers<void>();
|
||||
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, {
|
||||
|
||||
Reference in New Issue
Block a user