mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-08 03:07:53 +00:00
134 lines
4.1 KiB
TypeScript
134 lines
4.1 KiB
TypeScript
import { Anthropic } from "@anthropic-ai/sdk"
|
|
import OpenAI from "openai"
|
|
|
|
import {
|
|
vercelAiGatewayDefaultModelId,
|
|
vercelAiGatewayDefaultModelInfo,
|
|
VERCEL_AI_GATEWAY_DEFAULT_TEMPERATURE,
|
|
VERCEL_AI_GATEWAY_PROMPT_CACHING_MODELS,
|
|
} from "@roo-code/types"
|
|
|
|
import { ApiHandlerOptions } from "../../shared/api"
|
|
|
|
import { ApiStream } from "../transform/stream"
|
|
import { convertToOpenAiMessages } from "../transform/openai-format"
|
|
import { addCacheBreakpoints } from "../transform/caching/vercel-ai-gateway"
|
|
|
|
import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index"
|
|
import { RouterProvider } from "./router-provider"
|
|
|
|
// Extend OpenAI's CompletionUsage to include Vercel AI Gateway specific fields
|
|
interface VercelAiGatewayUsage extends OpenAI.CompletionUsage {
|
|
cache_creation_input_tokens?: number
|
|
cost?: number
|
|
}
|
|
|
|
export class VercelAiGatewayHandler extends RouterProvider implements SingleCompletionHandler {
|
|
constructor(options: ApiHandlerOptions) {
|
|
super({
|
|
options,
|
|
name: "vercel-ai-gateway",
|
|
baseURL: "https://ai-gateway.vercel.sh/v1",
|
|
apiKey: options.vercelAiGatewayApiKey,
|
|
modelId: options.vercelAiGatewayModelId,
|
|
defaultModelId: vercelAiGatewayDefaultModelId,
|
|
defaultModelInfo: vercelAiGatewayDefaultModelInfo,
|
|
})
|
|
}
|
|
|
|
override async *createMessage(
|
|
systemPrompt: string,
|
|
messages: Anthropic.Messages.MessageParam[],
|
|
metadata?: ApiHandlerCreateMessageMetadata,
|
|
): ApiStream {
|
|
const { id: modelId, info } = await this.fetchModel()
|
|
|
|
const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
|
|
{ role: "system", content: systemPrompt },
|
|
...convertToOpenAiMessages(messages),
|
|
]
|
|
|
|
if (VERCEL_AI_GATEWAY_PROMPT_CACHING_MODELS.has(modelId) && info.supportsPromptCache) {
|
|
addCacheBreakpoints(systemPrompt, openAiMessages)
|
|
}
|
|
|
|
const body: OpenAI.Chat.ChatCompletionCreateParams = {
|
|
model: modelId,
|
|
messages: openAiMessages,
|
|
temperature: this.supportsTemperature(modelId)
|
|
? (this.options.modelTemperature ?? VERCEL_AI_GATEWAY_DEFAULT_TEMPERATURE)
|
|
: undefined,
|
|
max_completion_tokens: info.maxTokens,
|
|
stream: true,
|
|
stream_options: { include_usage: true },
|
|
...(metadata?.tools && { tools: this.convertToolsForOpenAI(metadata.tools) }),
|
|
...(metadata?.tool_choice && { tool_choice: metadata.tool_choice }),
|
|
...(metadata?.toolProtocol === "native" && {
|
|
parallel_tool_calls: metadata.parallelToolCalls ?? false,
|
|
}),
|
|
}
|
|
|
|
const completion = await this.client.chat.completions.create(body)
|
|
|
|
for await (const chunk of completion) {
|
|
const delta = chunk.choices[0]?.delta
|
|
if (delta?.content) {
|
|
yield {
|
|
type: "text",
|
|
text: delta.content,
|
|
}
|
|
}
|
|
|
|
// Emit raw tool call chunks - NativeToolCallParser handles state management
|
|
if (delta?.tool_calls) {
|
|
for (const toolCall of delta.tool_calls) {
|
|
yield {
|
|
type: "tool_call_partial",
|
|
index: toolCall.index,
|
|
id: toolCall.id,
|
|
name: toolCall.function?.name,
|
|
arguments: toolCall.function?.arguments,
|
|
}
|
|
}
|
|
}
|
|
|
|
if (chunk.usage) {
|
|
const usage = chunk.usage as VercelAiGatewayUsage
|
|
yield {
|
|
type: "usage",
|
|
inputTokens: usage.prompt_tokens || 0,
|
|
outputTokens: usage.completion_tokens || 0,
|
|
cacheWriteTokens: usage.cache_creation_input_tokens || undefined,
|
|
cacheReadTokens: usage.prompt_tokens_details?.cached_tokens || undefined,
|
|
totalCost: usage.cost ?? 0,
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
async completePrompt(prompt: string): Promise<string> {
|
|
const { id: modelId, info } = await this.fetchModel()
|
|
|
|
try {
|
|
const requestOptions: OpenAI.Chat.ChatCompletionCreateParams = {
|
|
model: modelId,
|
|
messages: [{ role: "user", content: prompt }],
|
|
stream: false,
|
|
}
|
|
|
|
if (this.supportsTemperature(modelId)) {
|
|
requestOptions.temperature = this.options.modelTemperature ?? VERCEL_AI_GATEWAY_DEFAULT_TEMPERATURE
|
|
}
|
|
|
|
requestOptions.max_completion_tokens = info.maxTokens
|
|
|
|
const response = await this.client.chat.completions.create(requestOptions)
|
|
return response.choices[0]?.message.content || ""
|
|
} catch (error) {
|
|
if (error instanceof Error) {
|
|
throw new Error(`Vercel AI Gateway completion error: ${error.message}`)
|
|
}
|
|
throw error
|
|
}
|
|
}
|
|
}
|