feat(api): migrate Fireworks provider to AI SDK (#11118)

This commit is contained in:
Daniel 2026-01-30 14:52:49 -05:00 • committed by GitHub
parent 0cd257af89
commit b5ae557834
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 1399 additions and 916 deletions

30
pnpm-lock.yaml generated
View file

@ -752,6 +752,9 @@ importers:
'@ai-sdk/deepseek':
specifier: ^2.0.14
version: 2.0.14(zod@3.25.76)
'@ai-sdk/fireworks':
specifier: ^2.0.26
version: 2.0.26(zod@3.25.76)
'@ai-sdk/groq':
specifier: ^3.0.19
version: 3.0.19(zod@3.25.76)
@ -1411,6 +1414,12 @@ packages:
peerDependencies:
zod: 3.25.76
'@ai-sdk/fireworks@2.0.26':
resolution: {integrity: sha512-vBqSSksHhDGrSNYnmEmVGvLicHFjL4yAxFZfCb6ydrg+qgnlW2bdyTQDMI69BKG4spNZ1/iHMxRNIQpx19Yf6w==}
engines: {node: '>=18'}
peerDependencies:
zod: 3.25.76
'@ai-sdk/gateway@3.0.25':
resolution: {integrity: sha512-j0AQeA7hOVqwImykQlganf/Euj3uEXf0h3G0O4qKTDpEwE+EZGIPnVimCWht5W91lAetPZSfavDyvfpuPDd2PQ==}
engines: {node: '>=18'}
@ -1429,6 +1438,12 @@ packages:
peerDependencies:
zod: 3.25.76
'@ai-sdk/openai-compatible@2.0.24':
resolution: {integrity: sha512-3QrCKpQCn3g6sIMoFGuEroaqk7Xg+qfsohRp4dKszjto5stjBg4SdtOKqHg+CpE3X4woj2O62w2qr5dSekMZeQ==}
engines: {node: '>=18'}
peerDependencies:
zod: 3.25.76
'@ai-sdk/provider-utils@3.0.20':
resolution: {integrity: sha512-iXHVe0apM2zUEzauqJwqmpC37A5rihrStAih5Ks+JE32iTe4LZ58y17UGBjpQQTCRw9YxMeo2UFLxLpBluyvLQ==}
engines: {node: '>=18'}
@ -11042,6 +11057,13 @@ snapshots:
'@ai-sdk/provider-utils': 4.0.10(zod@3.25.76)
zod: 3.25.76
'@ai-sdk/fireworks@2.0.26(zod@3.25.76)':
dependencies:
'@ai-sdk/openai-compatible': 2.0.24(zod@3.25.76)
'@ai-sdk/provider': 3.0.6
'@ai-sdk/provider-utils': 4.0.11(zod@3.25.76)
zod: 3.25.76
'@ai-sdk/gateway@3.0.25(zod@3.25.76)':
dependencies:
'@ai-sdk/provider': 3.0.5
@ -11061,6 +11083,12 @@ snapshots:
'@ai-sdk/provider-utils': 3.0.20(zod@3.25.76)
zod: 3.25.76
'@ai-sdk/openai-compatible@2.0.24(zod@3.25.76)':
dependencies:
'@ai-sdk/provider': 3.0.6
'@ai-sdk/provider-utils': 4.0.11(zod@3.25.76)
zod: 3.25.76
'@ai-sdk/provider-utils@3.0.20(zod@3.25.76)':
dependencies:
'@ai-sdk/provider': 2.0.1
@ -15070,7 +15098,7 @@ snapshots:
sirv: 3.0.1
tinyglobby: 0.2.14
tinyrainbow: 2.0.0
vitest: 3.2.4(@types/debug@4.1.12)(@types/node@20.17.57)(@vitest/ui@3.2.4)(jiti@2.4.2)(jsdom@26.1.0)(lightningcss@1.30.1)(tsx@4.19.4)(yaml@2.8.0)
vitest: 3.2.4(@types/debug@4.1.12)(@types/node@20.17.50)(@vitest/ui@3.2.4)(jiti@2.4.2)(jsdom@26.1.0)(lightningcss@1.30.1)(tsx@4.19.4)(yaml@2.8.0)
'@vitest/utils@3.2.4':
dependencies:

View file

@ -452,51 +452,4 @@ describe("CerebrasHandler", () => {
expect(toolCallChunks.length).toBe(0)
})
})
describe("mapToolChoice", () => {
it("should handle string tool choices", () => {
class TestCerebrasHandler extends CerebrasHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestCerebrasHandler(mockOptions)
expect(testHandler.testMapToolChoice("auto")).toBe("auto")
expect(testHandler.testMapToolChoice("none")).toBe("none")
expect(testHandler.testMapToolChoice("required")).toBe("required")
expect(testHandler.testMapToolChoice("unknown")).toBe("auto")
})
it("should handle object tool choice with function name", () => {
class TestCerebrasHandler extends CerebrasHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestCerebrasHandler(mockOptions)
const result = testHandler.testMapToolChoice({
type: "function",
function: { name: "my_tool" },
})
expect(result).toEqual({ type: "tool", toolName: "my_tool" })
})
it("should return undefined for null or undefined", () => {
class TestCerebrasHandler extends CerebrasHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestCerebrasHandler(mockOptions)
expect(testHandler.testMapToolChoice(null)).toBeUndefined()
expect(testHandler.testMapToolChoice(undefined)).toBeUndefined()
})
})
})

