Roo-Code/src/api/providers/ollama.ts
roomote[bot] 65146b1b12
fix: add error transform to cryptic openAI SDK errors when API key is invalid (#7586)
Co-authored-by: Roo Code <roomote@roocode.com>
Co-authored-by: Daniel Riccio <ricciodaniel98@gmail.com>
2025-09-04 23:12:28 -04:00

137 lines
4 KiB
TypeScript

import { Anthropic } from "@anthropic-ai/sdk"
import OpenAI from "openai"
import { type ModelInfo, openAiModelInfoSaneDefaults, DEEP_SEEK_DEFAULT_TEMPERATURE } from "@roo-code/types"
import type { ApiHandlerOptions } from "../../shared/api"
import { XmlMatcher } from "../../utils/xml-matcher"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { convertToR1Format } from "../transform/r1-format"
import { ApiStream } from "../transform/stream"
import { BaseProvider } from "./base-provider"
import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index"
import { getApiRequestTimeout } from "./utils/timeout-config"
import { handleOpenAIError } from "./utils/openai-error-handler"
type CompletionUsage = OpenAI.Chat.Completions.ChatCompletionChunk["usage"]
export class OllamaHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: OpenAI
private readonly providerName = "Ollama"
constructor(options: ApiHandlerOptions) {
super()
this.options = options
// Use the API key if provided (for Ollama cloud or authenticated instances)
// Otherwise use "ollama" as a placeholder for local instances
const apiKey = this.options.ollamaApiKey || "ollama"
const headers: Record<string, string> = {}
if (this.options.ollamaApiKey) {
headers["Authorization"] = `Bearer ${this.options.ollamaApiKey}`
}
this.client = new OpenAI({
baseURL: (this.options.ollamaBaseUrl || "http://localhost:11434") + "/v1",
apiKey: apiKey,
timeout: getApiRequestTimeout(),
defaultHeaders: headers,
})
}
override async *createMessage(
systemPrompt: string,
messages: Anthropic.Messages.MessageParam[],
metadata?: ApiHandlerCreateMessageMetadata,
): ApiStream {
const modelId = this.getModel().id
const useR1Format = modelId.toLowerCase().includes("deepseek-r1")
const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
{ role: "system", content: systemPrompt },
...(useR1Format ? convertToR1Format(messages) : convertToOpenAiMessages(messages)),
]
let stream
try {
stream = await this.client.chat.completions.create({
model: this.getModel().id,
messages: openAiMessages,
temperature: this.options.modelTemperature ?? 0,
stream: true,
stream_options: { include_usage: true },
})
} catch (error) {
throw handleOpenAIError(error, this.providerName)
}
const matcher = new XmlMatcher(
"think",
(chunk) =>
({
type: chunk.matched ? "reasoning" : "text",
text: chunk.data,
}) as const,
)
let lastUsage: CompletionUsage | undefined
for await (const chunk of stream) {
const delta = chunk.choices[0]?.delta
if (delta?.content) {
for (const matcherChunk of matcher.update(delta.content)) {
yield matcherChunk
}
}
if (chunk.usage) {
lastUsage = chunk.usage
}
}
for (const chunk of matcher.final()) {
yield chunk
}
if (lastUsage) {
yield {
type: "usage",
inputTokens: lastUsage?.prompt_tokens || 0,
outputTokens: lastUsage?.completion_tokens || 0,
}
}
}
override getModel(): { id: string; info: ModelInfo } {
return {
id: this.options.ollamaModelId || "",
info: openAiModelInfoSaneDefaults,
}
}
async completePrompt(prompt: string): Promise<string> {
try {
const modelId = this.getModel().id
const useR1Format = modelId.toLowerCase().includes("deepseek-r1")
let response
try {
response = await this.client.chat.completions.create({
model: this.getModel().id,
messages: useR1Format
? convertToR1Format([{ role: "user", content: prompt }])
: [{ role: "user", content: prompt }],
temperature: this.options.modelTemperature ?? (useR1Format ? DEEP_SEEK_DEFAULT_TEMPERATURE : 0),
stream: false,
})
} catch (error) {
throw handleOpenAIError(error, this.providerName)
}
return response.choices[0]?.message.content || ""
} catch (error) {
if (error instanceof Error) {
throw new Error(`Ollama completion error: ${error.message}`)
}
throw error
}
}
}