From 68a1629e7e2a066607aefa2699430c6ac4711890 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Victor=20Ara=C3=BAjo?= Date: Thu, 9 Jul 2026 20:15:10 -0300 Subject: [PATCH] fix(ai): bound SuperGrok OAuth usage Keep optional identity enrichment bounded and let the dedicated bearer fall back past incompatible stored keys. --- packages/ai/src/auth-storage.ts | 4 +- .../oauth/__tests__/xai-oauth.test.ts | 51 ++++++++++++++++++- packages/ai/src/registry/oauth/xai-oauth.ts | 4 +- .../test/auth-storage-xai-oauth-usage.test.ts | 28 ++++++++++ 4 files changed, 84 insertions(+), 3 deletions(-) diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index 8e14c9a1d..0e04a3722 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -3230,14 +3230,16 @@ export class AuthStorage { } if (providerId === "xai-oauth") { + let hasUsableStoredOAuthCredential = false; for (const entry of entries) { if (entry.credential.type !== "oauth") continue; const request = this.#buildUsageRequestForOauth(provider, entry.credential, baseUrl); if (providerImpl.supports && !providerImpl.supports(request)) continue; requests.push(request); + hasUsableStoredOAuthCredential = true; } const oauthToken = $env.XAI_OAUTH_TOKEN?.trim(); - if (entries.length === 0 && oauthToken) { + if (!hasUsableStoredOAuthCredential && oauthToken) { const request = this.#buildUsageRequest(provider, { type: "oauth", accessToken: oauthToken }, baseUrl); if (!providerImpl.supports || providerImpl.supports(request)) requests.push(request); } diff --git a/packages/ai/src/registry/oauth/__tests__/xai-oauth.test.ts b/packages/ai/src/registry/oauth/__tests__/xai-oauth.test.ts index 9e5e80942..17bbb6ae2 100644 --- a/packages/ai/src/registry/oauth/__tests__/xai-oauth.test.ts +++ b/packages/ai/src/registry/oauth/__tests__/xai-oauth.test.ts @@ -162,6 +162,55 @@ describe("xAI OAuth helpers", () => { ); expect(requests[0]?.init?.redirect).toBe("error"); }); + + it("combines caller cancellation with the 15-second userinfo timeout", async () => { + const timeoutControllers: AbortController[] = []; + const timeoutSpy = vi.spyOn(AbortSignal, "timeout").mockImplementation(timeoutMs => { + expect(timeoutMs).toBe(15_000); + const controller = new AbortController(); + timeoutControllers.push(controller); + return controller.signal; + }); + const requests: RecordedRequest[] = []; + const fetchMock = vi.fn(async (input: string | URL | Request, init?: RequestInit) => { + requests.push({ + url: typeof input === "string" ? input : input instanceof Request ? input.url : input.toString(), + init, + }); + const { promise, reject } = Promise.withResolvers(); + const requestSignal = init?.signal; + if (!requestSignal) { + reject(new Error("expected userinfo request signal")); + } else if (requestSignal.aborted) { + reject(requestSignal.reason); + } else { + requestSignal.addEventListener("abort", () => reject(requestSignal.reason), { once: true }); + } + return promise; + }); + + const callerController = new AbortController(); + const callerCancelled = fetchXAIOAuthIdentity( + "access-token", + fetchMock as unknown as typeof fetch, + callerController.signal, + ); + callerController.abort(); + await expect(callerCancelled).resolves.toBeNull(); + expect(requests[0]?.init?.signal).not.toBe(callerController.signal); + + expect(timeoutSpy).toHaveBeenCalledWith(15_000); + const timeoutCancelled = fetchXAIOAuthIdentity( + "access-token", + fetchMock as unknown as typeof fetch, + new AbortController().signal, + ); + const timeoutController = timeoutControllers[1]; + expect(timeoutController).toBeDefined(); + timeoutController?.abort(); + await expect(timeoutCancelled).resolves.toBeNull(); + expect(requests[1]?.init?.signal).not.toBe(timeoutController?.signal); + }); }); describe("validateXAIEndpoint", () => { @@ -374,7 +423,7 @@ describe("loginXAIOAuth", () => { accountId: "jwt-sub", }); expect(requests.at(-1)?.url).toBe(USERINFO_URL); - expect(requests.at(-1)?.init?.signal).toBe(controller.signal); + expect(requests.at(-1)?.init?.signal).not.toBe(controller.signal); }); it("keeps the JWT subject when userinfo returns only an email", async () => { diff --git a/packages/ai/src/registry/oauth/xai-oauth.ts b/packages/ai/src/registry/oauth/xai-oauth.ts index 98ad6d1c4..faea9cb18 100644 --- a/packages/ai/src/registry/oauth/xai-oauth.ts +++ b/packages/ai/src/registry/oauth/xai-oauth.ts @@ -190,7 +190,9 @@ export async function fetchXAIOAuthIdentity( Accept: "application/json", }, redirect: "error", - signal: signal ?? AbortSignal.timeout(DISCOVERY_TIMEOUT_MS), + signal: signal + ? AbortSignal.any([signal, AbortSignal.timeout(DISCOVERY_TIMEOUT_MS)]) + : AbortSignal.timeout(DISCOVERY_TIMEOUT_MS), }); if (!response.ok) return null; const payload = (await response.json()) as unknown; diff --git a/packages/ai/test/auth-storage-xai-oauth-usage.test.ts b/packages/ai/test/auth-storage-xai-oauth-usage.test.ts index 7b6c885ca..d20317c77 100644 --- a/packages/ai/test/auth-storage-xai-oauth-usage.test.ts +++ b/packages/ai/test/auth-storage-xai-oauth-usage.test.ts @@ -100,6 +100,34 @@ describe("xAI OAuth environment usage", () => { }); }); + it("uses XAI_OAUTH_TOKEN when stored xAI credentials contain only an API key", async () => { + const calls: UsageFetchParams[] = []; + await withEnv({ XAI_OAUTH_TOKEN: "env-oauth-bearer" }, async () => { + const storage = new AuthStorage( + makeStore([ + { + id: 1, + provider: "xai-oauth", + credential: { + type: "api_key", + key: "stored-api-key", + }, + disabledCause: null, + }, + ]), + { + usageProviderResolver: provider => (provider === "xai-oauth" ? captureUsageProvider(calls) : undefined), + }, + ); + await storage.reload(); + + await storage.fetchUsageReports(); + }); + + expect(calls).toHaveLength(1); + expect(calls[0]?.credential).toEqual({ type: "oauth", accessToken: "env-oauth-bearer" }); + }); + it("does not send shared XAI_API_KEY to the SuperGrok usage endpoint", async () => { const calls: UsageFetchParams[] = []; await withEnv({ XAI_OAUTH_TOKEN: undefined, XAI_API_KEY: "paid-api-key" }, async () => {