View file

@ -733,51 +733,4 @@ describe("DeepSeekHandler", () => {
expect(result).toBe(8192)
})
})
describe("mapToolChoice", () => {
it("should handle string tool choices", () => {
class TestDeepSeekHandler extends DeepSeekHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestDeepSeekHandler(mockOptions)
expect(testHandler.testMapToolChoice("auto")).toBe("auto")
expect(testHandler.testMapToolChoice("none")).toBe("none")
expect(testHandler.testMapToolChoice("required")).toBe("required")
expect(testHandler.testMapToolChoice("unknown")).toBe("auto")
})
it("should handle object tool choice with function name", () => {
class TestDeepSeekHandler extends DeepSeekHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestDeepSeekHandler(mockOptions)
const result = testHandler.testMapToolChoice({
type: "function",
function: { name: "my_tool" },
})
expect(result).toEqual({ type: "tool", toolName: "my_tool" })
})
it("should return undefined for null or undefined", () => {
class TestDeepSeekHandler extends DeepSeekHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestDeepSeekHandler(mockOptions)
expect(testHandler.testMapToolChoice(null)).toBeUndefined()
expect(testHandler.testMapToolChoice(undefined)).toBeUndefined()
})
})
})

File diff suppressed because it is too large Load diff

View file

@ -575,51 +575,4 @@ describe("GroqHandler", () => {
expect(result).toBe(customMaxTokens)
})
})
describe("mapToolChoice", () => {
it("should handle string tool choices", () => {
class TestGroqHandler extends GroqHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestGroqHandler(mockOptions)
expect(testHandler.testMapToolChoice("auto")).toBe("auto")
expect(testHandler.testMapToolChoice("none")).toBe("none")
expect(testHandler.testMapToolChoice("required")).toBe("required")
expect(testHandler.testMapToolChoice("unknown")).toBe("auto")
})
it("should handle object tool choice with function name", () => {
class TestGroqHandler extends GroqHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestGroqHandler(mockOptions)
const result = testHandler.testMapToolChoice({
type: "function",
function: { name: "my_tool" },
})
expect(result).toEqual({ type: "tool", toolName: "my_tool" })
})
it("should return undefined for null or undefined", () => {
class TestGroqHandler extends GroqHandler {
public testMapToolChoice(toolChoice: any) {
return this.mapToolChoice(toolChoice)
}
}
const testHandler = new TestGroqHandler(mockOptions)
expect(testHandler.testMapToolChoice(null)).toBeUndefined()
expect(testHandler.testMapToolChoice(undefined)).toBeUndefined()
})
})
})

View file

