Roo-Code/src/api/providers/deepseek.ts

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
}
}