diff --git a/src/api/providers/__tests__/xai.spec.ts b/src/api/providers/__tests__/xai.spec.ts index d760c0fd47..763d10d027 100644 --- a/src/api/providers/__tests__/xai.spec.ts +++ b/src/api/providers/__tests__/xai.spec.ts @@ -239,12 +239,7 @@ describe("XAIHandler", () => { it("completePrompt should return text from Responses API", async () => { const expectedResponse = "This is a test response" mockResponsesCreate.mockResolvedValueOnce({ - output: [ - { - type: "message", - content: [{ type: "output_text", text: expectedResponse }], - }, - ], + output_text: expectedResponse, }) const result = await handler.completePrompt("test prompt") @@ -258,6 +253,41 @@ describe("XAIHandler", () => { await expect(handler.completePrompt("test prompt")).rejects.toThrow(`xAI completion error: ${errorMessage}`) }) + it("should include reasoning_effort for mini models", async () => { + const miniModelHandler = new XAIHandler({ + apiModelId: "grok-3-mini", + reasoningEffort: "high", + }) + + mockResponsesCreate.mockResolvedValueOnce(mockStream([])) + + const stream = miniModelHandler.createMessage("test prompt", []) + await stream.next() + + expect(mockResponsesCreate).toHaveBeenCalledWith( + expect.objectContaining({ + reasoning: expect.objectContaining({ + reasoning_effort: "high", + }), + }), + ) + }) + + it("should not include reasoning for non-mini models", async () => { + const regularHandler = new XAIHandler({ + apiModelId: "grok-3", + reasoningEffort: "high", + }) + + mockResponsesCreate.mockResolvedValueOnce(mockStream([])) + + const stream = regularHandler.createMessage("test prompt", []) + await stream.next() + + const callArgs = mockResponsesCreate.mock.calls[mockResponsesCreate.mock.calls.length - 1][0] + expect(callArgs).not.toHaveProperty("reasoning") + }) + it("should handle errors in createMessage", async () => { const errorMessage = "Stream error" mockResponsesCreate.mockRejectedValueOnce(new Error(errorMessage)) diff --git a/src/api/providers/xai.ts b/src/api/providers/xai.ts index daaad88d6c..0cd9cb0273 100644 --- a/src/api/providers/xai.ts +++ b/src/api/providers/xai.ts @@ -9,6 +9,7 @@ import type { ApiHandlerOptions } from "../../shared/api" import { ApiStream } from "../transform/stream" import { convertToResponsesApiInput } from "../transform/responses-api-input" import { processResponsesApiStream, createUsageNormalizer } from "../transform/responses-api-stream" +import { getModelParams } from "../transform/model-params" import { DEFAULT_HEADERS } from "./constants" import { BaseProvider } from "./base-provider" @@ -42,7 +43,15 @@ export class XAIHandler extends BaseProvider implements SingleCompletionHandler ? (this.options.apiModelId as XAIModelId) : xaiDefaultModelId - return { id, info: xaiModels[id] } + const info = xaiModels[id] + const params = getModelParams({ + format: "openai", + modelId: id, + model: info, + settings: this.options, + defaultTemperature: XAI_DEFAULT_TEMPERATURE, + }) + return { id, info, ...params } } /** @@ -54,10 +63,11 @@ export class XAIHandler extends BaseProvider implements SingleCompletionHandler * (additionalProperties: false, ensureAllRequired) and handles MCP tools. */ private mapResponseTools(tools?: any[]): any[] | undefined { - if (!tools?.length) { + const converted = this.convertToolsForOpenAI(tools) + if (!converted?.length) { return undefined } - return tools + return converted .filter((tool) => tool?.type === "function") .map((tool) => { const isMcp = isMcpTool(tool.function.name) @@ -84,22 +94,42 @@ export class XAIHandler extends BaseProvider implements SingleCompletionHandler const input = convertToResponsesApiInput(messages) const responseTools = this.mapResponseTools(metadata?.tools) - let stream + // Build request options + const requestBody: Record = { + model: model.id, + instructions: systemPrompt, + input: input, + stream: true, + store: false, // Don't store responses server-side for privacy + include: ["reasoning.encrypted_content"], + } + + if (model.maxTokens) { + requestBody.max_output_tokens = model.maxTokens + } + + if (model.temperature !== undefined) { + requestBody.temperature = model.temperature + } + + if (responseTools) { + requestBody.tools = responseTools + // Cast tool_choice since metadata uses Chat Completions types but Responses API has its own type + requestBody.tool_choice = (metadata?.tool_choice ?? "auto") as any + requestBody.parallel_tool_calls = metadata?.parallelToolCalls ?? true + } + + // Pass reasoning effort for models that support it (e.g., mini models) + if (model.reasoning) { + requestBody.reasoning = model.reasoning + } + + let stream: AsyncIterable try { - stream = await this.client.responses.create({ - model: model.id, - instructions: systemPrompt, - input: input, - max_output_tokens: model.info.maxTokens, - temperature: this.options.modelTemperature ?? XAI_DEFAULT_TEMPERATURE, + stream = (await this.client.responses.create({ + ...requestBody, stream: true, - store: false, // Don't store responses server-side for privacy - tools: responseTools, - // Cast tool_choice since metadata uses Chat Completions types but Responses API has its own type - tool_choice: (metadata?.tool_choice ?? (responseTools ? "auto" : undefined)) as any, - parallel_tool_calls: metadata?.parallelToolCalls ?? true, - include: ["reasoning.encrypted_content"], - }) + } as any)) as unknown as AsyncIterable } catch (error) { const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model.id, "createMessage") @@ -121,20 +151,8 @@ export class XAIHandler extends BaseProvider implements SingleCompletionHandler store: false, }) - // Extract text from the response output - const output = (response as any).output - if (Array.isArray(output)) { - for (const item of output) { - if (item.type === "message" && Array.isArray(item.content)) { - for (const content of item.content) { - if (content.type === "output_text" && content.text) { - return content.text - } - } - } - } - } - return (response as any).output_text || "" + // output_text is a convenience field on the Responses API response + return response.output_text || "" } catch (error) { const errorMessage = error instanceof Error ? error.message : String(error) const apiError = new ApiProviderError(errorMessage, this.providerName, model.id, "completePrompt") diff --git a/src/api/transform/responses-api-stream.ts b/src/api/transform/responses-api-stream.ts index 1ea9744180..884a755136 100644 --- a/src/api/transform/responses-api-stream.ts +++ b/src/api/transform/responses-api-stream.ts @@ -131,8 +131,8 @@ export function createUsageNormalizer( type: "usage", inputTokens, outputTokens, - cacheReadTokens, cacheWriteTokens, + cacheReadTokens, ...(typeof reasoningTokens === "number" ? { reasoningTokens } : {}), ...(typeof totalCost === "number" ? { totalCost } : {}), }