@ -6,7 +6,13 @@ import { cerebrasModels, cerebrasDefaultModelId, type CerebrasModelId, type Mode
import type { ApiHandlerOptions } from "../../shared/api"
import { convertToAiSdkMessages, convertToolsForAiSdk, processAiSdkStreamPart } from "../transform/ai-sdk"
import {
convertToAiSdkMessages,
convertToolsForAiSdk,
processAiSdkStreamPart,
mapToolChoice,
handleAiSdkError,
} from "../transform/ai-sdk"
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
import { getModelParams } from "../transform/model-params"
@ -75,40 +81,6 @@ export class CerebrasHandler extends BaseProvider implements SingleCompletionHan
}
}
/**
* Map OpenAI tool_choice to AI SDK toolChoice format.
*/
protected mapToolChoice(
toolChoice: any,
): "auto" | "none" | "required" | { type: "tool"; toolName: string } | undefined {
if (!toolChoice) {
return undefined
}
// Handle string values
if (typeof toolChoice === "string") {
switch (toolChoice) {
case "auto":
return "auto"
case "none":
return "none"
case "required":
return "required"
default:
return "auto"
}
}
// Handle object values (OpenAI ChatCompletionNamedToolChoice format)
if (typeof toolChoice === "object" && "type" in toolChoice) {
if (toolChoice.type === "function" && "function" in toolChoice && toolChoice.function?.name) {
return { type: "tool", toolName: toolChoice.function.name }
}
}
return undefined
}
/**
* Get the max tokens parameter to include in the request.
*/
@ -143,23 +115,28 @@ export class CerebrasHandler extends BaseProvider implements SingleCompletionHan
temperature: this.options.modelTemperature ?? temperature ?? CEREBRAS_DEFAULT_TEMPERATURE,
maxOutputTokens: this.getMaxOutputTokens(),
tools: aiSdkTools,
toolChoice: this.mapToolChoice(metadata?.tool_choice),
toolChoice: mapToolChoice(metadata?.tool_choice),
}
// Use streamText for streaming responses
const result = streamText(requestOptions)
// Process the full stream to get all events including reasoning
for await (const part of result.fullStream) {
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
try {
// Process the full stream to get all events including reasoning
for await (const part of result.fullStream) {
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
}
}
}
// Yield usage metrics at the end
const usage = await result.usage
if (usage) {
yield this.processUsageMetrics(usage)
// Yield usage metrics at the end
const usage = await result.usage
if (usage) {
yield this.processUsageMetrics(usage)
}
} catch (error) {
// Handle AI SDK errors (AI_RetryError, AI_APICallError, etc.)
throw handleAiSdkError(error, "Cerebras")
}
}

View file

@ -6,7 +6,13 @@ import { deepSeekModels, deepSeekDefaultModelId, DEEP_SEEK_DEFAULT_TEMPERATURE,
import type { ApiHandlerOptions } from "../../shared/api"
import { convertToAiSdkMessages, convertToolsForAiSdk, processAiSdkStreamPart } from "../transform/ai-sdk"
import {
convertToAiSdkMessages,
convertToolsForAiSdk,
processAiSdkStreamPart,
mapToolChoice,
handleAiSdkError,
} from "../transform/ai-sdk"
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
import { getModelParams } from "../transform/model-params"
@ -83,40 +89,6 @@ export class DeepSeekHandler extends BaseProvider implements SingleCompletionHan
}
}
/**
* Map OpenAI tool_choice to AI SDK toolChoice format.
*/
protected mapToolChoice(
toolChoice: any,
): "auto" | "none" | "required" | { type: "tool"; toolName: string } | undefined {
if (!toolChoice) {
return undefined
}
// Handle string values
if (typeof toolChoice === "string") {
switch (toolChoice) {
case "auto":
return "auto"
case "none":
return "none"
case "required":
return "required"
default:
return "auto"
}
}
// Handle object values (OpenAI ChatCompletionNamedToolChoice format)
if (typeof toolChoice === "object" && "type" in toolChoice) {
if (toolChoice.type === "function" && "function" in toolChoice && toolChoice.function?.name) {
return { type: "tool", toolName: toolChoice.function.name }
}
}
return undefined
}
/**
* Get the max tokens parameter to include in the request.
*/
@ -152,24 +124,29 @@ export class DeepSeekHandler extends BaseProvider implements SingleCompletionHan
temperature: this.options.modelTemperature ?? temperature ?? DEEP_SEEK_DEFAULT_TEMPERATURE,
maxOutputTokens: this.getMaxOutputTokens(),
tools: aiSdkTools,
toolChoice: this.mapToolChoice(metadata?.tool_choice),
toolChoice: mapToolChoice(metadata?.tool_choice),
}
// Use streamText for streaming responses
const result = streamText(requestOptions)
// Process the full stream to get all events including reasoning
for await (const part of result.fullStream) {
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
try {
// Process the full stream to get all events including reasoning
for await (const part of result.fullStream) {
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
}
}
}
// Yield usage metrics at the end, including cache metrics from providerMetadata
const usage = await result.usage
const providerMetadata = await result.providerMetadata
if (usage) {
yield this.processUsageMetrics(usage, providerMetadata as any)
// Yield usage metrics at the end, including cache metrics from providerMetadata
const usage = await result.usage
const providerMetadata = await result.providerMetadata
if (usage) {
yield this.processUsageMetrics(usage, providerMetadata as any)
}
} catch (error) {
// Handle AI SDK errors (AI_RetryError, AI_APICallError, etc.)
throw handleAiSdkError(error, "DeepSeek")
}
}

