From 5220aba61970fc9f284a7c71e95669b52b4f47a8 Mon Sep 17 00:00:00 2001 From: Milosz Jankiewicz <25470423+Jaaneek@users.noreply.github.com> Date: Thu, 16 Jul 2026 08:58:24 +0100 Subject: [PATCH] feat(ai): add xAI device OAuth and route grok-4.5 through Responses (#6651) * feat(ai): add xAI device OAuth and route grok-4.5 through Responses Add xAI device-code OAuth alongside XAI_API_KEY. Route only grok-4.5 through Responses with low/medium/high reasoning; other xAI models stay on Completions. * fix(ai): tolerate xAI device-code interval 0 and correct token poll error label --------- Co-authored-by: Jaaneek --- packages/ai/scripts/generate-models.ts | 15 +- packages/ai/src/api/openai-responses.ts | 1 + packages/ai/src/auth/oauth/load.ts | 3 + packages/ai/src/auth/oauth/xai.ts | 231 +++++++++++++++++++ packages/ai/src/providers/xai.models.ts | 7 +- packages/ai/src/providers/xai.ts | 16 +- packages/ai/test/oauth-auth.test.ts | 6 + packages/ai/test/xai-oauth.test.ts | 285 ++++++++++++++++++++++++ packages/ai/test/xai-responses.test.ts | 105 +++++++++ packages/coding-agent/docs/providers.md | 6 + 10 files changed, 667 insertions(+), 8 deletions(-) create mode 100644 packages/ai/src/auth/oauth/xai.ts create mode 100644 packages/ai/test/xai-oauth.test.ts create mode 100644 packages/ai/test/xai-responses.test.ts diff --git a/packages/ai/scripts/generate-models.ts b/packages/ai/scripts/generate-models.ts index 5fad8aaf..6f194818 100644 --- a/packages/ai/scripts/generate-models.ts +++ b/packages/ai/scripts/generate-models.ts @@ -255,6 +255,14 @@ const OPENAI_RESPONSES_NONE_REASONING_MODELS = new Set([ "gpt-5.6-terra", "gpt-5.6-luna", ]); +const XAI_RESPONSES_MODEL_ID = "grok-4.5"; +const XAI_RESPONSES_EFFORT_LEVEL_MAP = { + off: null, + minimal: null, +} as const; +const XAI_RESPONSES_COMPAT: OpenAIResponsesCompat = { + supportsLongCacheRetention: false, +}; const OPENCODE_OPENAI_COMPLETIONS_LONG_CACHE_RETENTION_UNSUPPORTED_MODELS = new Set([ "opencode:deepseek-v4-flash", @@ -538,6 +546,9 @@ function applyThinkingLevelMetadata(model: Model): void { ) { mergeThinkingLevelMap(model, { off: "none" }); } + if (model.provider === "xai" && model.api === "openai-responses" && model.id === XAI_RESPONSES_MODEL_ID) { + mergeThinkingLevelMap(model, XAI_RESPONSES_EFFORT_LEVEL_MAP); + } if (supportsOpenAiXhigh(model.id)) { mergeThinkingLevelMap(model, { xhigh: "xhigh" }); } @@ -1126,13 +1137,15 @@ async function loadModelsDevData(): Promise[]> { for (const [modelId, model] of Object.entries(data.xai.models)) { const m = model as ModelsDevModel; if (m.tool_call !== true) continue; + const useResponsesApi = modelId === XAI_RESPONSES_MODEL_ID; models.push({ id: modelId, name: m.name || modelId, - api: "openai-completions", + api: useResponsesApi ? "openai-responses" : "openai-completions", provider: "xai", baseUrl: "https://api.x.ai/v1", + ...(useResponsesApi ? { compat: { ...XAI_RESPONSES_COMPAT } } : {}), reasoning: m.reasoning === true, input: m.modalities?.input?.includes("image") ? ["text", "image"] : ["text"], cost: { diff --git a/packages/ai/src/api/openai-responses.ts b/packages/ai/src/api/openai-responses.ts index cb3a9c93..049e00ec 100644 --- a/packages/ai/src/api/openai-responses.ts +++ b/packages/ai/src/api/openai-responses.ts @@ -282,6 +282,7 @@ function buildParams(model: Model<"openai-responses">, context: Context, options effort: (model.thinkingLevelMap?.off ?? "none") as NonNullable["effort"], }; } + if (model.provider === "xai") params.include = ["reasoning.encrypted_content"]; } return params; diff --git a/packages/ai/src/auth/oauth/load.ts b/packages/ai/src/auth/oauth/load.ts index 6984da00..1e5fb862 100644 --- a/packages/ai/src/auth/oauth/load.ts +++ b/packages/ai/src/auth/oauth/load.ts @@ -20,6 +20,9 @@ export const loadOpenAICodexOAuth = async (): Promise => export const loadGitHubCopilotOAuth = async (): Promise => ((await importOAuthModule("./github-copilot.ts")) as { githubCopilotOAuth: OAuthAuth }).githubCopilotOAuth; +export const loadXaiOAuth = async (): Promise => + ((await importOAuthModule("./xai.ts")) as { xaiOAuth: OAuthAuth }).xaiOAuth; + export const loadRadiusOAuth = async (options: { name: string; gateway: string }): Promise => ( (await importOAuthModule("./radius.ts")) as { diff --git a/packages/ai/src/auth/oauth/xai.ts b/packages/ai/src/auth/oauth/xai.ts new file mode 100644 index 00000000..4c508b6b --- /dev/null +++ b/packages/ai/src/auth/oauth/xai.ts @@ -0,0 +1,231 @@ +/** + * xAI OAuth device-code flow. + */ + +import type { AuthInteraction, OAuthAuth, OAuthCredential } from "../types.ts"; +import { pollOAuthDeviceCodeFlow } from "./device-code.ts"; + +const XAI_CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828"; +const XAI_SCOPE = "openid profile email offline_access grok-cli:access api:access"; +const XAI_DEVICE_CODE_URL = "https://auth.x.ai/oauth2/device/code"; +const XAI_TOKEN_URL = "https://auth.x.ai/oauth2/token"; +// Refresh slightly before the reported expiry to avoid using a token that dies mid-request. +const REFRESH_SKEW_MS = 5 * 60 * 1000; +const DEFAULT_TOKEN_LIFETIME_SECONDS = 3600; + +type JsonObject = Record; + +type OAuthHttpResponse = { + ok: boolean; + status: number; + body: JsonObject; +}; + +type XaiDeviceCode = { + deviceCode: string; + userCode: string; + verificationUri: string; + intervalSeconds?: number; + expiresInSeconds: number; +}; + +function requiredString(body: JsonObject, field: string): string { + const value = body[field]; + if (typeof value !== "string" || value.length === 0) { + throw new Error(`Invalid xAI OAuth response field: ${field}`); + } + return value; +} + +function positiveNumber(body: JsonObject, field: string): number { + const value = body[field]; + if (typeof value !== "number" || !Number.isFinite(value) || value <= 0) { + throw new Error(`Invalid xAI OAuth response field: ${field}`); + } + return value; +} + +// The verification URI is opened in the user's browser; force it to be an https URL +// so a malicious response cannot make `open` launch something else. +function validateVerificationUri(raw: string): string { + let url: URL; + try { + url = new URL(raw); + } catch { + throw new Error("Untrusted verification URI in xAI OAuth response"); + } + if (url.protocol !== "https:") { + throw new Error("Untrusted verification URI in xAI OAuth response"); + } + return url.href; +} + +async function postForm(url: string, fields: Record, signal?: AbortSignal): Promise { + let response: Response; + try { + response = await fetch(url, { + method: "POST", + headers: { + Accept: "application/json", + "Content-Type": "application/x-www-form-urlencoded", + }, + body: new URLSearchParams(fields), + signal, + }); + } catch (error) { + if (signal?.aborted) { + throw new Error("Login cancelled"); + } + throw error; + } + + let body: JsonObject; + try { + const parsed = (await response.json()) as unknown; + body = parsed && typeof parsed === "object" && !Array.isArray(parsed) ? (parsed as JsonObject) : {}; + } catch { + if (signal?.aborted) { + throw new Error("Login cancelled"); + } + throw new Error(`xAI OAuth returned invalid JSON (HTTP ${response.status})`); + } + return { + ok: response.ok, + status: response.status, + body, + }; +} + +function requestFailure(action: string, response: OAuthHttpResponse): Error { + const error = typeof response.body.error === "string" ? response.body.error : undefined; + const description = + typeof response.body.error_description === "string" ? response.body.error_description : undefined; + const detail = [error, description].filter(Boolean).join(": "); + return new Error(`xAI OAuth ${action} failed (HTTP ${response.status})${detail ? `: ${detail}` : ""}`); +} + +function parseDeviceCode(body: JsonObject): XaiDeviceCode { + // RFC 8628 allows interval 0 (no minimum wait); fall back to the poller's + // default instead of failing on non-positive or malformed values. + const interval = body.interval; + const intervalSeconds = + typeof interval === "number" && Number.isFinite(interval) && interval > 0 ? interval : undefined; + return { + deviceCode: requiredString(body, "device_code"), + userCode: requiredString(body, "user_code"), + verificationUri: validateVerificationUri(requiredString(body, "verification_uri")), + intervalSeconds, + expiresInSeconds: positiveNumber(body, "expires_in"), + }; +} + +function credentialsFromTokenResponse(body: JsonObject, previousRefreshToken?: string): OAuthCredential { + const access = requiredString(body, "access_token"); + // xAI may omit refresh_token on refresh when the token is not rotated. + const refresh = + body.refresh_token === undefined && previousRefreshToken + ? previousRefreshToken + : requiredString(body, "refresh_token"); + const expiresInSeconds = + body.expires_in === undefined ? DEFAULT_TOKEN_LIFETIME_SECONDS : positiveNumber(body, "expires_in"); + return { + type: "oauth", + access, + refresh, + expires: Date.now() + expiresInSeconds * 1000 - REFRESH_SKEW_MS, + }; +} + +async function requestDeviceCode(signal?: AbortSignal): Promise { + const response = await postForm( + XAI_DEVICE_CODE_URL, + { + client_id: XAI_CLIENT_ID, + scope: XAI_SCOPE, + referrer: "pi", + }, + signal, + ); + if (!response.ok) { + throw requestFailure("device authorization", response); + } + return parseDeviceCode(response.body); +} + +async function pollForTokens(device: XaiDeviceCode, signal?: AbortSignal): Promise { + return pollOAuthDeviceCodeFlow({ + intervalSeconds: device.intervalSeconds, + expiresInSeconds: device.expiresInSeconds, + waitBeforeFirstPoll: true, + signal, + poll: async () => { + const response = await postForm( + XAI_TOKEN_URL, + { + grant_type: "urn:ietf:params:oauth:grant-type:device_code", + client_id: XAI_CLIENT_ID, + device_code: device.deviceCode, + }, + signal, + ); + + if (response.ok) { + return { status: "complete", value: credentialsFromTokenResponse(response.body) }; + } + + const error = response.body.error; + if (error === "authorization_pending") { + return { status: "pending" }; + } + if (error === "slow_down") { + const interval = response.body.interval; + return { status: "slow_down", intervalSeconds: typeof interval === "number" ? interval : undefined }; + } + if (error === "access_denied" || error === "authorization_denied") { + return { status: "failed", message: "xAI device authorization was denied" }; + } + if (error === "expired_token") { + return { status: "failed", message: "xAI device code expired" }; + } + return { status: "failed", message: requestFailure("device token polling", response).message }; + }, + }); +} + +async function loginXai(interaction: AuthInteraction): Promise { + const device = await requestDeviceCode(interaction.signal); + interaction.notify({ + type: "device_code", + userCode: device.userCode, + verificationUri: device.verificationUri, + intervalSeconds: device.intervalSeconds, + expiresInSeconds: device.expiresInSeconds, + }); + return pollForTokens(device, interaction.signal); +} + +async function refreshXaiToken(refreshToken: string, signal?: AbortSignal): Promise { + const response = await postForm( + XAI_TOKEN_URL, + { + grant_type: "refresh_token", + client_id: XAI_CLIENT_ID, + refresh_token: refreshToken, + }, + signal, + ); + if (!response.ok) { + throw requestFailure("token refresh", response); + } + return credentialsFromTokenResponse(response.body, refreshToken); +} + +export const xaiOAuth: OAuthAuth = { + name: "xAI (Grok/X subscription)", + login: loginXai, + refresh: (credential, signal) => refreshXaiToken(credential.refresh, signal), + + async toAuth(credential) { + return { apiKey: credential.access }; + }, +}; diff --git a/packages/ai/src/providers/xai.models.ts b/packages/ai/src/providers/xai.models.ts index 57580102..2c781afe 100644 --- a/packages/ai/src/providers/xai.models.ts +++ b/packages/ai/src/providers/xai.models.ts @@ -97,11 +97,12 @@ export const XAI_MODELS = { "grok-4.5": { id: "grok-4.5", name: "Grok 4.5", - api: "openai-completions", + api: "openai-responses", provider: "xai", baseUrl: "https://api.x.ai/v1", - compat: {"supportsStore":false,"supportsDeveloperRole":false,"supportsReasoningEffort":false}, + compat: {"supportsLongCacheRetention":false}, reasoning: true, + thinkingLevelMap: {"off":null,"minimal":null}, input: ["text", "image"], cost: { input: 2, @@ -111,7 +112,7 @@ export const XAI_MODELS = { }, contextWindow: 500000, maxTokens: 500000, - } satisfies Model<"openai-completions">, + } satisfies Model<"openai-responses">, "grok-build-0.1": { id: "grok-build-0.1", name: "Grok Build 0.1", diff --git a/packages/ai/src/providers/xai.ts b/packages/ai/src/providers/xai.ts index 3373fbf5..3a679a2f 100644 --- a/packages/ai/src/providers/xai.ts +++ b/packages/ai/src/providers/xai.ts @@ -1,15 +1,23 @@ import { openAICompletionsApi } from "../api/openai-completions.lazy.ts"; -import { envApiKeyAuth } from "../auth/helpers.ts"; +import { openAIResponsesApi } from "../api/openai-responses.lazy.ts"; +import { envApiKeyAuth, lazyOAuth } from "../auth/helpers.ts"; +import { loadXaiOAuth } from "../auth/oauth/load.ts"; import { createProvider, type Provider } from "../models.ts"; import { XAI_MODELS } from "./xai.models.ts"; -export function xaiProvider(): Provider<"openai-completions"> { +export function xaiProvider(): Provider<"openai-completions" | "openai-responses"> { return createProvider({ id: "xai", name: "xAI", baseUrl: "https://api.x.ai/v1", - auth: { apiKey: envApiKeyAuth("xAI API key", ["XAI_API_KEY"]) }, + auth: { + apiKey: envApiKeyAuth("xAI API key", ["XAI_API_KEY"]), + oauth: lazyOAuth({ name: "xAI (Grok/X subscription)", load: loadXaiOAuth }), + }, models: Object.values(XAI_MODELS), - api: openAICompletionsApi(), + api: { + "openai-completions": openAICompletionsApi(), + "openai-responses": openAIResponsesApi(), + }, }); } diff --git a/packages/ai/test/oauth-auth.test.ts b/packages/ai/test/oauth-auth.test.ts index 68534901..6620002d 100644 --- a/packages/ai/test/oauth-auth.test.ts +++ b/packages/ai/test/oauth-auth.test.ts @@ -3,6 +3,7 @@ import { InMemoryCredentialStore } from "../src/auth/credential-store.ts"; import { anthropicOAuth } from "../src/auth/oauth/anthropic.ts"; import { githubCopilotOAuth } from "../src/auth/oauth/github-copilot.ts"; import { openaiCodexOAuth } from "../src/auth/oauth/openai-codex.ts"; +import { xaiOAuth } from "../src/auth/oauth/xai.ts"; import { createModels } from "../src/models.ts"; import * as extensionOAuthCompatibility from "../src/oauth.ts"; import { anthropicProvider } from "../src/providers/anthropic.ts"; @@ -32,6 +33,11 @@ describe.sequential("OAuthAuth adapters", () => { expect(auth).toEqual({ apiKey: "token" }); }); + it("xAI toAuth derives the api key from the access token", async () => { + const auth = await xaiOAuth.toAuth({ type: "oauth", access: "token", refresh: "r", expires: 0 }); + expect(auth).toEqual({ apiKey: "token" }); + }); + it("github-copilot toAuth derives baseUrl from the token proxy endpoint", async () => { const access = "tid=abc;exp=123;proxy-ep=proxy.enterprise.example;rest"; const auth = await githubCopilotOAuth.toAuth({ type: "oauth", access, refresh: "r", expires: 0 }); diff --git a/packages/ai/test/xai-oauth.test.ts b/packages/ai/test/xai-oauth.test.ts new file mode 100644 index 00000000..fb71a621 --- /dev/null +++ b/packages/ai/test/xai-oauth.test.ts @@ -0,0 +1,285 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; +import { xaiOAuth } from "../src/auth/oauth/xai.ts"; +import type { OAuthCredential } from "../src/auth/types.ts"; + +function jsonResponse(body: unknown, status = 200): Response { + return new Response(JSON.stringify(body), { + status, + headers: { "Content-Type": "application/json" }, + }); +} + +function requestUrl(input: unknown): string { + if (typeof input === "string") return input; + if (input instanceof URL) return input.toString(); + if (input instanceof Request) return input.url; + throw new Error(`Unsupported request input: ${String(input)}`); +} + +function requestForm(init: RequestInit | undefined): URLSearchParams { + return new URLSearchParams(String(init?.body)); +} + +function deviceCodeResponse(overrides: Record = {}): Record { + return { + device_code: "device-code", + user_code: "ABCD-1234", + verification_uri: "https://accounts.x.ai/oauth2/device", + expires_in: 900, + interval: 5, + ...overrides, + }; +} + +function tokenResponse(overrides: Record = {}): Record { + return { + access_token: "access-token", + refresh_token: "refresh-token", + expires_in: 21_600, + token_type: "Bearer", + ...overrides, + }; +} + +type DeviceCodeInfo = { + userCode: string; + verificationUri: string; + intervalSeconds?: number; + expiresInSeconds?: number; +}; + +function loginXaiForTest(options: { + onDeviceCode: (info: DeviceCodeInfo) => void; + signal?: AbortSignal; +}): Promise { + return xaiOAuth.login({ + signal: options.signal, + prompt: () => { + throw new Error("Unexpected prompt"); + }, + notify: (event) => { + if (event.type === "device_code") { + const { type: _, ...info } = event; + options.onDeviceCode(info); + } + }, + }); +} + +function refreshXaiForTest(refreshToken: string): Promise { + return xaiOAuth.refresh({ type: "oauth", access: "old-access", refresh: refreshToken, expires: 0 }); +} + +describe("xAI OAuth device flow", () => { + afterEach(() => { + vi.restoreAllMocks(); + vi.unstubAllGlobals(); + vi.useRealTimers(); + }); + + it("uses the device grant, delays polling, and handles pending and slow_down", async () => { + vi.useFakeTimers(); + const startTime = new Date("2026-07-09T20:00:00Z"); + vi.setSystemTime(startTime); + const pollTimes: number[] = []; + const tokenReplies = [ + jsonResponse({ error: "authorization_pending" }, 400), + jsonResponse({ error: "slow_down", interval: 10 }, 400), + jsonResponse(tokenResponse()), + ]; + + const fetchMock = vi.fn(async (input: unknown, init?: RequestInit) => { + const url = requestUrl(input); + + if (url === "https://auth.x.ai/oauth2/device/code") { + const form = requestForm(init); + expect(form.get("client_id")).toBe("b1a00492-073a-47ea-816f-4c329264a828"); + expect(form.get("scope")).toBe("openid profile email offline_access grok-cli:access api:access"); + expect(form.get("referrer")).toBe("pi"); + return jsonResponse(deviceCodeResponse()); + } + + if (url === "https://auth.x.ai/oauth2/token") { + pollTimes.push(Date.now()); + const form = requestForm(init); + expect(form.get("grant_type")).toBe("urn:ietf:params:oauth:grant-type:device_code"); + expect(form.get("client_id")).toBe("b1a00492-073a-47ea-816f-4c329264a828"); + expect(form.get("device_code")).toBe("device-code"); + const reply = tokenReplies.shift(); + if (!reply) throw new Error("Unexpected token poll"); + return reply; + } + + throw new Error(`Unexpected request: ${url}`); + }); + vi.stubGlobal("fetch", fetchMock); + + const deviceCodes: DeviceCodeInfo[] = []; + const loginPromise = loginXaiForTest({ onDeviceCode: (info) => deviceCodes.push(info) }); + + await vi.advanceTimersByTimeAsync(0); + expect(deviceCodes).toEqual([ + { + userCode: "ABCD-1234", + verificationUri: "https://accounts.x.ai/oauth2/device", + intervalSeconds: 5, + expiresInSeconds: 900, + }, + ]); + expect(pollTimes).toEqual([]); + + await vi.advanceTimersByTimeAsync(5000); + expect(pollTimes).toEqual([startTime.getTime() + 5000]); + + // slow_down raised the interval to 10 seconds + await vi.advanceTimersByTimeAsync(5000); + expect(pollTimes).toEqual([startTime.getTime() + 5000, startTime.getTime() + 10_000]); + + await vi.advanceTimersByTimeAsync(10_000); + const credentials = await loginPromise; + expect(pollTimes).toEqual([ + startTime.getTime() + 5000, + startTime.getTime() + 10_000, + startTime.getTime() + 20_000, + ]); + expect(credentials).toEqual({ + type: "oauth", + access: "access-token", + refresh: "refresh-token", + expires: startTime.getTime() + 20_000 + 21_600_000 - 300_000, + }); + }); + + it("falls back to the default poll interval when the response reports interval 0", async () => { + vi.useFakeTimers(); + const startTime = new Date("2026-07-09T20:00:00Z"); + vi.setSystemTime(startTime); + const pollTimes: number[] = []; + vi.stubGlobal( + "fetch", + vi.fn(async (input: unknown) => { + if (requestUrl(input) === "https://auth.x.ai/oauth2/device/code") { + return jsonResponse(deviceCodeResponse({ interval: 0 })); + } + pollTimes.push(Date.now()); + return jsonResponse(tokenResponse()); + }), + ); + + const loginPromise = loginXaiForTest({ onDeviceCode: () => {} }); + // RFC 8628 default interval is 5 seconds when the server does not require a wait. + await vi.advanceTimersByTimeAsync(5000); + await loginPromise; + expect(pollTimes).toEqual([startTime.getTime() + 5000]); + }); + + it.each(["http://accounts.x.ai/oauth2/device", "file:///etc/passwd", "not a url"])( + "rejects a non-https verification URI: %s", + async (verificationUri) => { + vi.stubGlobal( + "fetch", + vi.fn(async () => jsonResponse(deviceCodeResponse({ verification_uri: verificationUri }))), + ); + + await expect(loginXaiForTest({ onDeviceCode: () => {} })).rejects.toThrow("Untrusted verification URI"); + }, + ); + + it.each(["access_denied", "authorization_denied"])( + "fails when device authorization is denied: %s", + async (error) => { + vi.useFakeTimers(); + let requestCount = 0; + vi.stubGlobal( + "fetch", + vi.fn(async () => { + requestCount += 1; + return requestCount === 1 + ? jsonResponse(deviceCodeResponse({ interval: 1 })) + : jsonResponse({ error }, 400); + }), + ); + + const loginPromise = loginXaiForTest({ onDeviceCode: () => {} }); + const assertion = expect(loginPromise).rejects.toThrow("xAI device authorization was denied"); + await vi.advanceTimersByTimeAsync(1000); + await assertion; + }, + ); + + it("cancels while waiting for the first token poll", async () => { + vi.useFakeTimers(); + const controller = new AbortController(); + const fetchMock = vi.fn(async () => jsonResponse(deviceCodeResponse())); + vi.stubGlobal("fetch", fetchMock); + + const loginPromise = loginXaiForTest({ + onDeviceCode: () => controller.abort(), + signal: controller.signal, + }); + + await expect(loginPromise).rejects.toThrow("Login cancelled"); + expect(fetchMock).toHaveBeenCalledTimes(1); + }); + + it("refreshes tokens and preserves an unrotated refresh token", async () => { + let requestCount = 0; + const fetchMock = vi.fn(async (input: unknown, init?: RequestInit) => { + expect(requestUrl(input)).toBe("https://auth.x.ai/oauth2/token"); + const form = requestForm(init); + expect(form.get("grant_type")).toBe("refresh_token"); + expect(form.get("client_id")).toBe("b1a00492-073a-47ea-816f-4c329264a828"); + requestCount += 1; + if (requestCount === 1) { + expect(form.get("refresh_token")).toBe("old-refresh"); + return jsonResponse(tokenResponse({ access_token: "new-access", refresh_token: "new-refresh" })); + } + expect(form.get("refresh_token")).toBe("keep-refresh"); + return jsonResponse(tokenResponse({ access_token: "newer-access", refresh_token: undefined })); + }); + vi.stubGlobal("fetch", fetchMock); + + const rotated = await refreshXaiForTest("old-refresh"); + const preserved = await refreshXaiForTest("keep-refresh"); + expect(rotated.type).toBe("oauth"); + expect(rotated.refresh).toBe("new-refresh"); + expect(rotated.access).toBe("new-access"); + expect(preserved.refresh).toBe("keep-refresh"); + expect(preserved.access).toBe("newer-access"); + expect(xaiOAuth.name).toBe("xAI (Grok/X subscription)"); + await expect(xaiOAuth.toAuth(preserved)).resolves.toEqual({ apiKey: "newer-access" }); + }); + + it("assumes a one-hour lifetime when expires_in is missing", async () => { + vi.useFakeTimers(); + const startTime = new Date("2026-07-09T20:00:00Z"); + vi.setSystemTime(startTime); + vi.stubGlobal( + "fetch", + vi.fn(async () => jsonResponse(tokenResponse({ expires_in: undefined }))), + ); + + const credentials = await refreshXaiForTest("old-refresh"); + expect(credentials.expires).toBe(startTime.getTime() + 3_600_000 - 300_000); + }); + + it("rejects token responses with missing fields", async () => { + vi.stubGlobal( + "fetch", + vi.fn(async () => jsonResponse(tokenResponse({ access_token: undefined }))), + ); + + await expect(refreshXaiForTest("old-refresh")).rejects.toThrow("Invalid xAI OAuth response field: access_token"); + }); + + it("surfaces the upstream error code and description on refresh failure", async () => { + vi.stubGlobal( + "fetch", + vi.fn(async () => jsonResponse({ error: "invalid_grant", error_description: "refresh token revoked" }, 400)), + ); + + await expect(refreshXaiForTest("old-refresh")).rejects.toThrow( + "xAI OAuth token refresh failed (HTTP 400): invalid_grant: refresh token revoked", + ); + }); +}); diff --git a/packages/ai/test/xai-responses.test.ts b/packages/ai/test/xai-responses.test.ts new file mode 100644 index 00000000..303c62c5 --- /dev/null +++ b/packages/ai/test/xai-responses.test.ts @@ -0,0 +1,105 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; +import type { OpenAIResponsesOptions } from "../src/api/openai-responses.ts"; +import { getSupportedThinkingLevels } from "../src/models.ts"; +import { XAI_MODELS } from "../src/providers/xai.models.ts"; +import { xaiProvider } from "../src/providers/xai.ts"; +import type { Context, Model } from "../src/types.ts"; + +type CapturedRequest = { + url: string; + headers: Headers; + body: Record; +}; + +function completedResponse(): Response { + const event = { + type: "response.completed", + sequence_number: 0, + response: { + id: "resp_xai_test", + status: "completed", + output: [], + usage: { + input_tokens: 1, + output_tokens: 1, + total_tokens: 2, + input_tokens_details: { cached_tokens: 0 }, + }, + }, + }; + return new Response(`data: ${JSON.stringify(event)}\n\ndata: [DONE]\n\n`, { + status: 200, + headers: { "content-type": "text/event-stream" }, + }); +} + +async function captureRequest( + model: Model<"openai-responses">, + context: Context, + options: OpenAIResponsesOptions, +): Promise { + let captured: CapturedRequest | undefined; + vi.spyOn(globalThis, "fetch").mockImplementation(async (input, init) => { + const request = new Request(input, init); + captured = { + url: request.url, + headers: request.headers, + body: JSON.parse(await request.clone().text()) as Record, + }; + return completedResponse(); + }); + + const result = await xaiProvider().stream(model, context, options).result(); + expect(result.stopReason, result.errorMessage).toBe("stop"); + expect(captured).toBeDefined(); + return captured!; +} + +describe("xAI Responses provider", () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("uses Responses with low/medium/high efforts only for Grok 4.5", () => { + expect(XAI_MODELS["grok-4.5"].api).toBe("openai-responses"); + expect(getSupportedThinkingLevels(XAI_MODELS["grok-4.5"])).toEqual(["low", "medium", "high"]); + expect(XAI_MODELS["grok-4.3"].api).toBe("openai-completions"); + }); + + it("uses /responses with bearer auth and xAI-compatible request fields", async () => { + const captured = await captureRequest( + XAI_MODELS["grok-4.5"], + { + systemPrompt: "You are a careful coding assistant.", + messages: [{ role: "user", content: "hello", timestamp: 1 }], + }, + { + apiKey: "xai-test-token", + sessionId: "pi-session-123", + cacheRetention: "long", + reasoningEffort: "medium", + }, + ); + + expect(captured.url).toBe("https://api.x.ai/v1/responses"); + expect(captured.headers.get("authorization")).toBe("Bearer xai-test-token"); + expect(captured.headers.get("session_id")).toBe("pi-session-123"); + expect(captured.body).toMatchObject({ + model: "grok-4.5", + store: false, + stream: true, + prompt_cache_key: "pi-session-123", + reasoning: { effort: "medium" }, + include: ["reasoning.encrypted_content"], + }); + expect(captured.body).not.toHaveProperty("prompt_cache_retention"); + expect(captured.body.input).toEqual( + expect.arrayContaining([ + expect.objectContaining({ + role: "developer", + content: "You are a careful coding assistant.", + }), + ]), + ); + }); +}); diff --git a/packages/coding-agent/docs/providers.md b/packages/coding-agent/docs/providers.md index f7d9983f..d3f832ef 100644 --- a/packages/coding-agent/docs/providers.md +++ b/packages/coding-agent/docs/providers.md @@ -18,6 +18,7 @@ Use `/login` in interactive mode, then select a provider: - ChatGPT Plus/Pro (Codex) - Claude Pro/Max - GitHub Copilot +- xAI (Grok/X subscription) - Radius Use `/logout` to clear credentials. Tokens are stored in `~/.pi/agent/auth.json` and auto-refresh when expired. @@ -36,6 +37,11 @@ Anthropic subscription auth is active for Claude Pro/Max accounts. Third-party h - Press Enter for github.com, or enter your GitHub Enterprise Server domain - If you get "model not supported", enable it in VS Code: Copilot Chat → model selector → select model → "Enable" +### xAI (Grok/X subscription) + +- Run `/login xai`, then select **Use a subscription** +- `XAI_API_KEY` remains available through **Use an API key** + ### Radius Radius is a dynamic `pi-messages` gateway. `/login radius` stores OAuth tokens in `auth.json`; the gateway catalog is refreshed independently and cached in `models-store.json`. Custom Radius gateways can be declared in `models.json` with `"oauth": "radius"` and a gateway `baseUrl`.