From b6e33f6ed3edcb7276533125fa0f6dfd88552ada Mon Sep 17 00:00:00 2001 From: can1357 Date: Sat, 4 Jul 2026 11:22:04 +0200 Subject: [PATCH] feat(ai): implemented persistent storage and synchronization for credential blocks - Implemented persistent storage for credential rate-limit blocks with automatic expiry and pruning. - Added broker API routes and client methods to manage, persist, and synchronize credential block states. - Integrated rate-limit checks into the credential selection logic, specifically refining Fable/Mythos tier exhaustion gating. - Extended schema versioning to include the new credential block table and verified persistence via comprehensive unit testing. --- packages/ai/CHANGELOG.md | 8 + packages/ai/src/auth-broker/client.ts | 53 ++- packages/ai/src/auth-broker/remote-store.ts | 198 ++++++++++- packages/ai/src/auth-broker/server.ts | 116 ++++++- packages/ai/src/auth-broker/types.ts | 23 +- packages/ai/src/auth-broker/wire-schemas.ts | 22 ++ packages/ai/src/auth-storage.ts | 314 ++++++++++++++++-- packages/ai/src/usage/claude.ts | 36 +- .../auth-storage-block-persistence.test.ts | 257 ++++++++++++++ ...auth-storage-claude-fable-fallback.test.ts | 159 ++++++++- packages/ai/test/remote-auth-store.test.ts | 108 ++++++ 11 files changed, 1224 insertions(+), 70 deletions(-) create mode 100644 packages/ai/test/auth-storage-block-persistence.test.ts diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index 9cbda6ba9..f455edca5 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -2,6 +2,14 @@ ## [Unreleased] +### Added + +- Persisted credential rate-limit blocks across processes: `auth_credential_blocks` (auth schema v5) stores per-credential blocks keyed by row id + provider key + block scope with MAX-upsert semantics, `AuthStorage` merges persisted and in-memory blocks on read, and auth-broker snapshots/SSE carry per-entry blocks with `POST /v1/credential/:id/block` and `DELETE /v1/credential/:id/blocks` endpoints so gateway and sibling omp processes stop re-discovering exhausted accounts by burning a 429 each. + +### Fixed + +- Fixed Anthropic credential selection sampling Fable/Mythos-exhausted accounts on every new session: a Fable/Mythos weekly cap now proactively hard-blocks the credential when confirmed exhausted (server `exhausted` status or used fraction >= 1) with a live `resetsAt`, and a live Fable 429 extends the reactive block to the confirmed tier reset instead of the 60s default. Unconfirmed rows (missing/expired reset, below cap) remain ranking hints only, preserving the false-100% guard. + ## [16.3.5] - 2026-07-04 ### Added diff --git a/packages/ai/src/auth-broker/client.ts b/packages/ai/src/auth-broker/client.ts index 6193d6a73..d873acfe6 100644 --- a/packages/ai/src/auth-broker/client.ts +++ b/packages/ai/src/auth-broker/client.ts @@ -9,6 +9,9 @@ import { readSseEvents } from "@oh-my-pi/pi-utils"; import { type } from "arktype"; import type { AuthCredential } from "../auth-storage"; import type { + CredentialBlockRequest, + CredentialBlockResponse, + CredentialBlocksDeleteResponse, CredentialDisableRequest, CredentialDisableResponse, CredentialRefreshResponse, @@ -20,6 +23,8 @@ import type { UsageResponse, } from "./types"; import { + credentialBlockResponseSchema, + credentialBlocksDeleteResponseSchema, credentialDisableResponseSchema, credentialRefreshResponseSchema, credentialUploadResponseSchema, @@ -106,7 +111,11 @@ export class AuthBrokerClient { } healthz(signal?: AbortSignal): Promise { - return this.#request("GET", "/v1/healthz", { schema: healthzResponseSchema, auth: false, signal }); + return this.#request("GET", "/v1/healthz", { + schema: healthzResponseSchema, + auth: false, + signal, + }); } async fetchSnapshot(opts: FetchSnapshotOptions = {}): Promise { @@ -230,19 +239,19 @@ export class AuthBrokerClient { // `metadata`) but leaves provider-specific extension fields permissive so // the broker can ship new shapes ahead of the client. `raw` is accepted // but normally stripped by the broker before send. - return this.#request("GET", "/v1/usage", { schema: usageResponseSchema, signal }) as Promise; + return this.#request("GET", "/v1/usage", { schema: usageResponseSchema, signal }); } async refreshCredential(id: number, signal?: AbortSignal): Promise { - return this.#request("POST", `/v1/credential/${id}/refresh`, { + return this.#request("POST", `/v1/credential/${id}/refresh`, { schema: credentialRefreshResponseSchema, signal, - }) as Promise; + }); } async disableCredential(id: number, cause: string, signal?: AbortSignal): Promise { const body: CredentialDisableRequest = { cause }; - return this.#request("POST", `/v1/credential/${id}/disable`, { + return this.#request("POST", `/v1/credential/${id}/disable`, { body, schema: credentialDisableResponseSchema, signal, @@ -255,18 +264,38 @@ export class AuthBrokerClient { signal?: AbortSignal, ): Promise { const body: CredentialUploadRequest = { provider, credential }; - return this.#request("POST", "/v1/credential", { + return this.#request("POST", "/v1/credential", { body, schema: credentialUploadResponseSchema, signal, - }) as Promise; + }); } - async #request( - method: "GET" | "POST", + async upsertCredentialBlock( + id: number, + block: CredentialBlockRequest, + signal?: AbortSignal, + ): Promise { + const body: CredentialBlockRequest = block; + return this.#request("POST", `/v1/credential/${id}/block`, { + body, + schema: credentialBlockResponseSchema, + signal, + }); + } + + async deleteCredentialBlocks(id: number, signal?: AbortSignal): Promise { + return this.#request("DELETE", `/v1/credential/${id}/blocks`, { + schema: credentialBlocksDeleteResponseSchema, + signal, + }); + } + + async #request( + method: "GET" | "POST" | "DELETE", path: string, opts: { schema: (input: unknown) => unknown; auth?: boolean; body?: unknown; signal?: AbortSignal }, - ): Promise { + ): Promise { const response = await this.#fetchRaw(method, path, opts); const text = await response.text(); const raw = this.#parseJson(text, response.status); @@ -277,7 +306,7 @@ export class AuthBrokerClient { body: validated.summary, }); } - return validated; + return validated as t; } #parseJson(text: string, status: number): unknown { @@ -293,7 +322,7 @@ export class AuthBrokerClient { } async #fetchRaw( - method: "GET" | "POST", + method: "GET" | "POST" | "DELETE", path: string, opts: { auth?: boolean; diff --git a/packages/ai/src/auth-broker/remote-store.ts b/packages/ai/src/auth-broker/remote-store.ts index bae50ac40..a5c89b108 100644 --- a/packages/ai/src/auth-broker/remote-store.ts +++ b/packages/ai/src/auth-broker/remote-store.ts @@ -16,13 +16,20 @@ import { type OAuthCredential, REMOTE_REFRESH_SENTINEL, type StoredAuthCredential, + type StoredCredentialBlock, } from "../auth-storage"; import * as AIError from "../error"; import type { OAuthCredentials } from "../registry/oauth/types"; import type { Provider } from "../types"; import type { UsageReport } from "../usage"; import { type AuthBrokerClient, AuthBrokerStreamUnsupportedError } from "./client"; -import type { RefresherSchedule, SnapshotEntry, SnapshotResponse, SnapshotStreamEvent } from "./types"; +import type { + CredentialBlockSnapshot, + RefresherSchedule, + SnapshotEntry, + SnapshotResponse, + SnapshotStreamEvent, +} from "./types"; /** * Client-side TTL for the aggregate `/v1/usage` response. The broker dedups @@ -38,6 +45,31 @@ const BACKGROUND_WAIT_MS = 30_000; const BACKGROUND_BACKOFF_INITIAL_MS = 500; const BACKGROUND_BACKOFF_MAX_MS = 30_000; +function compareCredentialBlockSnapshots(a: CredentialBlockSnapshot, b: CredentialBlockSnapshot): number { + const provider = a.providerKey.localeCompare(b.providerKey); + if (provider !== 0) return provider; + const scope = a.blockScope.localeCompare(b.blockScope); + if (scope !== 0) return scope; + return a.blockedUntilMs - b.blockedUntilMs; +} + +function toCredentialBlockSnapshot(block: StoredCredentialBlock): CredentialBlockSnapshot { + return { + providerKey: block.providerKey, + blockScope: block.blockScope, + blockedUntilMs: block.blockedUntilMs, + }; +} + +function credentialEntryWithBlocks( + entry: AuthCredentialSnapshotEntry, + blocks: readonly CredentialBlockSnapshot[] | undefined, +): SnapshotEntry { + const incoming: SnapshotEntry = { ...entry, rotatesInMs: null }; + if (blocks && blocks.length > 0) incoming.blocks = [...blocks].sort(compareCredentialBlockSnapshots); + return incoming; +} + function emptySnapshot(): SnapshotResponse { return { generation: 0, @@ -169,13 +201,17 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { } #applySnapshot(snapshot: SnapshotResponse, generation: number): void { - this.#snapshot = snapshot; + const nowMs = Date.now(); + this.#snapshot = { + ...snapshot, + credentials: snapshot.credentials.map(entry => this.#normalizeSnapshotEntryBlocks(entry, nowMs)), + }; this.#generation = generation; - this.#snapshotReceivedAt = Date.now(); + this.#snapshotReceivedAt = nowMs; const onSnapshot = this.#onSnapshot; if (!onSnapshot) return; try { - onSnapshot(snapshot, generation); + onSnapshot(this.#snapshot, generation); } catch (error) { logger.debug("auth-broker snapshot callback failed", { error: String(error) }); } @@ -266,11 +302,12 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { generation: number, serverNowMs: number, ): void { - const index = this.#snapshot.credentials.findIndex(candidate => candidate.id === entry.id); + const incoming = this.#normalizeSnapshotEntryBlocks(entry, Date.now()); + const index = this.#snapshot.credentials.findIndex(candidate => candidate.id === incoming.id); const credentials = index === -1 - ? [...this.#snapshot.credentials, entry] - : this.#snapshot.credentials.map((candidate, i) => (i === index ? entry : candidate)); + ? [...this.#snapshot.credentials, incoming] + : this.#snapshot.credentials.map((candidate, i) => (i === index ? incoming : candidate)); this.#snapshot = { ...this.#snapshot, generation, serverNowMs, refresher, credentials }; this.#generation = generation; this.#snapshotReceivedAt = Date.now(); @@ -304,6 +341,76 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { return out; } + getCredentialBlock(credentialId: number, providerKey: string, blockScope: string): number | undefined { + const nowMs = Date.now(); + this.cleanExpiredCredentialBlocks(nowMs); + const entry = this.#snapshot.credentials.find(candidate => candidate.id === credentialId); + if (!entry?.blocks) return undefined; + const block = entry.blocks.find( + candidate => candidate.providerKey === providerKey && candidate.blockScope === blockScope, + ); + if (!block || block.blockedUntilMs <= nowMs) return undefined; + return block.blockedUntilMs; + } + + listCredentialBlocks(credentialIds: readonly number[]): StoredCredentialBlock[] { + const nowMs = Date.now(); + this.cleanExpiredCredentialBlocks(nowMs); + const ids = new Set(credentialIds); + const blocks: StoredCredentialBlock[] = []; + for (const entry of this.#snapshot.credentials) { + if (!ids.has(entry.id) || !entry.blocks) continue; + for (const block of entry.blocks) { + if (block.blockedUntilMs <= nowMs) continue; + blocks.push({ + credentialId: entry.id, + providerKey: block.providerKey, + blockScope: block.blockScope, + blockedUntilMs: block.blockedUntilMs, + }); + } + } + blocks.sort((a, b) => a.credentialId - b.credentialId || compareCredentialBlockSnapshots(a, b)); + return blocks; + } + + upsertCredentialBlock(block: StoredCredentialBlock): void { + this.#upsertSnapshotBlock(block); + const body = toCredentialBlockSnapshot(block); + void this.#client + .upsertCredentialBlock(block.credentialId, body) + .then(() => { + this.#maybeRefreshSnapshot("credential block"); + }) + .catch(error => { + logger.warn("auth-broker credential block propagation failed", { + id: block.credentialId, + providerKey: block.providerKey, + blockScope: block.blockScope, + error: String(error), + }); + }); + } + + deleteCredentialBlocks(credentialId: number): void { + this.#deleteSnapshotBlocks(credentialId); + void this.#client + .deleteCredentialBlocks(credentialId) + .then(() => { + this.#maybeRefreshSnapshot("credential blocks delete"); + }) + .catch(error => { + logger.warn("auth-broker credential blocks delete propagation failed", { + id: credentialId, + error: String(error), + }); + }); + } + + cleanExpiredCredentialBlocks(nowMs: number): void { + this.#pruneExpiredCredentialBlocks(nowMs); + } + /** * In-memory update from a successful refresh through the broker. AuthStorage * calls this after `#replaceCredentialAt`; the broker already persisted the @@ -459,13 +566,19 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { // `entries` is the broker's authoritative post-upsert list of rows for // `provider`. Drop our existing rows for the same provider and splice in // the fresh set — preserving every other provider's rows in place. + const existingBlocks = new Map( + this.#snapshot.credentials + .filter(entry => entry.provider === provider && entry.blocks !== undefined) + .map(entry => [entry.id, entry.blocks] as const), + ); const others = this.#snapshot.credentials.filter(entry => entry.provider !== provider); - const incoming = entries.map(entry => ({ ...entry, rotatesInMs: null })); + const incoming = entries.map(entry => credentialEntryWithBlocks(entry, existingBlocks.get(entry.id))); this.#snapshot = { ...this.#snapshot, credentials: [...others, ...incoming] }; } #applyCredentialEntry(entry: AuthCredentialSnapshotEntry): void { - const incoming = { ...entry, rotatesInMs: null }; const index = this.#snapshot.credentials.findIndex(candidate => candidate.id === entry.id); + const existingBlocks = index === -1 ? undefined : this.#snapshot.credentials[index]?.blocks; + const incoming = credentialEntryWithBlocks(entry, existingBlocks); if (index === -1) { this.#snapshot = { ...this.#snapshot, credentials: [...this.#snapshot.credentials, incoming] }; return; @@ -485,6 +598,73 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { this.#snapshot = { ...this.#snapshot, credentials: next }; } + #normalizeSnapshotEntryBlocks(entry: SnapshotEntry, nowMs: number): SnapshotEntry { + if (!entry.blocks || entry.blocks.length === 0) return entry; + const blocks = entry.blocks + .filter(block => block.blockedUntilMs > nowMs) + .map(block => ({ + providerKey: block.providerKey, + blockScope: block.blockScope, + blockedUntilMs: block.blockedUntilMs, + })) + .sort(compareCredentialBlockSnapshots); + if (blocks.length > 0) return { ...entry, blocks }; + const next: SnapshotEntry = { ...entry }; + delete next.blocks; + return next; + } + + #upsertSnapshotBlock(block: StoredCredentialBlock): void { + const index = this.#snapshot.credentials.findIndex(entry => entry.id === block.credentialId); + if (index === -1) return; + const entry = this.#snapshot.credentials[index]!; + const incoming = toCredentialBlockSnapshot(block); + const blocks = entry.blocks ? [...entry.blocks] : []; + const blockIndex = blocks.findIndex( + candidate => candidate.providerKey === incoming.providerKey && candidate.blockScope === incoming.blockScope, + ); + if (blockIndex === -1) { + blocks.push(incoming); + } else { + const existing = blocks[blockIndex]!; + blocks[blockIndex] = { + ...existing, + blockedUntilMs: Math.max(existing.blockedUntilMs, incoming.blockedUntilMs), + }; + } + blocks.sort(compareCredentialBlockSnapshots); + const credentials = [...this.#snapshot.credentials]; + credentials[index] = { ...entry, blocks }; + this.#snapshot = { ...this.#snapshot, credentials }; + } + + #deleteSnapshotBlocks(credentialId: number): void { + const index = this.#snapshot.credentials.findIndex(entry => entry.id === credentialId); + if (index === -1) return; + const entry = this.#snapshot.credentials[index]!; + if (!entry.blocks || entry.blocks.length === 0) return; + const next: SnapshotEntry = { ...entry }; + delete next.blocks; + const credentials = [...this.#snapshot.credentials]; + credentials[index] = next; + this.#snapshot = { ...this.#snapshot, credentials }; + } + + #pruneExpiredCredentialBlocks(nowMs: number): void { + let changed = false; + const credentials = this.#snapshot.credentials.map(entry => { + if (!entry.blocks || entry.blocks.length === 0) return entry; + const blocks = entry.blocks.filter(block => block.blockedUntilMs > nowMs); + if (blocks.length === entry.blocks.length) return entry; + changed = true; + if (blocks.length > 0) return { ...entry, blocks }; + const next: SnapshotEntry = { ...entry }; + delete next.blocks; + return next; + }); + if (changed) this.#snapshot = { ...this.#snapshot, credentials }; + } + /** * Fire-and-forget `refreshSnapshot()` after a write. When the SSE stream is * active the broker will deliver the new generation push, so the extra GET diff --git a/packages/ai/src/auth-broker/server.ts b/packages/ai/src/auth-broker/server.ts index f74dd35e3..820ea8bd6 100644 --- a/packages/ai/src/auth-broker/server.ts +++ b/packages/ai/src/auth-broker/server.ts @@ -11,10 +11,13 @@ */ import { logger } from "@oh-my-pi/pi-utils"; import { type Type, type } from "arktype"; -import type { AuthStorage } from "../auth-storage"; +import type { AuthStorage, StoredCredentialBlock } from "../auth-storage"; import { parseBind } from "../utils/parse-bind"; import { AuthBrokerRefresher, type AuthBrokerRefresherSchedule } from "./refresher"; import type { + CredentialBlockResponse, + CredentialBlockSnapshot, + CredentialBlocksDeleteResponse, CredentialDisableResponse, CredentialRefreshResponse, CredentialUploadResponse, @@ -33,7 +36,11 @@ import { DEFAULT_SERVER_IDLE_TIMEOUT_S, DEFAULT_STREAM_KEEPALIVE_MS, } from "./types"; -import { credentialDisableRequestSchema, credentialUploadRequestSchema } from "./wire-schemas"; +import { + credentialBlockRequestSchema, + credentialDisableRequestSchema, + credentialUploadRequestSchema, +} from "./wire-schemas"; export interface AuthBrokerServerOptions { /** Underlying credential storage (wraps the local SQLite store on the broker). */ @@ -120,6 +127,8 @@ async function parseBody( const REFRESH_ROUTE = /^\/v1\/credential\/(\d+)\/refresh$/; const DISABLE_ROUTE = /^\/v1\/credential\/(\d+)\/disable$/; +const BLOCK_ROUTE = /^\/v1\/credential\/(\d+)\/block$/; +const BLOCKS_ROUTE = /^\/v1\/credential\/(\d+)\/blocks$/; const MAX_SNAPSHOT_WAIT_MS = 30_000; const DISABLED_NEXT_SWEEP_IN_MS = Number.MAX_SAFE_INTEGER; @@ -262,14 +271,48 @@ function computeRotatesInMs( return Math.max(0, rotatesAt - serverNowMs); } +function compareCredentialBlockSnapshots(a: CredentialBlockSnapshot, b: CredentialBlockSnapshot): number { + const provider = a.providerKey.localeCompare(b.providerKey); + if (provider !== 0) return provider; + const scope = a.blockScope.localeCompare(b.blockScope); + if (scope !== 0) return scope; + return a.blockedUntilMs - b.blockedUntilMs; +} + +function buildCredentialBlockGroups( + blocks: readonly StoredCredentialBlock[], + serverNowMs: number, +): Map { + const byCredentialId = new Map(); + for (const block of blocks) { + if (block.blockedUntilMs <= serverNowMs) continue; + const snapshotBlock: CredentialBlockSnapshot = { + providerKey: block.providerKey, + blockScope: block.blockScope, + blockedUntilMs: block.blockedUntilMs, + }; + const existing = byCredentialId.get(block.credentialId); + if (existing) { + existing.push(snapshotBlock); + } else { + byCredentialId.set(block.credentialId, [snapshotBlock]); + } + } + for (const credentialBlocks of byCredentialId.values()) credentialBlocks.sort(compareCredentialBlockSnapshots); + return byCredentialId; +} + function buildSnapshot(storage: AuthStorage, refresher: AuthBrokerRefresher | undefined): SnapshotResponse { const serverNowMs = Date.now(); const base = storage.exportSnapshot(); const { wire, nextSweepAt } = resolveRefresherSchedule(refresher, serverNowMs); - const credentials: SnapshotEntry[] = base.credentials.map(entry => ({ - ...entry, - rotatesInMs: computeRotatesInMs(entry, wire, nextSweepAt, serverNowMs), - })); + const credentialIds = base.credentials.map(entry => entry.id); + const blocksByCredentialId = buildCredentialBlockGroups(storage.listCredentialBlocks(credentialIds), serverNowMs); + const credentials: SnapshotEntry[] = base.credentials.map(entry => { + const blocks = blocksByCredentialId.get(entry.id); + const rotatesInMs = computeRotatesInMs(entry, wire, nextSweepAt, serverNowMs); + return blocks && blocks.length > 0 ? { ...entry, rotatesInMs, blocks } : { ...entry, rotatesInMs }; + }); return { generation: base.generation, generatedAt: base.generatedAt, @@ -336,7 +379,14 @@ async function serveSnapshot( * keep the stale projection. */ function fingerprintEntry(entry: SnapshotEntry): string { - return JSON.stringify([entry.id, entry.provider, entry.identityKey, entry.rotatesInMs, entry.credential]); + return JSON.stringify([ + entry.id, + entry.provider, + entry.identityKey, + entry.rotatesInMs, + entry.credential, + entry.blocks ?? [], + ]); } function sseEvent(event: string, body: unknown): string { @@ -593,6 +643,58 @@ export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServer const response: CredentialDisableResponse = { ok: true }; return json(200, response); } + const blockMatch = req.method === "POST" ? pathname.match(BLOCK_ROUTE) : null; + if (blockMatch) { + const id = Number.parseInt(blockMatch[1], 10); + const parsed = await parseBody(req, credentialBlockRequestSchema); + if (!parsed.ok) return parsed.response; + const block: StoredCredentialBlock = { + credentialId: id, + providerKey: parsed.data.providerKey, + blockScope: parsed.data.blockScope, + blockedUntilMs: parsed.data.blockedUntilMs, + }; + if (!opts.storage.exportSnapshot().credentials.some(entry => entry.id === id)) { + logger.info("auth-broker credential block miss", { id, peer }); + return json(404, { error: `No credential with id=${id}` }); + } + try { + opts.storage.upsertCredentialBlock(block); + const response: CredentialBlockResponse = { ok: true }; + logger.info("auth-broker credential block upserted", { + id, + peer, + providerKey: block.providerKey, + blockScope: block.blockScope, + blockedUntilMs: block.blockedUntilMs, + }); + return json(200, response); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-broker credential block upsert failed", { id, peer, error: message }); + const status = message.includes("No credential with id") ? 404 : 500; + return json(status, { error: message }); + } + } + const blocksDeleteMatch = req.method === "DELETE" ? pathname.match(BLOCKS_ROUTE) : null; + if (blocksDeleteMatch) { + const id = Number.parseInt(blocksDeleteMatch[1], 10); + if (!opts.storage.exportSnapshot().credentials.some(entry => entry.id === id)) { + logger.info("auth-broker credential blocks delete miss", { id, peer }); + return json(404, { error: `No credential with id=${id}` }); + } + try { + opts.storage.deleteCredentialBlocks(id); + const response: CredentialBlocksDeleteResponse = { ok: true }; + logger.info("auth-broker credential blocks deleted", { id, peer }); + return json(200, response); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-broker credential blocks delete failed", { id, peer, error: message }); + const status = message.includes("No credential with id") ? 404 : 500; + return json(status, { error: message }); + } + } if (req.method === "POST" && pathname === "/v1/credential") { const parsed = await parseBody(req, credentialUploadRequestSchema); if (!parsed.ok) return parsed.response; diff --git a/packages/ai/src/auth-broker/types.ts b/packages/ai/src/auth-broker/types.ts index 1387b6473..6a329d2ca 100644 --- a/packages/ai/src/auth-broker/types.ts +++ b/packages/ai/src/auth-broker/types.ts @@ -6,7 +6,12 @@ * credential expires or a 401 surfaces on a supposedly-fresh credential. */ -import type { AuthCredential, AuthCredentialSnapshot, AuthCredentialSnapshotEntry } from "../auth-storage"; +import type { + AuthCredential, + AuthCredentialSnapshot, + AuthCredentialSnapshotEntry, + StoredCredentialBlock, +} from "../auth-storage"; import type { UsageReport } from "../usage"; /** GET /v1/healthz response body. */ @@ -22,8 +27,11 @@ export interface RefresherSchedule { nextSweepInMs: number; } +export type CredentialBlockSnapshot = Omit; + export type SnapshotEntry = AuthCredentialSnapshotEntry & { rotatesInMs: number | null; + blocks?: CredentialBlockSnapshot[]; }; /** GET /v1/snapshot response body. */ @@ -54,6 +62,19 @@ export interface CredentialDisableResponse { ok: boolean; } +/** POST /v1/credential/:id/block request body. */ +export type CredentialBlockRequest = CredentialBlockSnapshot; + +/** POST /v1/credential/:id/block response body. */ +export interface CredentialBlockResponse { + ok: boolean; +} + +/** DELETE /v1/credential/:id/blocks response body. */ +export interface CredentialBlocksDeleteResponse { + ok: boolean; +} + /** * POST /v1/credential request body. The OAuth `refresh` must be the *real* * refresh token (not the sentinel) — the broker is the canonical writer. diff --git a/packages/ai/src/auth-broker/wire-schemas.ts b/packages/ai/src/auth-broker/wire-schemas.ts index 05b66499d..741a9ea11 100644 --- a/packages/ai/src/auth-broker/wire-schemas.ts +++ b/packages/ai/src/auth-broker/wire-schemas.ts @@ -69,6 +69,13 @@ export const credentialSnapshotEntrySchema = type({ identityKey: "string | null", }); +export const credentialBlockSnapshotSchema = type({ + "+": "reject", + providerKey: type("string").atLeastLength(1), + blockScope: "string", + blockedUntilMs: "number", +}); + export const snapshotEntrySchema = type({ "+": "reject", id: "number.integer", @@ -76,6 +83,7 @@ export const snapshotEntrySchema = type({ credential: snapshotCredentialSchema, identityKey: "string | null", rotatesInMs: "number | null", + "blocks?": credentialBlockSnapshotSchema.array(), }); export const refresherScheduleSchema = type({ @@ -235,6 +243,20 @@ export const credentialDisableResponseSchema = type({ ok: "boolean", }); +// ─── Credential blocks ────────────────────────────────────────────────────── + +export const credentialBlockRequestSchema = credentialBlockSnapshotSchema; + +export const credentialBlockResponseSchema = type({ + "+": "reject", + ok: "boolean", +}); + +export const credentialBlocksDeleteResponseSchema = type({ + "+": "reject", + ok: "boolean", +}); + // ─── Upload ──────────────────────────────────────────────────────────────── export const credentialUploadRequestSchema = type({ diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index fc71c9f70..a790ce5d5 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -123,6 +123,18 @@ export interface StoredAuthCredential { disabledCause: string | null; } +/** One persisted rate-limit block: credential row id + provider-type key + optional scope. */ +export interface StoredCredentialBlock { + /** SQLite row id of the credential (auth_credentials.id). */ + credentialId: number; + /** `${provider}:${credentialType}` — same value as AuthStorage's in-memory providerKey. */ + providerKey: string; + /** Block scope (e.g. "tier:fable"); empty string = unscoped. Never NUL-delimited. */ + blockScope: string; + /** Epoch milliseconds. */ + blockedUntilMs: number; +} + /** * Per-credential health record returned by {@link AuthStorage.checkCredentials}. * @@ -305,6 +317,16 @@ export interface AuthCredentialStore { getCache(key: string, options?: { includeExpired?: boolean }): string | null; setCache(key: string, value: string, expiresAtSec: number): void; cleanExpiredCache(): void; + /** Non-expired block for one (credential, providerKey, scope) key, or undefined. */ + getCredentialBlock?(credentialId: number, providerKey: string, blockScope: string): number | undefined; + /** Upsert with MAX semantics: keep the later blockedUntilMs on conflict. */ + upsertCredentialBlock?(block: StoredCredentialBlock): void; + /** Drop every block row for a credential (all providerKeys/scopes). */ + deleteCredentialBlocks?(credentialId: number): void; + /** Prune rows with blocked_until_ms <= nowMs. */ + cleanExpiredCredentialBlocks?(nowMs: number): void; + /** List non-expired blocks for broker snapshots. */ + listCredentialBlocks?(credentialIds: readonly number[]): StoredCredentialBlock[]; /** * Append usage-limit snapshots for trend history. Optional: stores without * durable storage (e.g. the broker remote store) omit it and recording is @@ -987,6 +1009,11 @@ export class AuthStorage { } catch { // Best-effort. } + try { + this.#store.cleanExpiredCredentialBlocks?.(Date.now()); + } catch { + // Best-effort. + } this.#usageFetch = options.usageFetch ?? fetch; this.#usageRequestTimeoutMs = options.usageRequestTimeoutMs ?? DEFAULT_USAGE_REQUEST_TIMEOUT_MS; this.#refreshOAuthCredentialOverride = options.refreshOAuthCredential; @@ -1305,13 +1332,13 @@ export class AuthStorage { return blockScope ? `${providerKey}\0${blockScope}` : providerKey; } - /** Returns block expiry timestamp for a credential/key pair, cleaning up expired entries. */ - #getCredentialBlockedUntilForKey(backoffKey: string, credentialIndex: number): number | undefined { + /** Returns in-memory block expiry timestamp for a credential/key pair, cleaning up expired entries. */ + #getCredentialBlockedUntilForKey(backoffKey: string, credentialIndex: number, nowMs: number): number | undefined { const backoffMap = this.#credentialBackoff.get(backoffKey); if (!backoffMap) return undefined; const blockedUntil = backoffMap.get(credentialIndex); if (!blockedUntil) return undefined; - if (blockedUntil <= Date.now()) { + if (blockedUntil <= nowMs) { backoffMap.delete(credentialIndex); if (backoffMap.size === 0) { this.#credentialBackoff.delete(backoffKey); @@ -1321,28 +1348,80 @@ export class AuthStorage { return blockedUntil; } - /** Returns block expiry timestamp for a credential, checking global then scoped blocks. */ + #readPersistedCredentialBlock( + credentialId: number, + providerKey: string, + blockScope: string | undefined, + ): number | undefined { + const getCredentialBlock = this.#store.getCredentialBlock?.bind(this.#store); + if (!getCredentialBlock) return undefined; + try { + return getCredentialBlock(credentialId, providerKey, blockScope ?? ""); + } catch (err) { + logger.debug("Failed to read credential block from persistent store", { + err, + credentialId, + providerKey, + blockScope, + }); + return undefined; + } + } + + /** Returns block expiry timestamp for a credential, checking unscoped and scoped blocks. */ #getCredentialBlockedUntil( + provider: string, providerKey: string, credentialIndex: number, blockScope: string | undefined = undefined, ): number | undefined { - const globalBlockedUntil = this.#getCredentialBlockedUntilForKey(providerKey, credentialIndex); - if (globalBlockedUntil !== undefined || !blockScope) return globalBlockedUntil; - return this.#getCredentialBlockedUntilForKey(this.#toScopedBackoffKey(providerKey, blockScope), credentialIndex); + const nowMs = Date.now(); + let blockedUntil = this.#getCredentialBlockedUntilForKey(providerKey, credentialIndex, nowMs); + if (blockScope) { + const scopedBlockedUntil = this.#getCredentialBlockedUntilForKey( + this.#toScopedBackoffKey(providerKey, blockScope), + credentialIndex, + nowMs, + ); + if (scopedBlockedUntil !== undefined && (blockedUntil === undefined || scopedBlockedUntil > blockedUntil)) { + blockedUntil = scopedBlockedUntil; + } + } + + const credentialId = this.#getStoredCredentials(provider)[credentialIndex]?.id; + if (credentialId === undefined) return blockedUntil; + const persistedGlobalBlockedUntil = this.#readPersistedCredentialBlock(credentialId, providerKey, ""); + if ( + persistedGlobalBlockedUntil !== undefined && + (blockedUntil === undefined || persistedGlobalBlockedUntil > blockedUntil) + ) { + blockedUntil = persistedGlobalBlockedUntil; + } + if (blockScope) { + const persistedScopedBlockedUntil = this.#readPersistedCredentialBlock(credentialId, providerKey, blockScope); + if ( + persistedScopedBlockedUntil !== undefined && + (blockedUntil === undefined || persistedScopedBlockedUntil > blockedUntil) + ) { + blockedUntil = persistedScopedBlockedUntil; + } + } + return blockedUntil; } /** Checks if a credential is temporarily blocked due to usage limits. */ #isCredentialBlocked( + provider: string, providerKey: string, credentialIndex: number, blockScope: string | undefined = undefined, ): boolean { - return this.#getCredentialBlockedUntil(providerKey, credentialIndex, blockScope) !== undefined; + return this.#getCredentialBlockedUntil(provider, providerKey, credentialIndex, blockScope) !== undefined; } /** Marks a credential as blocked until the specified time. */ #markCredentialBlocked( + provider: string, providerKey: string, credentialIndex: number, blockedUntilMs: number, @@ -1351,8 +1430,31 @@ export class AuthStorage { const backoffKey = this.#toScopedBackoffKey(providerKey, blockScope); const backoffMap = this.#credentialBackoff.get(backoffKey) ?? new Map(); const existing = backoffMap.get(credentialIndex) ?? 0; - backoffMap.set(credentialIndex, Math.max(existing, blockedUntilMs)); + const nextBlockedUntil = Math.max(existing, blockedUntilMs); + backoffMap.set(credentialIndex, nextBlockedUntil); this.#credentialBackoff.set(backoffKey, backoffMap); + + const upsertCredentialBlock = this.#store.upsertCredentialBlock?.bind(this.#store); + if (!upsertCredentialBlock) return; + const credentialId = this.#getStoredCredentials(provider)[credentialIndex]?.id; + if (credentialId === undefined) return; + try { + upsertCredentialBlock({ + credentialId, + providerKey, + blockScope: blockScope ?? "", + blockedUntilMs: nextBlockedUntil, + }); + } catch (err) { + logger.debug("Failed to persist credential block", { + err, + credentialId, + provider, + providerKey, + blockScope, + blockedUntilMs: nextBlockedUntil, + }); + } } /** Records which credential was used for a session (for rate-limit switching). */ @@ -1469,7 +1571,7 @@ export class AuthStorage { for (const idx of order) { const candidate = credentials[idx]; - if (!this.#isCredentialBlocked(providerKey, candidate.index)) { + if (!this.#isCredentialBlocked(provider, providerKey, candidate.index)) { return candidate; } } @@ -2983,7 +3085,7 @@ export class AuthStorage { } } - this.#markCredentialBlocked(providerKey, sessionCredential.index, blockedUntil, blockScope); + this.#markCredentialBlocked(provider, providerKey, sessionCredential.index, blockedUntil, blockScope); const remainingCredentials = this.#getCredentialsForProvider(provider) .map((credential, index) => ({ credential, index })) @@ -2994,7 +3096,12 @@ export class AuthStorage { let retryAtMs: number | undefined; for (const candidate of remainingCredentials) { - const candidateBlockedUntil = this.#getCredentialBlockedUntil(providerKey, candidate.index, blockScope); + const candidateBlockedUntil = this.#getCredentialBlockedUntil( + provider, + providerKey, + candidate.index, + blockScope, + ); if (candidateBlockedUntil === undefined) return { switched: true }; if (retryAtMs === undefined || candidateBlockedUntil < retryAtMs) retryAtMs = candidateBlockedUntil; } @@ -3174,7 +3281,12 @@ export class AuthStorage { args.order.map(async idx => { const selection = args.credentials[idx]; if (!selection) return null; - const blockedUntil = this.#getCredentialBlockedUntil(args.providerKey, selection.index, args.blockScope); + const blockedUntil = this.#getCredentialBlockedUntil( + args.provider, + args.providerKey, + selection.index, + args.blockScope, + ); if (blockedUntil !== undefined) return { selection, usage: null, usageChecked: false, blockedUntil }; const usage = await this.#getUsageReport(args.provider, selection.credential, { ...args.options, @@ -3211,7 +3323,13 @@ export class AuthStorage { if (!blocked && scopedLimits && this.#isUsageLimitReached(scopedLimits)) { const resetAtMs = this.#getUsageResetAtMs(scopedLimits, nowMs); blockedUntil = resetAtMs ?? Date.now() + AuthStorage.#defaultBackoffMs; - this.#markCredentialBlocked(args.providerKey, selection.index, blockedUntil, args.blockScope); + this.#markCredentialBlocked( + args.provider, + args.providerKey, + selection.index, + blockedUntil, + args.blockScope, + ); blocked = true; } const windows = usage ? strategy.findWindowLimits(usage, args.rankingContext) : undefined; @@ -3275,7 +3393,7 @@ export class AuthStorage { // with the most headroom proactively and fall back intelligently when rate-limited. const sessionPreferredIsAvailable = sessionPreferredIndex !== undefined && - !this.#isCredentialBlocked(providerKey, sessionPreferredIndex, blockScope); + !this.#isCredentialBlocked(provider, providerKey, sessionPreferredIndex, blockScope); const shouldRank = checkUsage && (!sessionPreferredIsAvailable || requiresProModel); const rankingOrder = shouldRank && sessionId ? credentials.map((_credential, index) => index) : order; const candidates = shouldRank @@ -3298,7 +3416,7 @@ export class AuthStorage { if (sessionPreferredIndex !== undefined && !requiresProModel) { const sessionPreferredCandidate = candidates.findIndex( candidate => - !this.#isCredentialBlocked(providerKey, candidate.selection.index, blockScope) && + !this.#isCredentialBlocked(provider, providerKey, candidate.selection.index, blockScope) && candidate.selection.index === sessionPreferredIndex, ); if (sessionPreferredCandidate > 0) { @@ -3403,7 +3521,7 @@ export class AuthStorage { if (resolved) return resolved; } - if (fallback && this.#isCredentialBlocked(providerKey, fallback.selection.index, blockScope)) { + if (fallback && this.#isCredentialBlocked(provider, providerKey, fallback.selection.index, blockScope)) { return this.#tryOAuthCredential(provider, fallback.selection, providerKey, sessionId, options, { checkUsage, allowBlocked: true, @@ -3566,7 +3684,7 @@ export class AuthStorage { blockScope, allowFallback = true, } = usageOptions; - if (!allowBlocked && this.#isCredentialBlocked(providerKey, selection.index, blockScope)) { + if (!allowBlocked && this.#isCredentialBlocked(provider, providerKey, selection.index, blockScope)) { return undefined; } @@ -3603,6 +3721,7 @@ export class AuthStorage { if (this.#isUsageLimitReached(scopedLimits)) { const resetAtMs = this.#getUsageResetAtMs(scopedLimits, Date.now()); this.#markCredentialBlocked( + provider, providerKey, selection.index, resetAtMs ?? Date.now() + AuthStorage.#defaultBackoffMs, @@ -3679,6 +3798,7 @@ export class AuthStorage { if (this.#isUsageLimitReached(scopedLimits)) { const resetAtMs = this.#getUsageResetAtMs(scopedLimits, Date.now()); this.#markCredentialBlocked( + provider, providerKey, selection.index, resetAtMs ?? Date.now() + AuthStorage.#defaultBackoffMs, @@ -3755,7 +3875,7 @@ export class AuthStorage { } } else { // Block temporarily for transient failures (5 minutes) - this.#markCredentialBlocked(providerKey, selection.index, Date.now() + 5 * 60 * 1000); + this.#markCredentialBlocked(provider, providerKey, selection.index, Date.now() + 5 * 60 * 1000); } } @@ -4154,6 +4274,15 @@ export class AuthStorage { * that `markUsageLimitReached` set for the now-obsolete reset time. */ #clearCredentialBlocks(provider: string, credentialId: number): void { + const deleteCredentialBlocks = this.#store.deleteCredentialBlocks?.bind(this.#store); + if (deleteCredentialBlocks) { + try { + deleteCredentialBlocks(credentialId); + } catch (err) { + logger.debug("Failed to clear persisted credential blocks", { err, provider, credentialId }); + } + } + const index = this.#getStoredCredentials(provider).findIndex(entry => entry.id === credentialId); if (index < 0) return; const providerKey = this.#getProviderTypeKey(provider, "oauth"); @@ -4213,6 +4342,7 @@ export class AuthStorage { this.#clearSessionCredential(provider, sessionId); this.#markCredentialBlocked( + provider, this.#getProviderTypeKey(provider, matched.type), matched.index, Date.now() + AuthStorage.#defaultBackoffMs, @@ -4275,11 +4405,16 @@ export class AuthStorage { (credential, index) => credential.type === sessionCredential.type && index !== sessionCredential.index && - !this.#isCredentialBlocked(providerKey, index), + !this.#isCredentialBlocked(provider, providerKey, index), ); const target = this.#getStoredCredentials(provider)[sessionCredential.index]; this.#clearSessionCredential(provider, sessionId); - this.#markCredentialBlocked(providerKey, sessionCredential.index, Date.now() + AuthStorage.#defaultBackoffMs); + this.#markCredentialBlocked( + provider, + providerKey, + sessionCredential.index, + Date.now() + AuthStorage.#defaultBackoffMs, + ); if (target) { const markSuspect = this.#store.markCredentialSuspect?.bind(this.#store); @@ -4505,6 +4640,33 @@ export class AuthStorage { }); } + /** + * Broker-server seam: list non-expired persisted blocks for snapshot entries. + */ + listCredentialBlocks(credentialIds: readonly number[]): StoredCredentialBlock[] { + return this.#store.listCredentialBlocks?.(credentialIds) ?? []; + } + + /** + * Broker-server seam: persist one credential block and notify snapshot waiters. + */ + upsertCredentialBlock(block: StoredCredentialBlock): void { + const upsertCredentialBlock = this.#store.upsertCredentialBlock?.bind(this.#store); + if (!upsertCredentialBlock) return; + upsertCredentialBlock(block); + this.#bumpGeneration("credential-block"); + } + + /** + * Broker-server seam: clear all persisted blocks for one credential and notify snapshot waiters. + */ + deleteCredentialBlocks(credentialId: number): void { + const deleteCredentialBlocks = this.#store.deleteCredentialBlocks?.bind(this.#store); + if (!deleteCredentialBlocks) return; + deleteCredentialBlocks(credentialId); + this.#bumpGeneration("credential-block"); + } + /** * Describe where the active credential for a provider came from. * @@ -4570,13 +4732,20 @@ type AuthRow = { identity_key: string | null; }; +type CredentialBlockRow = { + credential_id: number; + provider_key: string; + block_scope: string; + blocked_until_ms: number; +}; + type SerializedCredentialRecord = { credentialType: AuthCredential["type"]; data: string; identityKey: string | null; }; -const AUTH_SCHEMA_VERSION = 4; +const AUTH_SCHEMA_VERSION = 5; const SQLITE_NOW_EPOCH = "CAST(strftime('%s','now') AS INTEGER)"; /** @@ -4776,6 +4945,11 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { #getCacheIncludingExpiredStmt: Statement; #upsertCacheStmt: Statement; #deleteExpiredCacheStmt: Statement; + #getCredentialBlockStmt: Statement; + #listCredentialBlocksByCredentialStmt: Statement; + #upsertCredentialBlockStmt: Statement; + #deleteCredentialBlocksStmt: Statement; + #deleteExpiredCredentialBlocksStmt: Statement; #insertUsageHistoryStmt: Statement; #insertUsageCostStmt: Statement; #listUsageCostsStmt: Statement; @@ -4821,6 +4995,23 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { "INSERT INTO cache (key, value, expires_at) VALUES (?, ?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value, expires_at = excluded.expires_at", ); this.#deleteExpiredCacheStmt = this.#db.prepare(`DELETE FROM cache WHERE expires_at <= ${SQLITE_NOW_EPOCH}`); + this.#getCredentialBlockStmt = this.#db.prepare( + "SELECT blocked_until_ms FROM auth_credential_blocks WHERE credential_id = ? AND provider_key = ? AND block_scope = ? AND blocked_until_ms > ?", + ); + this.#listCredentialBlocksByCredentialStmt = this.#db.prepare( + "SELECT credential_id, provider_key, block_scope, blocked_until_ms FROM auth_credential_blocks WHERE credential_id = ? AND blocked_until_ms > ? ORDER BY provider_key ASC, block_scope ASC", + ); + this.#upsertCredentialBlockStmt = this.#db.prepare( + `INSERT INTO auth_credential_blocks (credential_id, provider_key, block_scope, blocked_until_ms, updated_at) + VALUES (?, ?, ?, ?, ${SQLITE_NOW_EPOCH}) + ON CONFLICT(credential_id, provider_key, block_scope) DO UPDATE SET + blocked_until_ms = MAX(blocked_until_ms, excluded.blocked_until_ms), + updated_at = excluded.updated_at`, + ); + this.#deleteCredentialBlocksStmt = this.#db.prepare("DELETE FROM auth_credential_blocks WHERE credential_id = ?"); + this.#deleteExpiredCredentialBlocksStmt = this.#db.prepare( + "DELETE FROM auth_credential_blocks WHERE blocked_until_ms <= ?", + ); this.#insertUsageHistoryStmt = this.#db.prepare( "INSERT INTO usage_history (recorded_at, provider, account_key, email, account_id, limit_id, label, window_label, used_fraction, status, resets_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", ); @@ -4932,6 +5123,7 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { if (!this.#authCredentialsTableExists()) { this.#createAuthCredentialsTable(); + this.#createAuthCredentialBlocksTable(); this.#writeAuthSchemaVersion(AUTH_SCHEMA_VERSION); return; } @@ -4948,6 +5140,7 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { } this.#createAuthCredentialIndexes(); + this.#createAuthCredentialBlocksTable(); this.#backfillCredentialIdentityKeys(); // Rewriting an already-current version row is a no-op write transaction // on every boot; only persist when the recorded version actually changes. @@ -5031,6 +5224,20 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { `); } + #createAuthCredentialBlocksTable(): void { + this.#db.run(` + CREATE TABLE IF NOT EXISTS auth_credential_blocks ( + credential_id INTEGER NOT NULL, + provider_key TEXT NOT NULL, + block_scope TEXT NOT NULL DEFAULT '', + blocked_until_ms INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (credential_id, provider_key, block_scope) + ); + CREATE INDEX IF NOT EXISTS idx_auth_credential_blocks_expires ON auth_credential_blocks(blocked_until_ms); + `); + } + #migrateAuthSchema(fromVersion: number): void { if (fromVersion < 1) { this.#migrateAuthSchemaV0ToV1(); @@ -5041,6 +5248,9 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { if (fromVersion < 4) { this.#migrateAuthSchemaV3ToV4(); } + if (fromVersion < 5) { + this.#migrateAuthSchemaV4ToV5(); + } } #migrateAuthSchemaV0ToV1(): void { @@ -5127,6 +5337,13 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { migrate(); } + #migrateAuthSchemaV4ToV5(): void { + const migrate = this.#db.transaction(() => { + this.#createAuthCredentialBlocksTable(); + }); + migrate(); + } + #backfillCredentialIdentityKeys(): void { const selectRowsStmt = this.#db.prepare( "SELECT id, provider, credential_type, data, disabled_cause, identity_key FROM auth_credentials WHERE identity_key IS NULL ORDER BY id ASC", @@ -5392,6 +5609,54 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { } } + getCredentialBlock(credentialId: number, providerKey: string, blockScope: string): number | undefined { + const nowMs = Date.now(); + this.#deleteExpiredCredentialBlocksStmt.run(nowMs); + const row = this.#getCredentialBlockStmt.get(credentialId, providerKey, blockScope, nowMs) as + | { blocked_until_ms?: number } + | undefined; + return typeof row?.blocked_until_ms === "number" ? row.blocked_until_ms : undefined; + } + + upsertCredentialBlock(block: StoredCredentialBlock): void { + this.#upsertCredentialBlockStmt.run( + block.credentialId, + block.providerKey, + block.blockScope, + block.blockedUntilMs, + ); + } + + deleteCredentialBlocks(credentialId: number): void { + this.#deleteCredentialBlocksStmt.run(credentialId); + } + + cleanExpiredCredentialBlocks(nowMs: number): void { + this.#deleteExpiredCredentialBlocksStmt.run(nowMs); + } + + listCredentialBlocks(credentialIds: readonly number[]): StoredCredentialBlock[] { + if (credentialIds.length === 0) return []; + const nowMs = Date.now(); + this.cleanExpiredCredentialBlocks(nowMs); + const seenCredentialIds = new Set(); + const blocks: StoredCredentialBlock[] = []; + for (const credentialId of credentialIds) { + if (seenCredentialIds.has(credentialId)) continue; + seenCredentialIds.add(credentialId); + const rows = this.#listCredentialBlocksByCredentialStmt.all(credentialId, nowMs) as CredentialBlockRow[]; + for (const row of rows) { + blocks.push({ + credentialId: row.credential_id, + providerKey: row.provider_key, + blockScope: row.block_scope, + blockedUntilMs: row.blocked_until_ms, + }); + } + } + return blocks; + } + recordUsageSnapshots(entries: UsageHistoryEntry[]): void { try { for (const entry of entries) { @@ -5585,6 +5850,11 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { this.#getCacheIncludingExpiredStmt.finalize(); this.#upsertCacheStmt.finalize(); this.#deleteExpiredCacheStmt.finalize(); + this.#getCredentialBlockStmt.finalize(); + this.#listCredentialBlocksByCredentialStmt.finalize(); + this.#upsertCredentialBlockStmt.finalize(); + this.#deleteCredentialBlocksStmt.finalize(); + this.#deleteExpiredCredentialBlocksStmt.finalize(); this.#insertUsageHistoryStmt.finalize(); this.#lastUsageHistoryStmt.finalize(); this.#listUsageHistoryStmt.finalize(); diff --git a/packages/ai/src/usage/claude.ts b/packages/ai/src/usage/claude.ts index d5ca399d0..c346115ee 100644 --- a/packages/ai/src/usage/claude.ts +++ b/packages/ai/src/usage/claude.ts @@ -615,20 +615,40 @@ function scopeClaudeLimitsForModel(report: UsageReport, context: CredentialRanki } /** - * Exclude Fable and Mythos tier weekly caps from proactive hard-blocking - * (gating) because they are notoriously unreliable (they report 100% exhausted - * while the account can still serve requests). They stay available for ranking - * pressure in findWindowLimits via scopeClaudeLimitsForModel. + * A Fable/Mythos weekly row is trusted for gating only at full exhaustion + * (server `exhausted` status or used fraction >= 1) with a live reset + * timestamp. Anything below that stays untrusted: the counters are + * notoriously unreliable short of the cap (they report high utilization + * while the account can still serve requests). + */ +function isConfirmedExhaustedTierRow(limit: UsageLimit, nowMs: number): boolean { + const resetsAt = limit.window?.resetsAt; + if (typeof resetsAt !== "number" || !Number.isFinite(resetsAt) || resetsAt <= nowMs) return false; + if (limit.status === "exhausted") return true; + const fraction = resolveUsedFraction(limit); + return typeof fraction === "number" && fraction >= 1; +} + +/** + * Scope limits for proactive hard-blocking (gating). Fable and Mythos tier + * weekly caps participate only when {@link isConfirmedExhaustedTierRow} + * confirms them, so a confirmed-dead account is skipped up front and a + * reactive 429 block extends to the tier reset in markUsageLimitReached, + * while unconfirmed rows remain ranking pressure only via + * scopeClaudeLimitsForModel. */ function scopeClaudeLimitsForModelHardBlock( report: UsageReport, context: CredentialRankingContext | undefined, ): UsageLimit[] { const kind = getClaudeModelKind(context); - const excludeHardBlock = kind === "fable" || kind === "mythos"; - return report.limits.filter( - limit => limit.scope.shared === true || (kind !== undefined && limit.scope.tier === kind && !excludeHardBlock), - ); + const requireConfirmedTierRow = kind === "fable" || kind === "mythos"; + const nowMs = Date.now(); + return report.limits.filter(limit => { + if (limit.scope.shared === true) return true; + if (kind === undefined || limit.scope.tier !== kind) return false; + return !requireConfirmedTierRow || isConfirmedExhaustedTierRow(limit, nowMs); + }); } function rankingUsedFraction(limit: UsageLimit): number { diff --git a/packages/ai/test/auth-storage-block-persistence.test.ts b/packages/ai/test/auth-storage-block-persistence.test.ts new file mode 100644 index 000000000..5c460373b --- /dev/null +++ b/packages/ai/test/auth-storage-block-persistence.test.ts @@ -0,0 +1,257 @@ +import { Database } from "bun:sqlite"; +import { afterEach, beforeEach, describe, expect, it } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { AuthStorage, type OAuthCredential, SqliteAuthCredentialStore } from "@oh-my-pi/pi-ai"; +import { removeWithRetries } from "../../utils/src/temp"; + +const PROVIDER = "anthropic"; +const PROVIDER_KEY = "anthropic:oauth"; +const FUTURE_BLOCK_MS = 1_899_999_999_000; +const EXPIRED_BLOCK_MS = 1; +const LEGACY_TIMESTAMP = 1_700_000_000; + +function oauthCredential(suffix: string): OAuthCredential { + return { + type: "oauth", + access: `access-${suffix}`, + refresh: `refresh-${suffix}`, + expires: Date.now() + 3_600_000, + accountId: `account-${suffix}`, + email: `${suffix}@example.com`, + }; +} + +function readAuthSchemaVersion(dbPath: string): number | null { + const db = new Database(dbPath, { readonly: true }); + try { + const row = db.prepare("SELECT version FROM auth_schema_version WHERE id = 1").get() as + | { version?: number } + | undefined; + return typeof row?.version === "number" ? row.version : null; + } finally { + db.close(); + } +} + +function tableExists(dbPath: string, tableName: string): boolean { + const db = new Database(dbPath, { readonly: true }); + try { + const row = db + .prepare("SELECT 1 AS present FROM sqlite_master WHERE type = 'table' AND name = ?") + .get(tableName) as { present?: number } | undefined; + return row?.present === 1; + } finally { + db.close(); + } +} + +describe("AuthStorage credential block persistence", () => { + let tempDir = ""; + let dbPath = ""; + + beforeEach(async () => { + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-blocks-")); + dbPath = path.join(tempDir, "agent.db"); + }); + + afterEach(async () => { + dbPath = ""; + if (tempDir) { + await removeWithRetries(tempDir); + tempDir = ""; + } + }); + + it("honors scoped and unscoped blocks written by a previous AuthStorage instance", async () => { + const firstStore = await SqliteAuthCredentialStore.open(dbPath); + firstStore.saveOAuth(PROVIDER, oauthCredential("1")); + firstStore.saveOAuth(PROVIDER, oauthCredential("2")); + firstStore.saveOAuth(PROVIDER, oauthCredential("3")); + const rows = firstStore.listAuthCredentials(PROVIDER); + const firstStorage = new AuthStorage(firstStore); + await firstStorage.reload(); + try { + firstStorage.upsertCredentialBlock({ + credentialId: rows[0]!.id, + providerKey: PROVIDER_KEY, + blockScope: "tier:fable", + blockedUntilMs: FUTURE_BLOCK_MS, + }); + firstStorage.upsertCredentialBlock({ + credentialId: rows[1]!.id, + providerKey: PROVIDER_KEY, + blockScope: "", + blockedUntilMs: FUTURE_BLOCK_MS, + }); + } finally { + firstStorage.close(); + } + + const reopenedStore = await SqliteAuthCredentialStore.open(dbPath); + const reopenedStorage = new AuthStorage(reopenedStore); + await reopenedStorage.reload(); + try { + const fableKey = await reopenedStorage.getApiKey(PROVIDER, "session-3", { modelId: "claude-fable-5" }); + expect(fableKey).toBe("access-3"); + } finally { + reopenedStorage.close(); + } + }); + + it("keeps the later expiry when a shorter block is upserted for the same key", async () => { + const store = await SqliteAuthCredentialStore.open(dbPath); + store.saveOAuth(PROVIDER, oauthCredential("1")); + const [row] = store.listAuthCredentials(PROVIDER); + if (!row) throw new Error("expected credential row"); + const storage = new AuthStorage(store); + await storage.reload(); + try { + const longerBlock = FUTURE_BLOCK_MS + 60_000; + storage.upsertCredentialBlock({ + credentialId: row.id, + providerKey: PROVIDER_KEY, + blockScope: "tier:fable", + blockedUntilMs: longerBlock, + }); + storage.upsertCredentialBlock({ + credentialId: row.id, + providerKey: PROVIDER_KEY, + blockScope: "tier:fable", + blockedUntilMs: FUTURE_BLOCK_MS, + }); + + expect(storage.listCredentialBlocks([row.id])).toEqual([ + { credentialId: row.id, providerKey: PROVIDER_KEY, blockScope: "tier:fable", blockedUntilMs: longerBlock }, + ]); + } finally { + storage.close(); + } + }); + + it("drops expired rows from reads and clears persisted blocks through the public delete wrapper", async () => { + const store = await SqliteAuthCredentialStore.open(dbPath); + store.saveOAuth(PROVIDER, oauthCredential("1")); + const [row] = store.listAuthCredentials(PROVIDER); + if (!row) throw new Error("expected credential row"); + const storage = new AuthStorage(store); + await storage.reload(); + try { + storage.upsertCredentialBlock({ + credentialId: row.id, + providerKey: PROVIDER_KEY, + blockScope: "tier:fable", + blockedUntilMs: FUTURE_BLOCK_MS, + }); + storage.upsertCredentialBlock({ + credentialId: row.id, + providerKey: PROVIDER_KEY, + blockScope: "", + blockedUntilMs: EXPIRED_BLOCK_MS, + }); + + expect(storage.listCredentialBlocks([row.id])).toEqual([ + { + credentialId: row.id, + providerKey: PROVIDER_KEY, + blockScope: "tier:fable", + blockedUntilMs: FUTURE_BLOCK_MS, + }, + ]); + + const generationBeforeDelete = storage.getGeneration(); + storage.deleteCredentialBlocks(row.id); + expect(storage.listCredentialBlocks([row.id])).toEqual([]); + expect(storage.getGeneration()).toBe(generationBeforeDelete + 1); + } finally { + storage.close(); + } + }); + + it("keeps a block attached to the same credential row after a sibling is disabled", async () => { + const store = await SqliteAuthCredentialStore.open(dbPath); + store.saveOAuth(PROVIDER, oauthCredential("1")); + store.saveOAuth(PROVIDER, oauthCredential("2")); + store.saveOAuth(PROVIDER, oauthCredential("3")); + const rows = store.listAuthCredentials(PROVIDER); + const storage = new AuthStorage(store); + await storage.reload(); + try { + storage.upsertCredentialBlock({ + credentialId: rows[1]!.id, + providerKey: PROVIDER_KEY, + blockScope: "", + blockedUntilMs: FUTURE_BLOCK_MS, + }); + } finally { + storage.close(); + } + + const disablingStore = await SqliteAuthCredentialStore.open(dbPath); + disablingStore.deleteAuthCredential(rows[0]!.id, "disabled for test"); + disablingStore.close(); + + const reopenedStore = await SqliteAuthCredentialStore.open(dbPath); + const reopenedStorage = new AuthStorage(reopenedStore); + await reopenedStorage.reload(); + try { + const key = await reopenedStorage.getApiKey(PROVIDER, "a"); + expect(key).toBe("access-3"); + } finally { + reopenedStorage.close(); + } + }); + + it("migrates a v4 auth database to v5 without dropping credential rows", async () => { + const legacyDb = new Database(dbPath); + legacyDb.run(` + CREATE TABLE auth_schema_version ( + id INTEGER PRIMARY KEY CHECK (id = 1), + version INTEGER NOT NULL + ); + INSERT INTO auth_schema_version(id, version) VALUES (1, 4); + CREATE TABLE auth_credentials ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider TEXT NOT NULL, + credential_type TEXT NOT NULL, + data TEXT NOT NULL, + disabled_cause TEXT DEFAULT NULL, + identity_key TEXT DEFAULT NULL, + created_at INTEGER NOT NULL DEFAULT (CAST(strftime('%s','now') AS INTEGER)), + updated_at INTEGER NOT NULL DEFAULT (CAST(strftime('%s','now') AS INTEGER)) + ); + `); + legacyDb + .prepare( + "INSERT INTO auth_credentials (provider, credential_type, data, disabled_cause, identity_key, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?)", + ) + .run( + PROVIDER, + "oauth", + JSON.stringify({ + access: "legacy-access", + refresh: "legacy-refresh", + expires: Date.now() + 3_600_000, + accountId: "legacy-account", + email: "legacy@example.com", + }), + null, + "email:legacy@example.com", + LEGACY_TIMESTAMP, + LEGACY_TIMESTAMP, + ); + legacyDb.close(); + + const migratedStore = await SqliteAuthCredentialStore.open(dbPath); + try { + const rows = migratedStore.listAuthCredentials(PROVIDER); + expect(rows).toHaveLength(1); + expect(rows[0]!.credential).toMatchObject({ type: "oauth", access: "legacy-access" }); + expect(readAuthSchemaVersion(dbPath)).toBe(5); + expect(tableExists(dbPath, "auth_credential_blocks")).toBe(true); + } finally { + migratedStore.close(); + } + }); +}); diff --git a/packages/ai/test/auth-storage-claude-fable-fallback.test.ts b/packages/ai/test/auth-storage-claude-fable-fallback.test.ts index 927c99269..a94395807 100644 --- a/packages/ai/test/auth-storage-claude-fable-fallback.test.ts +++ b/packages/ai/test/auth-storage-claude-fable-fallback.test.ts @@ -5,7 +5,7 @@ import { AuthStorage, type StoredAuthCredential, } from "@oh-my-pi/pi-ai/auth-storage"; -import type { UsageReport } from "@oh-my-pi/pi-ai/usage"; +import type { UsageLimit, UsageReport } from "@oh-my-pi/pi-ai/usage"; import * as claudeUsage from "@oh-my-pi/pi-ai/usage/claude"; interface ObservableStore extends AuthCredentialStore { @@ -83,7 +83,11 @@ function baseReport(email: string): UsageReport { }; } -function withFable(report: UsageReport, usedFraction: number): UsageReport { +function withFable( + report: UsageReport, + usedFraction: number, + options: { resetsAt?: number; status?: UsageLimit["status"] } = {}, +): UsageReport { return { ...report, limits: [ @@ -92,9 +96,13 @@ function withFable(report: UsageReport, usedFraction: number): UsageReport { id: "anthropic:7d:fable", label: "Claude 7 Day (Fable)", scope: { provider: "anthropic", windowId: "7d", tier: "fable" }, - window: { id: "7d", label: "7 Day" }, + window: { + id: "7d", + label: "7 Day", + ...(options.resetsAt === undefined ? {} : { resetsAt: options.resetsAt }), + }, amount: { used: usedFraction * 100, limit: 100, usedFraction, unit: "percent" }, - status: usedFraction >= 1 ? "exhausted" : "ok", + status: options.status ?? (usedFraction >= 1 ? "exhausted" : "ok"), }, ], }; @@ -133,10 +141,10 @@ describe("AuthStorage Claude Fable tier fallback", () => { }); it("does not block OAuth credentials just because the Fable tier is not reported", async () => { - // All three credentials lack a Fable-specific bucket. Per the user's - // intent, unknown headroom is not treated as exhausted; the selector - // still picks the first credential in hashed order and lets the live - // request decide if the account can serve Fable. + // All three credentials lack a Fable-specific bucket. Unknown headroom is + // not treated as exhausted; the selector still picks the first credential + // in hashed order and lets the live request decide if the account can + // serve Fable. const reportsByAccess: Record = { "oat-1": baseReport("a@example.com"), "oat-2": baseReport("b@example.com"), @@ -149,13 +157,108 @@ describe("AuthStorage Claude Fable tier fallback", () => { return reportsByAccess[access] ?? null; }); - // With Fable tier excluded from proactive hard blocks, it should still select the first available key. + // Unknown Fable headroom is not a proactive hard block. const key = await storage.getApiKey("anthropic", "session-3", { modelId: "claude-fable-5" }); expect(key).toBe("oat-1"); }); - it("uses explicit exhausted Fable tier rows as ranking hints instead of hard blockers", async () => { + it("skips a Fable credential only when the exhausted tier row has a future reset", async () => { + const now = Date.now(); + const reportsByAccess: Record = { + "oat-1": withFable(baseReport("a@example.com"), 1.0, { resetsAt: now + 3_600_000 }), + "oat-2": withFable(baseReport("b@example.com"), 0.4, { resetsAt: now + 3_600_000 }), + "oat-3": withFable(baseReport("c@example.com"), 0.4, { resetsAt: now + 3_600_000 }), + }; + + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + const access = params.credential.type === "oauth" ? params.credential.accessToken : undefined; + if (!access) return null; + return reportsByAccess[access] ?? null; + }); + + const key = await storage.getApiKey("anthropic", "session-3", { modelId: "claude-fable-5" }); + + expect(key).toBe("oat-2"); + }); + + it("keeps a full Fable tier row eligible when the reset timestamp is missing", async () => { + const now = Date.now(); + const reportsByAccess: Record = { + "oat-1": withFable(baseReport("a@example.com"), 1.0), + "oat-2": withFable(baseReport("b@example.com"), 1.0, { resetsAt: now + 3_600_000 }), + "oat-3": withFable(baseReport("c@example.com"), 1.0, { resetsAt: now + 3_600_000 }), + }; + + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + const access = params.credential.type === "oauth" ? params.credential.accessToken : undefined; + if (!access) return null; + return reportsByAccess[access] ?? null; + }); + + const key = await storage.getApiKey("anthropic", "session-3", { modelId: "claude-fable-5" }); + + expect(key).toBe("oat-1"); + }); + + it("keeps a full Fable tier row eligible when the reset timestamp is already elapsed", async () => { + const now = Date.now(); + const reportsByAccess: Record = { + "oat-1": withFable(baseReport("a@example.com"), 1.0, { resetsAt: now - 1 }), + "oat-2": withFable(baseReport("b@example.com"), 1.0, { resetsAt: now + 3_600_000 }), + "oat-3": withFable(baseReport("c@example.com"), 1.0, { resetsAt: now + 3_600_000 }), + }; + + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + const access = params.credential.type === "oauth" ? params.credential.accessToken : undefined; + if (!access) return null; + return reportsByAccess[access] ?? null; + }); + + const key = await storage.getApiKey("anthropic", "session-3", { modelId: "claude-fable-5" }); + + expect(key).toBe("oat-1"); + }); + + it("keeps a near-cap Fable tier row eligible even with a future reset timestamp", async () => { + const now = Date.now(); + const reportsByAccess: Record = { + "oat-1": withFable(baseReport("a@example.com"), 0.97, { resetsAt: now + 3_600_000 }), + "oat-2": withFable(baseReport("b@example.com"), 1.0, { resetsAt: now + 3_600_000 }), + "oat-3": withFable(baseReport("c@example.com"), 1.0, { resetsAt: now + 3_600_000 }), + }; + + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + const access = params.credential.type === "oauth" ? params.credential.accessToken : undefined; + if (!access) return null; + return reportsByAccess[access] ?? null; + }); + + const key = await storage.getApiKey("anthropic", "session-3", { modelId: "claude-fable-5" }); + + expect(key).toBe("oat-1"); + }); + + it("treats exhausted status plus a future reset as a confirmed Fable hard block", async () => { + const now = Date.now(); + const reportsByAccess: Record = { + "oat-1": withFable(baseReport("a@example.com"), 0.97, { resetsAt: now + 3_600_000, status: "exhausted" }), + "oat-2": withFable(baseReport("b@example.com"), 0.4, { resetsAt: now + 3_600_000 }), + "oat-3": withFable(baseReport("c@example.com"), 0.4, { resetsAt: now + 3_600_000 }), + }; + + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + const access = params.credential.type === "oauth" ? params.credential.accessToken : undefined; + if (!access) return null; + return reportsByAccess[access] ?? null; + }); + + const key = await storage.getApiKey("anthropic", "session-3", { modelId: "claude-fable-5" }); + + expect(key).toBe("oat-2"); + }); + + it("uses unconfirmed exhausted Fable tier rows as ranking hints instead of hard blockers", async () => { const reportsByAccess: Record = { "oat-1": withFable(baseReport("a@example.com"), 1.0), "oat-2": withFable(baseReport("b@example.com"), 1.0), @@ -173,7 +276,7 @@ describe("AuthStorage Claude Fable tier fallback", () => { expect(key).toBe("oat-3"); }); - it("rotates after a live Fable 429 even when sibling Fable tier rows are exhausted", async () => { + it("rotates after a live Fable 429 when sibling Fable tier rows are unconfirmed", async () => { const reportsByAccess: Record = { "oat-1": withFable(baseReport("a@example.com"), 1.0), "oat-2": withFable(baseReport("b@example.com"), 1.0), @@ -197,6 +300,40 @@ describe("AuthStorage Claude Fable tier fallback", () => { expect(["oat-2", "oat-3"]).toContain(retryKey as string); }); + it("extends a live Fable rate-limit block to the confirmed Fable reset", async () => { + const startNow = Date.now(); + let now = startNow; + vi.spyOn(Date, "now").mockImplementation(() => now); + const fableReset = startNow + 10 * 60_000; + const reportsByAccess: Record = { + "oat-1": withFable(baseReport("a@example.com"), 1.0, { resetsAt: fableReset }), + "oat-2": withFable(baseReport("b@example.com"), 0.2, { resetsAt: fableReset }), + "oat-3": withFable(baseReport("c@example.com"), 0.2, { resetsAt: fableReset }), + }; + + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + const access = params.credential.type === "oauth" ? params.credential.accessToken : undefined; + if (!access) return null; + return reportsByAccess[access] ?? null; + }); + + const firstKey = await storage.getApiKey("anthropic", "session-3"); + expect(firstKey).toBe("oat-1"); + + const result = await storage.markUsageLimitReached("anthropic", "session-3", { + modelId: "claude-fable-5", + retryAfterMs: 1_000, + }); + expect(result.switched).toBe(true); + + reportsByAccess["oat-1"] = withFable(baseReport("a@example.com"), 0.2, { resetsAt: fableReset }); + store.cache.clear(); + now = startNow + 60_001; + + const retryKey = await storage.getApiKey("anthropic", "session-3", { modelId: "claude-fable-5" }); + expect(retryKey).toBe("oat-2"); + }); + it("still blocks OAuth credentials with exhausted shared Anthropic limits", async () => { const reportsByAccess: Record = { "oat-1": withSharedUsage(withFable(baseReport("a@example.com"), 0.1), "7d", 1.0), diff --git a/packages/ai/test/remote-auth-store.test.ts b/packages/ai/test/remote-auth-store.test.ts index 1f424a678..876b34344 100644 --- a/packages/ai/test/remote-auth-store.test.ts +++ b/packages/ai/test/remote-auth-store.test.ts @@ -6,11 +6,15 @@ import { AuthStorage, REMOTE_REFRESH_SENTINEL, SqliteAuthCredentialStore } from import { AuthBrokerClient, type AuthBrokerServerHandle, + type CredentialBlockResponse, + type FetchSnapshotResult, RemoteAuthCredentialStore, startAuthBroker, } from "@oh-my-pi/pi-ai/auth-broker"; +import { snapshotResponseSchema } from "@oh-my-pi/pi-ai/auth-broker/wire-schemas"; import * as oauthUtils from "@oh-my-pi/pi-ai/registry/oauth"; import type { UsageLimit, UsageReport } from "@oh-my-pi/pi-ai/usage"; +import { type } from "arktype"; import { removeWithRetries } from "../../utils/src/temp"; function requireLimit(report: UsageReport, id: string): UsageLimit { @@ -267,6 +271,110 @@ describe("RemoteAuthCredentialStore + AuthStorage integration", () => { remoteStore.close(); }); + test("snapshot wire schema accepts entries with and without credential blocks", () => { + const futureBlock = Date.now() + 60_000; + const validated = snapshotResponseSchema({ + generation: 1, + generatedAt: Date.now(), + serverNowMs: Date.now(), + refresher: { enabled: false, intervalMs: 0, skewMs: 0, nextSweepInMs: Number.MAX_SAFE_INTEGER }, + credentials: [ + { + id: 1, + provider: "anthropic", + credential: { + type: "oauth", + access: "access-without-blocks", + refresh: REMOTE_REFRESH_SENTINEL, + expires: futureBlock, + accountId: "account-without-blocks", + email: "without-blocks@example.com", + }, + identityKey: "email:without-blocks@example.com", + rotatesInMs: null, + }, + { + id: 2, + provider: "anthropic", + credential: { + type: "oauth", + access: "access-with-blocks", + refresh: REMOTE_REFRESH_SENTINEL, + expires: futureBlock, + accountId: "account-with-blocks", + email: "with-blocks@example.com", + }, + identityKey: "email:with-blocks@example.com", + rotatesInMs: null, + blocks: [{ providerKey: "anthropic:oauth", blockScope: "tier:fable", blockedUntilMs: futureBlock }], + }, + ], + }); + + expect(validated).not.toBeInstanceOf(type.errors); + if (validated instanceof type.errors) throw new Error("expected valid snapshot"); + expect(validated.credentials[0]!.blocks).toBeUndefined(); + expect(validated.credentials[1]!.blocks).toEqual([ + { providerKey: "anthropic:oauth", blockScope: "tier:fable", blockedUntilMs: futureBlock }, + ]); + }); + + test("RemoteAuthCredentialStore reads snapshot blocks and applies upserts before broker acknowledgement", () => { + const futureBlock = Date.now() + 60_000; + const laterBlock = futureBlock + 60_000; + const brokerClient = new AuthBrokerClient({ url: "http://127.0.0.1:9", token: "unused" }); + const fetchSnapshotPending = Promise.withResolvers(); + vi.spyOn(brokerClient, "fetchSnapshot").mockReturnValue(fetchSnapshotPending.promise); + const upsertPending = Promise.withResolvers(); + const upsertSpy = vi.spyOn(brokerClient, "upsertCredentialBlock").mockReturnValue(upsertPending.promise); + const remoteStore = new RemoteAuthCredentialStore({ + client: brokerClient, + streamSnapshots: false, + initialSnapshot: { + generation: 1, + generatedAt: Date.now(), + serverNowMs: Date.now(), + refresher: { enabled: false, intervalMs: 0, skewMs: 0, nextSweepInMs: Number.MAX_SAFE_INTEGER }, + credentials: [ + { + id: 7, + provider: "anthropic", + credential: { + type: "oauth", + access: "remote-access", + refresh: REMOTE_REFRESH_SENTINEL, + expires: futureBlock, + accountId: "remote-account", + email: "remote@example.com", + }, + identityKey: "email:remote@example.com", + rotatesInMs: null, + blocks: [{ providerKey: "anthropic:oauth", blockScope: "tier:fable", blockedUntilMs: futureBlock }], + }, + ], + }, + }); + try { + expect(remoteStore.getCredentialBlock(7, "anthropic:oauth", "tier:fable")).toBe(futureBlock); + + remoteStore.upsertCredentialBlock({ + credentialId: 7, + providerKey: "anthropic:oauth", + blockScope: "tier:fable", + blockedUntilMs: laterBlock, + }); + + expect(remoteStore.getCredentialBlock(7, "anthropic:oauth", "tier:fable")).toBe(laterBlock); + expect(upsertSpy).toHaveBeenCalledWith(7, { + providerKey: "anthropic:oauth", + blockScope: "tier:fable", + blockedUntilMs: laterBlock, + }); + } finally { + remoteStore.close(); + } + }); + test("ingestUsageReport overlays only the matching Anthropic report and getUsageReport returns the overlaid Fable row", async () => { const brokerClient = new AuthBrokerClient({ url: handle!.url, token }); const remoteStore = new RemoteAuthCredentialStore({