View file

@ -1,19 +1,175 @@
import { type FireworksModelId, fireworksDefaultModelId, fireworksModels } from "@roo-code/types"
import { Anthropic } from "@anthropic-ai/sdk"
import { createFireworks } from "@ai-sdk/fireworks"
import { streamText, generateText, ToolSet } from "ai"
import { fireworksModels, fireworksDefaultModelId, type ModelInfo } from "@roo-code/types"
import type { ApiHandlerOptions } from "../../shared/api"
import { BaseOpenAiCompatibleProvider } from "./base-openai-compatible-provider"
import {
convertToAiSdkMessages,
convertToolsForAiSdk,
processAiSdkStreamPart,
mapToolChoice,
handleAiSdkError,
} from "../transform/ai-sdk"
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
import { getModelParams } from "../transform/model-params"
import { DEFAULT_HEADERS } from "./constants"
import { BaseProvider } from "./base-provider"
import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index"
const FIREWORKS_DEFAULT_TEMPERATURE = 0.5
/**
* Fireworks provider using the dedicated @ai-sdk/fireworks package.
* Provides native support for various models including reasoning models.
*/
export class FireworksHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
protected provider: ReturnType<typeof createFireworks>
export class FireworksHandler extends BaseOpenAiCompatibleProvider<FireworksModelId> {
constructor(options: ApiHandlerOptions) {
super({
...options,
providerName: "Fireworks",
super()
this.options = options
// Create the Fireworks provider using AI SDK
this.provider = createFireworks({
baseURL: "https://api.fireworks.ai/inference/v1",
apiKey: options.fireworksApiKey,
defaultProviderModelId: fireworksDefaultModelId,
providerModels: fireworksModels,
defaultTemperature: 0.5,
apiKey: options.fireworksApiKey ?? "not-provided",
headers: DEFAULT_HEADERS,
})
}
override getModel(): { id: string; info: ModelInfo; maxTokens?: number; temperature?: number } {
const id = this.options.apiModelId ?? fireworksDefaultModelId
const info = fireworksModels[id as keyof typeof fireworksModels] || fireworksModels[fireworksDefaultModelId]
const params = getModelParams({
format: "openai",
modelId: id,
model: info,
settings: this.options,
defaultTemperature: FIREWORKS_DEFAULT_TEMPERATURE,
})
return { id, info, ...params }
}
/**
* Get the language model for the configured model ID.
*/
protected getLanguageModel() {
const { id } = this.getModel()
return this.provider(id)
}
/**
* Process usage metrics from the AI SDK response.
*/
protected processUsageMetrics(
usage: {
inputTokens?: number
outputTokens?: number
details?: {
cachedInputTokens?: number
reasoningTokens?: number
}
},
providerMetadata?: {
fireworks?: {
promptCacheHitTokens?: number
promptCacheMissTokens?: number
}
},
): ApiStreamUsageChunk {
// Extract cache metrics from Fireworks' providerMetadata if available
const cacheReadTokens = providerMetadata?.fireworks?.promptCacheHitTokens ?? usage.details?.cachedInputTokens
const cacheWriteTokens = providerMetadata?.fireworks?.promptCacheMissTokens
return {
type: "usage",
inputTokens: usage.inputTokens || 0,
outputTokens: usage.outputTokens || 0,
cacheReadTokens,
cacheWriteTokens,
reasoningTokens: usage.details?.reasoningTokens,
}
}
/**
* Get the max tokens parameter to include in the request.
*/
protected getMaxOutputTokens(): number | undefined {
const { info } = this.getModel()
return this.options.modelMaxTokens || info.maxTokens || undefined
}
/**
* Create a message stream using the AI SDK.
*/
override async *createMessage(
systemPrompt: string,
messages: Anthropic.Messages.MessageParam[],
metadata?: ApiHandlerCreateMessageMetadata,
): ApiStream {
const { temperature } = this.getModel()
const languageModel = this.getLanguageModel()
// Convert messages to AI SDK format
const aiSdkMessages = convertToAiSdkMessages(messages)
// Convert tools to OpenAI format first, then to AI SDK format
const openAiTools = this.convertToolsForOpenAI(metadata?.tools)
const aiSdkTools = convertToolsForAiSdk(openAiTools) as ToolSet | undefined
// Build the request options
const requestOptions: Parameters<typeof streamText>[0] = {
model: languageModel,
system: systemPrompt,
messages: aiSdkMessages,
temperature: this.options.modelTemperature ?? temperature ?? FIREWORKS_DEFAULT_TEMPERATURE,
maxOutputTokens: this.getMaxOutputTokens(),
tools: aiSdkTools,
toolChoice: mapToolChoice(metadata?.tool_choice),
}
// Use streamText for streaming responses
const result = streamText(requestOptions)
try {
// Process the full stream to get all events including reasoning
for await (const part of result.fullStream) {
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
}
}
// Yield usage metrics at the end, including cache metrics from providerMetadata
const usage = await result.usage
const providerMetadata = await result.providerMetadata
if (usage) {
yield this.processUsageMetrics(usage, providerMetadata as any)
}
} catch (error) {
// Handle AI SDK errors (AI_RetryError, AI_APICallError, etc.)
throw handleAiSdkError(error, "Fireworks")
}
}
/**
* Complete a prompt using the AI SDK generateText.
*/
async completePrompt(prompt: string): Promise<string> {
const { temperature } = this.getModel()
const languageModel = this.getLanguageModel()
const { text } = await generateText({
model: languageModel,
prompt,
maxOutputTokens: this.getMaxOutputTokens(),
temperature: this.options.modelTemperature ?? temperature ?? FIREWORKS_DEFAULT_TEMPERATURE,
})
return text
}
}

View file

@ -6,7 +6,13 @@ import { groqModels, groqDefaultModelId, type ModelInfo } from "@roo-code/types"
import type { ApiHandlerOptions } from "../../shared/api"
import { convertToAiSdkMessages, convertToolsForAiSdk, processAiSdkStreamPart } from "../transform/ai-sdk"
import {
convertToAiSdkMessages,
convertToolsForAiSdk,
processAiSdkStreamPart,
mapToolChoice,
handleAiSdkError,
} from "../transform/ai-sdk"
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
import { getModelParams } from "../transform/model-params"
@ -91,40 +97,6 @@ export class GroqHandler extends BaseProvider implements SingleCompletionHandler
}
}
/**
* Map OpenAI tool_choice to AI SDK toolChoice format.
*/
protected mapToolChoice(
toolChoice: any,
): "auto" | "none" | "required" | { type: "tool"; toolName: string } | undefined {
if (!toolChoice) {
return undefined
}
// Handle string values
if (typeof toolChoice === "string") {
switch (toolChoice) {
case "auto":
return "auto"
case "none":
return "none"
case "required":
return "required"
default:
return "auto"
}
}
// Handle object values (OpenAI ChatCompletionNamedToolChoice format)
if (typeof toolChoice === "object" && "type" in toolChoice) {
if (toolChoice.type === "function" && "function" in toolChoice && toolChoice.function?.name) {
return { type: "tool", toolName: toolChoice.function.name }
}
}
return undefined
}
/**
* Get the max tokens parameter to include in the request.
*/
@ -160,24 +132,29 @@ export class GroqHandler extends BaseProvider implements SingleCompletionHandler
temperature: this.options.modelTemperature ?? temperature ?? GROQ_DEFAULT_TEMPERATURE,
maxOutputTokens: this.getMaxOutputTokens(),
tools: aiSdkTools,
toolChoice: this.mapToolChoice(metadata?.tool_choice),
toolChoice: mapToolChoice(metadata?.tool_choice),
}
// Use streamText for streaming responses
const result = streamText(requestOptions)
// Process the full stream to get all events including reasoning
for await (const part of result.fullStream) {
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
try {
// Process the full stream to get all events including reasoning
for await (const part of result.fullStream) {
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
}
}
}
// Yield usage metrics at the end, including cache metrics from providerMetadata
const usage = await result.usage
const providerMetadata = await result.providerMetadata
if (usage) {
yield this.processUsageMetrics(usage, providerMetadata as any)
// Yield usage metrics at the end, including cache metrics from providerMetadata
const usage = await result.usage
const providerMetadata = await result.providerMetadata
if (usage) {
yield this.processUsageMetrics(usage, providerMetadata as any)
}
} catch (error) {
// Handle AI SDK errors (AI_RetryError, AI_APICallError, etc.)
throw handleAiSdkError(error, "Groq")
}
}

