mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-09-17 23:51:08 +00:00
187 lines
5.9 KiB
TypeScript
187 lines
5.9 KiB
TypeScript
import { Anthropic } from "@anthropic-ai/sdk"
|
|
import { createDeepSeek } from "@ai-sdk/deepseek"
|
|
import { streamText, generateText, ToolSet, ModelMessage } from "ai"
|
|
|
|
import { deepSeekModels, deepSeekDefaultModelId, DEEP_SEEK_DEFAULT_TEMPERATURE, type ModelInfo } from "@roo-code/types"
|
|
|
|
import type { ApiHandlerOptions } from "../../shared/api"
|
|
|
|
import {
|
|
convertToAiSdkMessages,
|
|
convertToolsForAiSdk,
|
|
consumeAiSdkStream,
|
|
mapToolChoice,
|
|
handleAiSdkError,
|
|
} from "../transform/ai-sdk"
|
|
import { applyPromptCacheToMessages, mergeProviderOptions } from "../transform/prompt-cache"
|
|
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
|
|
import { getModelParams } from "../transform/model-params"
|
|
import { normalizeProviderUsage } from "./utils/normalize-provider-usage"
|
|
|
|
import { DEFAULT_HEADERS } from "./constants"
|
|
import { BaseProvider } from "./base-provider"
|
|
import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index"
|
|
import type { RooMessage } from "../../core/task-persistence/rooMessage"
|
|
|
|
/**
|
|
* DeepSeek provider using the dedicated @ai-sdk/deepseek package.
|
|
* Provides native support for reasoning (deepseek-reasoner) and prompt caching.
|
|
*/
|
|
export class DeepSeekHandler extends BaseProvider implements SingleCompletionHandler {
|
|
protected options: ApiHandlerOptions
|
|
protected provider: ReturnType<typeof createDeepSeek>
|
|
|
|
constructor(options: ApiHandlerOptions) {
|
|
super()
|
|
this.options = options
|
|
|
|
// Create the DeepSeek provider using AI SDK
|
|
this.provider = createDeepSeek({
|
|
baseURL: options.deepSeekBaseUrl || "https://api.deepseek.com/v1",
|
|
apiKey: options.deepSeekApiKey ?? "not-provided",
|
|
headers: DEFAULT_HEADERS,
|
|
})
|
|
}
|
|
|
|
override getModel(): { id: string; info: ModelInfo; maxTokens?: number; temperature?: number } {
|
|
const id = this.options.apiModelId ?? deepSeekDefaultModelId
|
|
const info = deepSeekModels[id as keyof typeof deepSeekModels] || deepSeekModels[deepSeekDefaultModelId]
|
|
const params = getModelParams({
|
|
format: "openai",
|
|
modelId: id,
|
|
model: info,
|
|
settings: this.options,
|
|
defaultTemperature: DEEP_SEEK_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, including DeepSeek's cache metrics.
|
|
* DeepSeek provides cache hit/miss info via providerMetadata.
|
|
*/
|
|
protected processUsageMetrics(
|
|
usage: {
|
|
inputTokens?: number
|
|
outputTokens?: number
|
|
details?: {
|
|
cachedInputTokens?: number
|
|
reasoningTokens?: number
|
|
}
|
|
},
|
|
providerMetadata?: {
|
|
deepseek?: {
|
|
promptCacheHitTokens?: number
|
|
promptCacheMissTokens?: number
|
|
}
|
|
},
|
|
): ApiStreamUsageChunk {
|
|
const { chunk } = normalizeProviderUsage({
|
|
provider: "deepseek",
|
|
apiProtocol: "openai",
|
|
usage: usage as any,
|
|
providerMetadata: providerMetadata as Record<string, unknown> | undefined,
|
|
modelInfo: this.getModel().info,
|
|
})
|
|
return chunk
|
|
}
|
|
|
|
/**
|
|
* 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.
|
|
* The AI SDK automatically handles reasoning for deepseek-reasoner model.
|
|
*/
|
|
override async *createMessage(
|
|
systemPrompt: string,
|
|
messages: RooMessage[],
|
|
metadata?: ApiHandlerCreateMessageMetadata,
|
|
): ApiStream {
|
|
const { temperature, info } = this.getModel()
|
|
const languageModel = this.getLanguageModel()
|
|
|
|
// Convert messages to AI SDK format
|
|
const aiSdkMessages = messages as ModelMessage[]
|
|
|
|
const promptCache = applyPromptCacheToMessages({
|
|
adapter: "ai-sdk",
|
|
overrideKey: "deepseek",
|
|
messages: aiSdkMessages,
|
|
modelInfo: {
|
|
supportsPromptCache: info.supportsPromptCache,
|
|
promptCacheRetention: info.promptCacheRetention,
|
|
},
|
|
settings: this.options,
|
|
})
|
|
|
|
// Convert tools to OpenAI format first, then to AI SDK format
|
|
const openAiTools = this.convertToolsForOpenAI(metadata?.tools)
|
|
const aiSdkTools = convertToolsForAiSdk(openAiTools, {
|
|
functionToolProviderOptions: promptCache.toolProviderOptions,
|
|
}) as ToolSet | undefined
|
|
|
|
const providerOptions = mergeProviderOptions(undefined, promptCache.providerOptionsPatch)
|
|
|
|
// Build the request options
|
|
const requestOptions: Parameters<typeof streamText>[0] = {
|
|
model: languageModel,
|
|
system: promptCache.systemProviderOptions
|
|
? ({ role: "system", content: systemPrompt, providerOptions: promptCache.systemProviderOptions } as any)
|
|
: systemPrompt,
|
|
messages: aiSdkMessages,
|
|
temperature: this.options.modelTemperature ?? temperature ?? DEEP_SEEK_DEFAULT_TEMPERATURE,
|
|
maxOutputTokens: this.getMaxOutputTokens(),
|
|
tools: aiSdkTools,
|
|
toolChoice: mapToolChoice(metadata?.tool_choice),
|
|
...(providerOptions ? ({ providerOptions } as Record<string, unknown>) : {}),
|
|
}
|
|
|
|
// Use streamText for streaming responses
|
|
const result = streamText(requestOptions)
|
|
|
|
try {
|
|
const processUsage = this.processUsageMetrics.bind(this)
|
|
yield* consumeAiSdkStream(result, async function* () {
|
|
const [usage, providerMetadata] = await Promise.all([result.usage, result.providerMetadata])
|
|
yield processUsage(usage, providerMetadata as Parameters<typeof processUsage>[1])
|
|
})
|
|
} catch (error) {
|
|
throw handleAiSdkError(error, "DeepSeek")
|
|
}
|
|
}
|
|
|
|
/**
|
|
* 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 ?? DEEP_SEEK_DEFAULT_TEMPERATURE,
|
|
})
|
|
|
|
return text
|
|
}
|
|
|
|
override isAiSdkProvider(): boolean {
|
|
return true
|
|
}
|
|
}
|