diff --git a/packages/ai/src/auth-broker/remote-store.ts b/packages/ai/src/auth-broker/remote-store.ts index 25a0a73f7..04e67dffa 100644 --- a/packages/ai/src/auth-broker/remote-store.ts +++ b/packages/ai/src/auth-broker/remote-store.ts @@ -191,8 +191,14 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { } async markCredentialSuspect(credentialId: number, opts: { signal?: AbortSignal } = {}): Promise { - await this.#client.refreshCredential(credentialId, opts.signal); - await this.waitForFreshSnapshot(MAX_WAIT_MS, opts); + const { entry } = await this.#client.refreshCredential(credentialId, opts.signal); + if (entry.credential.type !== "oauth") { + throw new Error(`Broker returned non-OAuth credential for id=${credentialId}`); + } + this.#applyCredentialEntry(entry); + void this.refreshSnapshot().catch(error => { + logger.debug("auth-broker snapshot refresh after suspect credential refresh failed", { error: String(error) }); + }); } replaceAuthCredentialsForProvider(_provider: string, _credentials: AuthCredential[]): StoredAuthCredential[] { @@ -295,6 +301,17 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { const incoming = entries.map(entry => ({ ...entry, rotatesInMs: null })); 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); + if (index === -1) { + this.#snapshot = { ...this.#snapshot, credentials: [...this.#snapshot.credentials, incoming] }; + return; + } + const credentials = [...this.#snapshot.credentials]; + credentials[index] = incoming; + this.#snapshot = { ...this.#snapshot, credentials }; + } #removeProviderEntries(provider: string): void { const next = this.#snapshot.credentials.filter(entry => entry.provider !== provider); diff --git a/packages/ai/test/remote-auth-store.test.ts b/packages/ai/test/remote-auth-store.test.ts index 471631cf4..7c9e6c13f 100644 --- a/packages/ai/test/remote-auth-store.test.ts +++ b/packages/ai/test/remote-auth-store.test.ts @@ -105,6 +105,39 @@ describe("RemoteAuthCredentialStore + AuthStorage integration", () => { expect(refreshSpy).toHaveBeenCalledTimes(1); clientStorage.close(); }); + test("suspect credential refresh updates the client snapshot from the broker response", async () => { + const rotated = { + access: "server-access-after-401", + refresh: "server-refresh-after-401", + expires: Date.now() + 120_000, + accountId: "account-1", + email: "a@example.com", + }; + const refreshSpy = vi.spyOn(oauthUtils, "refreshOAuthToken").mockResolvedValue(rotated); + + const brokerClient = new AuthBrokerClient({ url: handle!.url, token }); + const initialResult = await brokerClient.fetchSnapshot(); + if (initialResult.status !== 200) throw new Error("expected snapshot"); + const initialEntry = initialResult.snapshot.credentials[0]; + if (!initialEntry) throw new Error("expected credential"); + + const remoteStore = new RemoteAuthCredentialStore({ + client: brokerClient, + initialSnapshot: initialResult.snapshot, + }); + + await remoteStore.markCredentialSuspect(initialEntry.id); + const rows = remoteStore.listAuthCredentials("anthropic"); + + expect(rows).toHaveLength(1); + expect(rows[0]?.credential.type).toBe("oauth"); + if (rows[0]?.credential.type === "oauth") { + expect(rows[0].credential.access).toBe("server-access-after-401"); + expect(rows[0].credential.refresh).toBe(REMOTE_REFRESH_SENTINEL); + } + expect(refreshSpy).toHaveBeenCalledTimes(1); + remoteStore.close(); + }); test("RemoteAuthCredentialStore rejects writes from the client", () => { const remoteStore = new RemoteAuthCredentialStore({