View file

@ -12,7 +12,13 @@ import type { ModelInfo } from "@roo-code/types"
import type { ApiHandlerOptions } from "../../shared/api"
import { convertToAiSdkMessages, convertToolsForAiSdk, processAiSdkStreamPart } from "../transform/ai-sdk"
import {
convertToAiSdkMessages,
convertToolsForAiSdk,
processAiSdkStreamPart,
mapToolChoice,
handleAiSdkError,
} from "../transform/ai-sdk"
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
import { DEFAULT_HEADERS } from "./constants"
@ -103,40 +109,6 @@ export abstract class OpenAICompatibleHandler extends BaseProvider implements Si
}
}
/**
* Map OpenAI tool_choice to AI SDK toolChoice format.
*/
protected mapToolChoice(
toolChoice: OpenAI.Chat.ChatCompletionCreateParams["tool_choice"],
): "auto" | "none" | "required" | { type: "tool"; toolName: string } | undefined {
if (!toolChoice) {
return undefined
}
// Handle string values
if (typeof toolChoice === "string") {
switch (toolChoice) {
case "auto":
return "auto"
case "none":
return "none"
case "required":
return "required"
default:
return "auto"
}
}
// Handle object values (OpenAI ChatCompletionNamedToolChoice format)
if (typeof toolChoice === "object" && "type" in toolChoice) {
if (toolChoice.type === "function" && "function" in toolChoice && toolChoice.function?.name) {
return { type: "tool", toolName: toolChoice.function.name }
}
}
return undefined
}
/**
* Get the max tokens parameter to include in the request.
*/
@ -173,24 +145,29 @@ export abstract class OpenAICompatibleHandler extends BaseProvider implements Si
temperature: model.temperature ?? this.config.temperature ?? 0,
maxOutputTokens: this.getMaxOutputTokens(),
tools: aiSdkTools,
toolChoice: this.mapToolChoice(metadata?.tool_choice),
toolChoice: mapToolChoice(metadata?.tool_choice),
}
// Use streamText for streaming responses
const result = streamText(requestOptions)
// Process the full stream to get all events
for await (const part of result.fullStream) {
// Use the processAiSdkStreamPart utility to convert stream parts
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
try {
// Process the full stream to get all events
for await (const part of result.fullStream) {
// Use the processAiSdkStreamPart utility to convert stream parts
for (const chunk of processAiSdkStreamPart(part)) {
yield chunk
}
}
}
// Yield usage metrics at the end
const usage = await result.usage
if (usage) {
yield this.processUsageMetrics(usage)
// Yield usage metrics at the end
const usage = await result.usage
if (usage) {
yield this.processUsageMetrics(usage)
}
} catch (error) {
// Handle AI SDK errors (AI_RetryError, AI_APICallError, etc.)
throw handleAiSdkError(error, this.config.providerName)
}
}

