fix(ai): bound OAuth refresh duration

closes #7508
This commit is contained in:
Vegard Stikbakke
2026-08-04 12:12:32 +02:00
parent f7ea2ef38c
commit acbdc0d25e
3 changed files with 51 additions and 1 deletions
+1
View File
@@ -75,6 +75,7 @@
### Fixed
- Bounded OAuth token refreshes so stalled requests release the credential-store lock ([#7508](https://github.com/earendil-works/pi/issues/7508)).
- Fixed tool argument validation to preserve values that already match an `anyOf`/`oneOf` union arm before attempting coercion, avoiding nullable unions converting `null` to another primitive value ([#7328](https://github.com/earendil-works/pi/issues/7328)).
- Fixed cancellation of model catalog refreshes so callers stop waiting even when a custom provider ignores its abort signal ([#7027](https://github.com/earendil-works/pi/issues/7027)).
- Fixed auth resolution, availability checks, OAuth refreshes, provider login, and in-memory credential queue waits to honor caller cancellation.
+6 -1
View File
@@ -117,6 +117,7 @@ function overlayEnvAuthContext(base: AuthContext, env: ProviderEnv): AuthContext
}
const DEFAULT_OAUTH_MINIMUM_VALIDITY_MS = 5 * 60 * 1000;
const DEFAULT_OAUTH_REFRESH_TIMEOUT_MS = 15_000;
/**
* OAuth resolution with double-checked locking: tokens with less than five
@@ -145,7 +146,11 @@ async function resolveStoredOAuth(
if (current?.type !== "oauth") return undefined; // logged out meanwhile
if (!expiresSoon(current)) return undefined; // another process/request refreshed
try {
return await oauth.refresh(current, signal);
const refreshSignal = AbortSignal.any([
signal,
AbortSignal.timeout(DEFAULT_OAUTH_REFRESH_TIMEOUT_MS),
]);
return await oauth.refresh(current, refreshSignal);
} catch (error) {
throw new ModelsError("oauth", `OAuth refresh failed for ${providerId}`, { cause: error });
}
+44
View File
@@ -5,6 +5,8 @@ import { githubCopilotOAuth } from "../src/auth/oauth/github-copilot.ts";
import { openaiCodexOAuth } from "../src/auth/oauth/openai-codex.ts";
import { openRouterOAuth } from "../src/auth/oauth/openrouter.ts";
import { xaiOAuth } from "../src/auth/oauth/xai.ts";
import { resolveProviderAuth } from "../src/auth/resolve.ts";
import type { OAuthAuth, OAuthCredential } from "../src/auth/types.ts";
import { createModels } from "../src/models.ts";
import * as extensionOAuthCompatibility from "../src/oauth.ts";
import { anthropicProvider } from "../src/providers/anthropic.ts";
@@ -24,6 +26,7 @@ describe.sequential("OAuthAuth adapters", () => {
afterEach(() => {
vi.unstubAllGlobals();
vi.restoreAllMocks();
});
it("anthropic toAuth derives the api key from the access token", async () => {
@@ -116,6 +119,47 @@ describe.sequential("OAuthAuth adapters", () => {
expect(refreshed.enterpriseUrl).toBe("company.ghe.com");
expect(fetchedUrls[0]).toContain("api.company.ghe.com");
});
it("bounds OAuth refreshes and releases the credential-store lock after timeout", async () => {
const timeoutController = new AbortController();
const timeout = vi.spyOn(AbortSignal, "timeout").mockReturnValue(timeoutController.signal);
let markRefreshStarted: (() => void) | undefined;
const refreshStarted = new Promise<void>((resolve) => {
markRefreshStarted = resolve;
});
const credential: OAuthCredential = {
type: "oauth",
access: "old-access",
refresh: "refresh-token",
expires: 0,
};
const oauth: OAuthAuth = {
name: "Stalled OAuth",
login: async () => credential,
refresh: async (_current, signal) => {
markRefreshStarted?.();
return new Promise<OAuthCredential>((_resolve, reject) => {
const rejectAborted = () => reject(signal.reason);
signal.addEventListener("abort", rejectAborted, { once: true });
if (signal.aborted) rejectAborted();
});
},
toAuth: async (current) => ({ apiKey: current.access }),
};
const credentials = new InMemoryCredentialStore();
await credentials.modify("stalled", async () => credential);
const auth = resolveProviderAuth({ id: "stalled", auth: { oauth } }, credentials, {
env: async () => undefined,
fileExists: async () => false,
});
await refreshStarted;
expect(timeout).toHaveBeenCalledWith(15_000);
timeoutController.abort(new DOMException("The operation was aborted due to timeout", "TimeoutError"));
await expect(auth).rejects.toMatchObject({ name: "ModelsError", code: "oauth" });
await expect(credentials.modify("stalled", async (current) => current)).resolves.toEqual(credential);
});
});
describe("OAuth through Models.getAuth (lazy load chain)", () => {