diff --git a/src/platform/credentials.ts b/src/platform/credentials.ts index 926278e..265f295 100644 --- a/src/platform/credentials.ts +++ b/src/platform/credentials.ts @@ -1,4 +1,9 @@ -import type { Credential, CredentialInfo, CredentialStore } from "@earendil-works/pi-ai"; +import type { + AuthOperationOptions, + Credential, + CredentialInfo, + CredentialStore, +} from "@earendil-works/pi-ai"; const SECRET_PREFIX = "provider:"; @@ -65,21 +70,33 @@ export class PortableCredentialStore implements CredentialStore { modify( providerId: string, fn: (current: Credential | undefined) => Promise, + options?: AuthOperationOptions, ): Promise { - return this.#enqueue(providerId, async () => { - const current = await this.read(providerId); - const next = await fn(current); - if (next === undefined) return current; - await this.#write(providerId, next); - return next; - }); + const signal = options?.signal; + return this.#enqueue( + providerId, + async () => { + const current = await this.read(providerId); + signal?.throwIfAborted(); + // Once fn starts, finish persistence: a refresh may already have rotated tokens. + const next = await fn(current); + if (next === undefined) return current; + await this.#write(providerId, next); + return next; + }, + signal, + ); } - delete(providerId: string): Promise { - return this.#enqueue(providerId, async () => { - this.#memory.delete(providerId); - if (this.#ctx) await this.#ctx.setSecret(`${SECRET_PREFIX}${providerId}`, ""); - }); + delete(providerId: string, options?: AuthOperationOptions): Promise { + return this.#enqueue( + providerId, + async () => { + this.#memory.delete(providerId); + if (this.#ctx) await this.#ctx.setSecret(`${SECRET_PREFIX}${providerId}`, ""); + }, + options?.signal, + ); } async setApiKey(providerId: string, key: string): Promise { @@ -98,13 +115,26 @@ export class PortableCredentialStore implements CredentialStore { } } - #enqueue(providerId: string, task: () => Promise): Promise { + #enqueue(providerId: string, task: () => Promise, signal?: AbortSignal): Promise { + if (signal?.aborted) return Promise.reject(signal.reason); const previous = this.#chains.get(providerId) ?? Promise.resolve(); - const run = previous.then(task, task); + let onAbort: (() => void) | undefined; + const start = () => { + if (onAbort) signal?.removeEventListener("abort", onAbort); + signal?.throwIfAborted(); + return task(); + }; + const run = previous.then(start, start); + // Caller cancellation must not release the serialized queue. this.#chains.set( providerId, run.catch(() => undefined), ); - return run; + if (!signal) return run; + return new Promise((resolve, reject) => { + onAbort = () => reject(signal.reason); + signal.addEventListener("abort", onAbort, { once: true }); + void run.then(resolve, reject); + }); } } diff --git a/tests/credentials.test.ts b/tests/credentials.test.ts new file mode 100644 index 0000000..931ff5d --- /dev/null +++ b/tests/credentials.test.ts @@ -0,0 +1,411 @@ +import { + createModels, + createProvider, + type Credential, + type CredentialStore, + type OAuthCredential, +} from "@earendil-works/pi-ai"; +import { afterEach, expect, onTestFinished, test, vi } from "vitest"; +import { PortableCredentialStore } from "../src/platform/credentials.ts"; + +afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllGlobals(); +}); + +test.each([ + { name: "saves a queued login after the earlier mutation finishes", cancel: false }, + { name: "does not save a queued login after it is cancelled", cancel: true }, +])("$name", async ({ cancel }) => { + const providerId = "synthetic-provider"; + const retained: OAuthCredential = { + type: "oauth", + access: "synthetic-access-a", + refresh: "synthetic-refresh-a", + expires: Number.MAX_SAFE_INTEGER, + }; + const incoming: OAuthCredential = { + type: "oauth", + access: "synthetic-access-b", + refresh: "synthetic-refresh-b", + expires: Number.MAX_SAFE_INTEGER, + }; + const credentials: CredentialStore = new PortableCredentialStore(null); + const modify = credentials.modify.bind(credentials); + const activeStarted = deferred(); + const releaseActive = deferred(); + const loginQueued = deferred(); + const abort = new AbortController(); + const network = vi.fn(() => { + throw new Error("This test must not make network requests."); + }); + vi.stubGlobal("fetch", network); + + const loginProvider = vi.fn(async () => incoming); + const models = createModels({ + credentials, + authContext: { + async env() { + return undefined; + }, + async fileExists() { + return false; + }, + }, + }); + models.setProvider( + createProvider({ + id: providerId, + name: "Synthetic provider", + models: [], + auth: { + oauth: { + name: "Synthetic OAuth", + login: loginProvider, + async refresh(credential) { + return credential; + }, + async toAuth(credential) { + return { apiKey: credential.access }; + }, + }, + }, + api: { + stream() { + throw new Error("This test must not start inference."); + }, + streamSimple() { + throw new Error("This test must not start inference."); + }, + }, + }), + ); + + const active = modify(providerId, async () => { + activeStarted.resolve(); + await releaseActive.promise; + return retained; + }); + + onTestFinished(async () => { + releaseActive.resolve(); + await active; + await modify(providerId, async () => undefined); + }); + + await activeStarted.promise; + const modifySpy = vi.spyOn(credentials, "modify").mockImplementation((id, update, options) => { + const pending = modify(id, update, options); + loginQueued.resolve(options?.signal); + return pending; + }); + const outcome = models + .login(providerId, "oauth", { + signal: abort.signal, + async prompt() { + throw new Error("The synthetic login must not prompt."); + }, + notify() {}, + }) + .then( + (credential) => ({ status: "fulfilled", credential }), + (reason: unknown) => ({ status: "rejected", reason }), + ); + + const queuedSignal = await Promise.race([ + loginQueued.promise, + outcome.then((result) => { + throw new Error("Login settled before its credential mutation was queued.", { + cause: result, + }); + }), + ]); + expect(loginProvider).toHaveBeenCalledOnce(); + expect(modifySpy).toHaveBeenCalledOnce(); + expect(queuedSignal?.aborted).toBe(false); + expect(await credentials.read(providerId)).toBeUndefined(); + + if (cancel) { + abort.abort(); + expect(queuedSignal?.aborted).toBe(true); + expect(await outcome).toMatchObject({ + status: "rejected", + reason: { name: "AbortError" }, + }); + } + + releaseActive.resolve(); + await active; + await modify(providerId, async () => undefined); + + if (!cancel) { + expect(await outcome).toEqual({ status: "fulfilled", credential: incoming }); + } + expect(network).not.toHaveBeenCalled(); + expect(await credentials.read(providerId)).toEqual(cancel ? retained : incoming); +}); + +test.each(["modify", "delete"] as const)("does not start a pre-aborted %s", async (operation) => { + const { store, values, getSecret, setSecret } = secretStore(); + const retained = syntheticCredential("retained"); + values.set("provider:synthetic-provider", JSON.stringify(retained)); + const update = vi.fn(async () => syntheticCredential("cancelled")); + const abort = new AbortController(); + const reason = new Error("Synthetic cancellation"); + abort.abort(reason); + + const pending = + operation === "modify" + ? store.modify("synthetic-provider", update, { signal: abort.signal }) + : store.delete("synthetic-provider", { signal: abort.signal }); + + await expect(pending).rejects.toBe(reason); + expect(update).not.toHaveBeenCalled(); + expect(getSecret).not.toHaveBeenCalled(); + expect(setSecret).not.toHaveBeenCalled(); + expect(values.get("provider:synthetic-provider")).toBe(JSON.stringify(retained)); +}); + +test.each([ + { operation: "modify", preAborted: false }, + { operation: "delete", preAborted: false }, + { operation: "modify", preAborted: true }, + { operation: "delete", preAborted: true }, +])( + "a cancelled queued $operation rejects before its predecessor settles (pre-aborted: $preAborted)", + async ({ operation, preAborted }) => { + const { store, values, getSecret, setSecret } = secretStore(); + const retained = syntheticCredential("retained"); + const entered = deferred(); + const release = deferred(); + const abort = new AbortController(); + const reason = new Error("Synthetic queued cancellation"); + const active = store.modify("synthetic-provider", async () => { + entered.resolve(); + await release.promise; + return retained; + }); + onTestFinished(async () => { + release.resolve(); + await active; + await store.modify("synthetic-provider", async () => undefined); + }); + await entered.promise; + getSecret.mockClear(); + + if (preAborted) abort.abort(reason); + const update = vi.fn(async () => syntheticCredential("cancelled")); + const cancelled = + operation === "modify" + ? store.modify("synthetic-provider", update, { signal: abort.signal }) + : store.delete("synthetic-provider", { signal: abort.signal }); + let settled = false; + let rejection: unknown; + const outcome = cancelled.then( + () => { + settled = true; + }, + (error: unknown) => { + settled = true; + rejection = error; + }, + ); + if (!preAborted) abort.abort(reason); + const observe = vi.fn(async () => undefined); + const following = store.modify("synthetic-provider", observe); + await new Promise((resolve) => setImmediate(resolve)); + + expect(settled).toBe(true); + expect(rejection).toBe(reason); + expect(update).not.toHaveBeenCalled(); + expect(observe).not.toHaveBeenCalled(); + expect(getSecret).not.toHaveBeenCalled(); + expect(setSecret).not.toHaveBeenCalled(); + + release.resolve(); + await active; + await outcome; + await following; + expect(update).not.toHaveBeenCalled(); + expect(observe).toHaveBeenCalledWith(retained); + expect(getSecret).toHaveBeenCalledOnce(); + expect(setSecret).toHaveBeenCalledOnce(); + expect(values.get("provider:synthetic-provider")).toBe(JSON.stringify(retained)); + }, +); + +test("cancellation during the initial secret read prevents the mutation callback", async () => { + const { store, values, getSecret, setSecret } = secretStore(); + const retained = syntheticCredential("retained"); + values.set("provider:synthetic-provider", JSON.stringify(retained)); + const entered = deferred(); + const release = deferred(); + getSecret.mockImplementationOnce(async () => { + entered.resolve(); + await release.promise; + return JSON.stringify(retained); + }); + onTestFinished(() => release.resolve()); + const update = vi.fn(async () => syntheticCredential("cancelled")); + const abort = new AbortController(); + const pending = store.modify("synthetic-provider", update, { signal: abort.signal }); + const result = expect(pending).rejects.toMatchObject({ name: "AbortError" }); + await entered.promise; + abort.abort(); + release.resolve(); + await result; + + expect(update).not.toHaveBeenCalled(); + expect(setSecret).not.toHaveBeenCalled(); + expect(await store.read("synthetic-provider")).toEqual(retained); +}); + +test("cancellation after the mutation starts preserves its rotated credential and queue ownership", async () => { + const { store, values } = secretStore(); + const retained = syntheticCredential("retained"); + const rotated = syntheticCredential("rotated"); + values.set("provider:synthetic-provider", JSON.stringify(retained)); + const entered = deferred(); + const release = deferred(); + const abort = new AbortController(); + const events: string[] = []; + const active = store.modify( + "synthetic-provider", + async (current) => { + expect(current).toEqual(retained); + entered.resolve(); + await release.promise; + events.push("mutation finished"); + return rotated; + }, + { signal: abort.signal }, + ); + onTestFinished(async () => { + release.resolve(); + await active; + await store.modify("synthetic-provider", async () => undefined); + }); + await entered.promise; + abort.abort(); + const observe = vi.fn(async () => { + events.push("following mutation"); + return undefined; + }); + const following = store.modify("synthetic-provider", observe); + await store.modify("other-provider", async () => undefined); + expect(observe).not.toHaveBeenCalled(); + release.resolve(); + + await expect(active).resolves.toEqual(rotated); + await following; + expect(observe).toHaveBeenCalledWith(rotated); + expect(events).toEqual(["mutation finished", "following mutation"]); + expect(await store.read("synthetic-provider")).toEqual(rotated); +}); + +test.each([ + { operation: "modify", fails: false }, + { operation: "modify", fails: true }, + { operation: "delete", fails: false }, + { operation: "delete", fails: true }, +])( + "an active $operation retains the queue after cancellation (write fails: $fails)", + async ({ operation, fails }) => { + const { store, values, setSecret } = secretStore(); + const retained = syntheticCredential("retained"); + const incoming = syntheticCredential("incoming"); + values.set("provider:synthetic-provider", JSON.stringify(retained)); + const entered = deferred(); + const release = deferred(); + const storageError = new Error("Synthetic storage failure"); + const events: string[] = []; + setSecret.mockImplementationOnce(async (key, value) => { + events.push("write started"); + entered.resolve(); + await release.promise; + events.push("write settled"); + if (fails) throw storageError; + values.set(key, value); + }); + const abort = new AbortController(); + const expected = operation === "modify" ? incoming : undefined; + const active = + operation === "modify" + ? store.modify("synthetic-provider", async () => incoming, { signal: abort.signal }) + : store.delete("synthetic-provider", { signal: abort.signal }); + const outcome = active.then( + (credential) => ({ status: "fulfilled", credential }), + (reason: unknown) => ({ status: "rejected", reason }), + ); + onTestFinished(async () => { + release.resolve(); + await outcome; + await store.modify("synthetic-provider", async () => undefined); + }); + await entered.promise; + abort.abort(); + const observe = vi.fn(async () => { + events.push("following mutation"); + return undefined; + }); + const following = store.modify("synthetic-provider", observe); + await store.modify("other-provider", async () => undefined); + expect(observe).not.toHaveBeenCalled(); + release.resolve(); + + expect(await outcome).toEqual( + fails + ? { status: "rejected", reason: storageError } + : { status: "fulfilled", credential: expected }, + ); + await following; + expect(observe).toHaveBeenCalledWith(fails ? retained : expected); + expect(events).toEqual(["write started", "write settled", "following mutation"]); + expect(await store.read("synthetic-provider")).toEqual(fails ? retained : expected); + }, +); + +test("a failed mutation does not poison the following operation", async () => { + const { store, values } = secretStore(); + const retained = syntheticCredential("retained"); + values.set("provider:synthetic-provider", JSON.stringify(retained)); + const failure = new Error("Synthetic mutation failure"); + const failed = store.modify("synthetic-provider", async () => { + throw failure; + }); + const result = expect(failed).rejects.toBe(failure); + const observe = vi.fn(async () => undefined); + const following = store.modify("synthetic-provider", observe); + + await result; + await following; + expect(observe).toHaveBeenCalledWith(retained); +}); + +function secretStore() { + const values = new Map(); + const getSecret = vi.fn(async (key: string, fallback = "") => values.get(key) ?? fallback); + const setSecret = vi.fn(async (key: string, value: string) => { + values.set(key, value); + }); + const context = { getSecret, setSecret } as unknown as Acode.PluginContext; + const store: CredentialStore = new PortableCredentialStore(context); + return { store, values, getSecret, setSecret }; +} + +function syntheticCredential(suffix: string): Credential { + return { + type: "oauth", + access: `synthetic-access-${suffix}`, + refresh: `synthetic-refresh-${suffix}`, + expires: Number.MAX_SAFE_INTEGER, + }; +} + +function deferred() { + let resolve!: (value: T | PromiseLike) => void; + const promise = new Promise((done) => { + resolve = done; + }); + return { promise, resolve }; +}