View file

@ -1,6 +1,13 @@
import { Anthropic } from "@anthropic-ai/sdk"
import OpenAI from "openai"
import { convertToAiSdkMessages, convertToolsForAiSdk, processAiSdkStreamPart } from "../ai-sdk"
import {
convertToAiSdkMessages,
convertToolsForAiSdk,
processAiSdkStreamPart,
mapToolChoice,
extractAiSdkErrorMessage,
handleAiSdkError,
} from "../ai-sdk"
vitest.mock("ai", () => ({
tool: vitest.fn((t) => t),
@ -486,4 +493,155 @@ describe("AI SDK conversion utilities", () => {
}
})
})
describe("mapToolChoice", () => {
it("should return undefined for null or undefined", () => {
expect(mapToolChoice(null)).toBeUndefined()
expect(mapToolChoice(undefined)).toBeUndefined()
})
it("should handle string tool choices", () => {
expect(mapToolChoice("auto")).toBe("auto")
expect(mapToolChoice("none")).toBe("none")
expect(mapToolChoice("required")).toBe("required")
})
it("should return auto for unknown string values", () => {
expect(mapToolChoice("unknown")).toBe("auto")
expect(mapToolChoice("invalid")).toBe("auto")
})
it("should handle object tool choice with function name", () => {
const result = mapToolChoice({
type: "function",
function: { name: "my_tool" },
})
expect(result).toEqual({ type: "tool", toolName: "my_tool" })
})
it("should return undefined for object without function name", () => {
const result = mapToolChoice({
type: "function",
function: {},
})
expect(result).toBeUndefined()
})
it("should return undefined for object with non-function type", () => {
const result = mapToolChoice({
type: "other",
function: { name: "my_tool" },
})
expect(result).toBeUndefined()
})
})
describe("extractAiSdkErrorMessage", () => {
it("should return 'Unknown error' for null/undefined", () => {
expect(extractAiSdkErrorMessage(null)).toBe("Unknown error")
expect(extractAiSdkErrorMessage(undefined)).toBe("Unknown error")
})
it("should extract message from AI_RetryError", () => {
const retryError = {
name: "AI_RetryError",
message: "Failed after 3 attempts",
errors: [new Error("Error 1"), new Error("Error 2"), new Error("Too Many Requests")],
lastError: { message: "Too Many Requests", status: 429 },
}
const result = extractAiSdkErrorMessage(retryError)
expect(result).toBe("Failed after 3 attempts (429): Too Many Requests")
})
it("should handle AI_RetryError without status", () => {
const retryError = {
name: "AI_RetryError",
message: "Failed after 2 attempts",
errors: [new Error("Error 1"), new Error("Connection failed")],
lastError: { message: "Connection failed" },
}
const result = extractAiSdkErrorMessage(retryError)
expect(result).toBe("Failed after 2 attempts: Connection failed")
})
it("should extract message from AI_APICallError", () => {
const apiError = {
name: "AI_APICallError",
message: "Rate limit exceeded",
status: 429,
}
const result = extractAiSdkErrorMessage(apiError)
expect(result).toBe("API Error (429): Rate limit exceeded")
})
it("should handle AI_APICallError without status", () => {
const apiError = {
name: "AI_APICallError",
message: "Connection timeout",
}
const result = extractAiSdkErrorMessage(apiError)
expect(result).toBe("Connection timeout")
})
it("should extract message from standard Error", () => {
const error = new Error("Something went wrong")
expect(extractAiSdkErrorMessage(error)).toBe("Something went wrong")
})
it("should convert non-Error to string", () => {
expect(extractAiSdkErrorMessage("string error")).toBe("string error")
expect(extractAiSdkErrorMessage({ custom: "object" })).toBe("[object Object]")
})
})
describe("handleAiSdkError", () => {
it("should wrap error with provider name", () => {
const error = new Error("API Error")
const result = handleAiSdkError(error, "Fireworks")
expect(result.message).toBe("Fireworks: API Error")
})
it("should preserve status code from AI_RetryError", () => {
const retryError = {
name: "AI_RetryError",
errors: [new Error("Too Many Requests")],
lastError: { message: "Too Many Requests", status: 429 },
}
const result = handleAiSdkError(retryError, "Groq")
expect(result.message).toContain("Groq:")
expect(result.message).toContain("429")
expect((result as any).status).toBe(429)
})
it("should preserve status code from AI_APICallError", () => {
const apiError = {
name: "AI_APICallError",
message: "Unauthorized",
status: 401,
}
const result = handleAiSdkError(apiError, "DeepSeek")
expect(result.message).toContain("DeepSeek:")
expect(result.message).toContain("401")
expect((result as any).status).toBe(401)
})
it("should preserve original error as cause", () => {
const originalError = new Error("Original error")
const result = handleAiSdkError(originalError, "Cerebras")
expect((result as any).cause).toBe(originalError)
})
})
})

View file

@ -273,3 +273,125 @@ export function* processAiSdkStreamPart(part: ExtendedStreamPart): Generator<Api
break
}
}
/**
* Type for AI SDK tool choice format.
*/
export type AiSdkToolChoice = "auto" | "none" | "required" | { type: "tool"; toolName: string } | undefined
/**
* Map OpenAI-style tool_choice to AI SDK toolChoice format.
* This is a shared utility to avoid duplication across providers.
*
* @param toolChoice - OpenAI-style tool choice (string or object)
* @returns AI SDK toolChoice format
*/
export function mapToolChoice(toolChoice: any): AiSdkToolChoice {
if (!toolChoice) {
return undefined
}
// Handle string values
if (typeof toolChoice === "string") {
switch (toolChoice) {
case "auto":
return "auto"
case "none":
return "none"
case "required":
return "required"
default:
return "auto"
}
}
// Handle object values (OpenAI ChatCompletionNamedToolChoice format)
if (typeof toolChoice === "object" && "type" in toolChoice) {
if (toolChoice.type === "function" && "function" in toolChoice && toolChoice.function?.name) {
return { type: "tool", toolName: toolChoice.function.name }
}
}
return undefined
}
/**
* Extract a user-friendly error message from AI SDK errors.
* The AI SDK wraps errors in types like AI_RetryError and AI_APICallError
* which need to be unwrapped to get the actual error message.
*
* @param error - The error to extract the message from
* @returns A user-friendly error message
*/
export function extractAiSdkErrorMessage(error: unknown): string {
if (!error) {
return "Unknown error"
}
// Cast to access AI SDK error properties
const anyError = error as any
// AI_RetryError has a lastError property with the actual error
if (anyError.name === "AI_RetryError") {
const retryCount = anyError.errors?.length || 0
const lastError = anyError.lastError
const lastErrorMessage = lastError?.message || lastError?.toString() || "Unknown error"
// Extract status code if available
const statusCode =
lastError?.status || lastError?.statusCode || anyError.status || anyError.statusCode || undefined
if (statusCode) {
return `Failed after ${retryCount} attempts (${statusCode}): ${lastErrorMessage}`
}
return `Failed after ${retryCount} attempts: ${lastErrorMessage}`
}
// AI_APICallError has message and optional status
if (anyError.name === "AI_APICallError") {
const statusCode = anyError.status || anyError.statusCode
if (statusCode) {
return `API Error (${statusCode}): ${anyError.message}`
}
return anyError.message || "API call failed"
}
// Standard Error
if (error instanceof Error) {
return error.message
}
// Fallback for non-Error objects
return String(error)
}
/**
* Handle AI SDK errors by extracting the message and preserving status codes.
* Returns an Error object with proper status preserved for retry logic.
*
* @param error - The AI SDK error to handle
* @param providerName - The name of the provider for context
* @returns An Error with preserved status code
*/
export function handleAiSdkError(error: unknown, providerName: string): Error {
const message = extractAiSdkErrorMessage(error)
const wrappedError = new Error(`${providerName}: ${message}`)
// Preserve status code for retry logic
const anyError = error as any
const statusCode =
anyError?.lastError?.status ||
anyError?.lastError?.statusCode ||
anyError?.status ||
anyError?.statusCode ||
undefined
if (statusCode) {
;(wrappedError as any).status = statusCode
}
// Preserve the original error for debugging
;(wrappedError as any).cause = error
return wrappedError
}

View file

@ -452,6 +452,7 @@
"dependencies": {
"@ai-sdk/cerebras": "^1.0.0",
"@ai-sdk/deepseek": "^2.0.14",
"@ai-sdk/fireworks": "^2.0.26",
"@ai-sdk/groq": "^3.0.19",
"@anthropic-ai/bedrock-sdk": "^0.10.2",
"@anthropic-ai/sdk": "^0.37.0",