diff --git a/packages/types/src/provider-settings.ts b/packages/types/src/provider-settings.ts index dc51188df9..db919e224b 100644 --- a/packages/types/src/provider-settings.ts +++ b/packages/types/src/provider-settings.ts @@ -36,6 +36,7 @@ export const providerNames = [ "huggingface", "cerebras", "sambanova", + "tars", "zai", "fireworks", ] as const @@ -268,6 +269,12 @@ const fireworksSchema = apiModelIdProviderModelSchema.extend({ fireworksApiKey: z.string().optional(), }) +const tarsSchema = baseProviderSettingsSchema.extend({ + tarsApiKey: z.string().optional(), + tarsModelId: z.string().optional(), + tarsBaseUrl: z.string().optional(), +}) + const defaultSchema = z.object({ apiProvider: z.undefined(), }) @@ -301,6 +308,7 @@ export const providerSettingsSchemaDiscriminated = z.discriminatedUnion("apiProv litellmSchema.merge(z.object({ apiProvider: z.literal("litellm") })), cerebrasSchema.merge(z.object({ apiProvider: z.literal("cerebras") })), sambaNovaSchema.merge(z.object({ apiProvider: z.literal("sambanova") })), + tarsSchema.merge(z.object({ apiProvider: z.literal("tars") })), zaiSchema.merge(z.object({ apiProvider: z.literal("zai") })), fireworksSchema.merge(z.object({ apiProvider: z.literal("fireworks") })), defaultSchema, @@ -336,6 +344,7 @@ export const providerSettingsSchema = z.object({ ...litellmSchema.shape, ...cerebrasSchema.shape, ...sambaNovaSchema.shape, + ...tarsSchema.shape, ...zaiSchema.shape, ...fireworksSchema.shape, ...codebaseIndexProviderSchema.shape, @@ -363,6 +372,7 @@ export const MODEL_ID_KEYS: Partial[] = [ "requestyModelId", "litellmModelId", "huggingFaceModelId", + "tarsModelId", ] export const getModelId = (settings: ProviderSettings): string | undefined => { diff --git a/packages/types/src/providers/index.ts b/packages/types/src/providers/index.ts index 0ab27ea3dc..bd55a4f3e8 100644 --- a/packages/types/src/providers/index.ts +++ b/packages/types/src/providers/index.ts @@ -17,6 +17,7 @@ export * from "./openai.js" export * from "./openrouter.js" export * from "./requesty.js" export * from "./sambanova.js" +export * from "./tars.js" export * from "./unbound.js" export * from "./vertex.js" export * from "./vscode-llm.js" diff --git a/packages/types/src/providers/tars.ts b/packages/types/src/providers/tars.ts new file mode 100644 index 0000000000..85b178e3c1 --- /dev/null +++ b/packages/types/src/providers/tars.ts @@ -0,0 +1,30 @@ +import type { ModelInfo } from "../model.js" + +// TARS is a router service similar to OpenRouter, so we'll follow a similar pattern +export const tarsDefaultModelId = "anthropic/claude-3-5-sonnet-20241022" + +export const tarsDefaultModelInfo: ModelInfo = { + maxTokens: 8192, + contextWindow: 200_000, + supportsImages: true, + supportsComputerUse: false, + supportsPromptCache: true, + inputPrice: 3.0, + outputPrice: 15.0, + cacheWritesPrice: 3.75, + cacheReadsPrice: 0.3, + description: + "Claude 3.5 Sonnet delivers strong performance on tasks requiring visual reasoning, like interpreting charts, graphs, or diagrams. It's a versatile, balanced model that handles both text and image inputs effectively.", +} + +export const TARS_DEFAULT_PROVIDER_NAME = "[default]" + +// Models that support prompt caching through TARS +export const TARS_PROMPT_CACHING_MODELS = new Set([ + "anthropic/claude-3-haiku-20240307", + "anthropic/claude-3-opus-20240229", + "anthropic/claude-3-sonnet-20240229", + "anthropic/claude-3-5-haiku-20241022", + "anthropic/claude-3-5-sonnet-20240620", + "anthropic/claude-3-5-sonnet-20241022", +]) diff --git a/src/api/index.ts b/src/api/index.ts index 57b06f7bbd..64e95f9da0 100644 --- a/src/api/index.ts +++ b/src/api/index.ts @@ -32,6 +32,7 @@ import { LiteLLMHandler, ClaudeCodeHandler, SambaNovaHandler, + TarsHandler, DoubaoHandler, ZAiHandler, FireworksHandler, @@ -126,6 +127,8 @@ export function buildApiHandler(configuration: ProviderSettings): ApiHandler { return new CerebrasHandler(options) case "sambanova": return new SambaNovaHandler(options) + case "tars": + return new TarsHandler(options) case "zai": return new ZAiHandler(options) case "fireworks": diff --git a/src/api/providers/__tests__/tars.spec.ts b/src/api/providers/__tests__/tars.spec.ts new file mode 100644 index 0000000000..a40764dc67 --- /dev/null +++ b/src/api/providers/__tests__/tars.spec.ts @@ -0,0 +1,224 @@ +import { describe, it, expect, vitest, beforeEach } from "vitest" +import OpenAI from "openai" + +import { tarsDefaultModelId, tarsDefaultModelInfo } from "@roo-code/types" + +import { TarsHandler } from "../tars" +import { ApiHandlerOptions } from "../../../shared/api" + +// Mock OpenAI +vitest.mock("openai", () => { + const mockCreate = vitest.fn() + const mockChat = { + completions: { + create: mockCreate, + }, + } + const MockOpenAI = vitest.fn(() => ({ + chat: mockChat, + })) + return { default: MockOpenAI } +}) + +describe("TarsHandler", () => { + const mockOptions: ApiHandlerOptions = { + tarsApiKey: "test-key", + tarsModelId: "anthropic/claude-3-5-sonnet-20241022", + tarsBaseUrl: "https://api.tetrate.io/v1", + } + + beforeEach(() => { + vitest.clearAllMocks() + }) + + it("initializes with correct options", () => { + const handler = new TarsHandler(mockOptions) + expect(handler).toBeInstanceOf(TarsHandler) + + // Verify OpenAI client was initialized with correct parameters + expect(OpenAI).toHaveBeenCalledWith({ + baseURL: "https://api.tetrate.io/v1", + apiKey: "test-key", + defaultHeaders: expect.any(Object), + }) + }) + + it("uses default base URL when not provided", () => { + const handler = new TarsHandler({ tarsApiKey: "test-key" }) + expect(handler).toBeInstanceOf(TarsHandler) + + expect(OpenAI).toHaveBeenCalledWith({ + baseURL: "https://api.tetrate.io/v1", + apiKey: "test-key", + defaultHeaders: expect.any(Object), + }) + }) + + describe("getModel", () => { + it("returns correct model info when options are provided", () => { + const handler = new TarsHandler(mockOptions) + const result = handler.getModel() + + expect(result).toEqual({ + id: "anthropic/claude-3-5-sonnet-20241022", + info: tarsDefaultModelInfo, + }) + }) + + it("returns default model info when options are not provided", () => { + const handler = new TarsHandler({}) + const result = handler.getModel() + + expect(result).toEqual({ + id: tarsDefaultModelId, + info: tarsDefaultModelInfo, + }) + }) + }) + + describe("createMessage", () => { + it("generates correct stream chunks", async () => { + const mockStream = { + [Symbol.asyncIterator]: async function* () { + yield { + choices: [{ delta: { content: "Hello" } }], + usage: null, + } + yield { + choices: [{ delta: { content: " world" } }], + usage: { + prompt_tokens: 10, + completion_tokens: 2, + prompt_tokens_details: { cached_tokens: 5 }, + }, + } + }, + } + + const mockCreate = vitest.fn().mockResolvedValue(mockStream) + const mockOpenAI = vitest.fn(() => ({ + chat: { completions: { create: mockCreate } }, + })) + ;(OpenAI as any).mockImplementation(mockOpenAI) + + const handler = new TarsHandler(mockOptions) + + const chunks = [] + const generator = handler.createMessage("System prompt", [{ role: "user", content: "Hello" }]) + + for await (const chunk of generator) { + chunks.push(chunk) + } + + expect(chunks).toEqual([ + { type: "text", text: "Hello" }, + { type: "text", text: " world" }, + { + type: "usage", + inputTokens: 10, + outputTokens: 2, + cacheReadTokens: 5, + totalCost: 0, + }, + ]) + + // The messages will have cache control added + expect(mockCreate).toHaveBeenCalledWith({ + model: "anthropic/claude-3-5-sonnet-20241022", + max_tokens: 8192, + temperature: 0, + messages: expect.arrayContaining([ + expect.objectContaining({ role: "system" }), + expect.objectContaining({ role: "user" }), + ]), + stream: true, + stream_options: { include_usage: true }, + }) + }) + + it("adds cache control for supported models", async () => { + const mockStream = { + [Symbol.asyncIterator]: async function* () { + yield { + choices: [{ delta: { content: "test" } }], + usage: null, + } + }, + } + + const mockCreate = vitest.fn().mockResolvedValue(mockStream) + const mockOpenAI = vitest.fn(() => ({ + chat: { completions: { create: mockCreate } }, + })) + ;(OpenAI as any).mockImplementation(mockOpenAI) + + const handler = new TarsHandler({ + ...mockOptions, + tarsModelId: "anthropic/claude-3-5-sonnet-20241022", + }) + + const generator = handler.createMessage("System prompt", [{ role: "user", content: "Hello" }]) + + for await (const chunk of generator) { + // Consume the generator + } + + const call = mockCreate.mock.calls[0][0] + + // The cache breakpoints function should have been called + expect(call.messages.length).toBe(2) + expect(call.messages[0].role).toBe("system") + expect(call.messages[1].role).toBe("user") + + // Messages should have cache control structure + expect(call.messages[0].content).toEqual( + expect.arrayContaining([ + expect.objectContaining({ + type: "text", + text: "System prompt", + cache_control: expect.objectContaining({ type: "ephemeral" }), + }), + ]), + ) + }) + }) + + describe("completePrompt", () => { + it("returns correct response", async () => { + const mockResponse = { choices: [{ message: { content: "test completion" } }] } + + const mockCreate = vitest.fn().mockResolvedValue(mockResponse) + const mockOpenAI = vitest.fn(() => ({ + chat: { completions: { create: mockCreate } }, + })) + ;(OpenAI as any).mockImplementation(mockOpenAI) + + const handler = new TarsHandler(mockOptions) + const result = await handler.completePrompt("test prompt") + + expect(result).toBe("test completion") + expect(mockCreate).toHaveBeenCalledWith({ + model: "anthropic/claude-3-5-sonnet-20241022", + max_tokens: 8192, + temperature: 0, + messages: [{ role: "user", content: "test prompt" }], + stream: false, + }) + }) + + it("handles empty response", async () => { + const mockResponse = { choices: [{ message: { content: null } }] } + + const mockCreate = vitest.fn().mockResolvedValue(mockResponse) + const mockOpenAI = vitest.fn(() => ({ + chat: { completions: { create: mockCreate } }, + })) + ;(OpenAI as any).mockImplementation(mockOpenAI) + + const handler = new TarsHandler(mockOptions) + const result = await handler.completePrompt("test prompt") + + expect(result).toBe("") + }) + }) +}) diff --git a/src/api/providers/index.ts b/src/api/providers/index.ts index 890999aa25..9ecc744631 100644 --- a/src/api/providers/index.ts +++ b/src/api/providers/index.ts @@ -22,6 +22,7 @@ export { OpenAiHandler } from "./openai" export { OpenRouterHandler } from "./openrouter" export { RequestyHandler } from "./requesty" export { SambaNovaHandler } from "./sambanova" +export { TarsHandler } from "./tars" export { UnboundHandler } from "./unbound" export { VertexHandler } from "./vertex" export { VsCodeLmHandler } from "./vscode-lm" diff --git a/src/api/providers/tars.ts b/src/api/providers/tars.ts new file mode 100644 index 0000000000..6c999053ac --- /dev/null +++ b/src/api/providers/tars.ts @@ -0,0 +1,112 @@ +import { Anthropic } from "@anthropic-ai/sdk" +import OpenAI from "openai" + +import { tarsDefaultModelId, tarsDefaultModelInfo, TARS_PROMPT_CACHING_MODELS } from "@roo-code/types" + +import type { ApiHandlerOptions } from "../../shared/api" + +import { convertToOpenAiMessages } from "../transform/openai-format" +import { ApiStreamChunk } from "../transform/stream" +import { addCacheBreakpoints as addAnthropicCacheBreakpoints } from "../transform/caching/anthropic" + +import { DEFAULT_HEADERS } from "./constants" +import { BaseProvider } from "./base-provider" +import type { SingleCompletionHandler } from "../index" + +export class TarsHandler extends BaseProvider implements SingleCompletionHandler { + protected options: ApiHandlerOptions + private client: OpenAI + + constructor(options: ApiHandlerOptions) { + super() + this.options = options + + const baseURL = this.options.tarsBaseUrl || "https://api.tetrate.io/v1" + const apiKey = this.options.tarsApiKey ?? "not-provided" + + this.client = new OpenAI({ baseURL, apiKey, defaultHeaders: DEFAULT_HEADERS }) + } + + override async *createMessage( + systemPrompt: string, + messages: Anthropic.Messages.MessageParam[], + ): AsyncGenerator { + const model = this.getModel() + + const { id: modelId, info: modelInfo } = model + const maxTokens = + this.options.includeMaxTokens !== false && modelInfo.maxTokens ? modelInfo.maxTokens : undefined + const temperature = this.options.modelTemperature ?? 0 + + // Convert Anthropic messages to OpenAI format + const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [ + { role: "system", content: systemPrompt }, + ...convertToOpenAiMessages(messages), + ] + + // Add prompt caching for supported models + if (TARS_PROMPT_CACHING_MODELS.has(modelId)) { + addAnthropicCacheBreakpoints(systemPrompt, openAiMessages) + } + + const completionParams: OpenAI.Chat.ChatCompletionCreateParams = { + model: modelId, + ...(maxTokens && { max_tokens: maxTokens }), + temperature, + messages: openAiMessages, + stream: true, + stream_options: { include_usage: true }, + } + + const stream = await this.client.chat.completions.create(completionParams) + + let lastUsage: OpenAI.CompletionUsage | undefined = undefined + + for await (const chunk of stream) { + const delta = chunk.choices[0]?.delta + + if (delta?.content) { + yield { type: "text", text: delta.content } + } + + if (chunk.usage) { + lastUsage = chunk.usage + } + } + + if (lastUsage) { + yield { + type: "usage", + inputTokens: lastUsage.prompt_tokens || 0, + outputTokens: lastUsage.completion_tokens || 0, + cacheReadTokens: lastUsage.prompt_tokens_details?.cached_tokens, + totalCost: 0, // TARS doesn't provide cost information in the API response + } + } + } + + override getModel() { + const id = this.options.tarsModelId ?? tarsDefaultModelId + const info = tarsDefaultModelInfo + + return { id, info } + } + + async completePrompt(prompt: string) { + const model = this.getModel() + const { id: modelId, info: modelInfo } = model + const maxTokens = modelInfo.maxTokens + const temperature = this.options.modelTemperature ?? 0 + + const completionParams: OpenAI.Chat.ChatCompletionCreateParams = { + model: modelId, + max_tokens: maxTokens, + temperature, + messages: [{ role: "user", content: prompt }], + stream: false, + } + + const response = await this.client.chat.completions.create(completionParams) + return response.choices[0]?.message?.content || "" + } +} diff --git a/webview-ui/src/components/settings/ApiOptions.tsx b/webview-ui/src/components/settings/ApiOptions.tsx index 204abe9c0f..4e369a33af 100644 --- a/webview-ui/src/components/settings/ApiOptions.tsx +++ b/webview-ui/src/components/settings/ApiOptions.tsx @@ -31,6 +31,7 @@ import { internationalZAiDefaultModelId, mainlandZAiDefaultModelId, fireworksDefaultModelId, + tarsDefaultModelId, } from "@roo-code/types" import { vscode } from "@src/utils/vscode" @@ -78,6 +79,7 @@ import { OpenRouter, Requesty, SambaNova, + Tars, Unbound, Vertex, VSCodeLM, @@ -319,6 +321,7 @@ const ApiOptions = ({ : internationalZAiDefaultModelId, }, fireworks: { field: "apiModelId", default: fireworksDefaultModelId }, + tars: { field: "tarsModelId", default: tarsDefaultModelId }, openai: { field: "openAiModelId" }, ollama: { field: "ollamaModelId" }, lmstudio: { field: "lmStudioModelId" }, @@ -543,6 +546,10 @@ const ApiOptions = ({ )} + {selectedProvider === "tars" && ( + + )} + {selectedProvider === "zai" && ( )} diff --git a/webview-ui/src/components/settings/constants.ts b/webview-ui/src/components/settings/constants.ts index 90192f372b..6a24ef43bd 100644 --- a/webview-ui/src/components/settings/constants.ts +++ b/webview-ui/src/components/settings/constants.ts @@ -67,6 +67,7 @@ export const PROVIDERS = [ { value: "chutes", label: "Chutes AI" }, { value: "litellm", label: "LiteLLM" }, { value: "sambanova", label: "SambaNova" }, + { value: "tars", label: "TARS (Tetrate Agent Router Service)" }, { value: "zai", label: "Z AI" }, { value: "fireworks", label: "Fireworks AI" }, ].sort((a, b) => a.label.localeCompare(b.label)) diff --git a/webview-ui/src/components/settings/providers/Tars.tsx b/webview-ui/src/components/settings/providers/Tars.tsx new file mode 100644 index 0000000000..41195b8f93 --- /dev/null +++ b/webview-ui/src/components/settings/providers/Tars.tsx @@ -0,0 +1,57 @@ +import { useCallback } from "react" +import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react" + +import { type ProviderSettings, tarsDefaultModelId } from "@roo-code/types" + +import { useAppTranslation } from "@src/i18n/TranslationContext" + +import { inputEventTransform } from "../transforms" + +type TarsProps = { + apiConfiguration: ProviderSettings + setApiConfigurationField: (field: keyof ProviderSettings, value: ProviderSettings[keyof ProviderSettings]) => void +} + +export const Tars = ({ apiConfiguration, setApiConfigurationField }: TarsProps) => { + const { t } = useAppTranslation() + + const handleInputChange = useCallback( + ( + field: K, + transform: (event: E) => ProviderSettings[K] = inputEventTransform, + ) => + (event: E | Event) => { + setApiConfigurationField(field, transform(event as E)) + }, + [setApiConfigurationField], + ) + + return ( + <> + +
+ +
+
+
+ {t("settings:providers.apiKeyStorageNotice")} +
+
{t("settings:providers.tarsDescription")}
+ + + +
+ {t("settings:providers.tarsModelDescription")} +
+ + ) +} diff --git a/webview-ui/src/components/settings/providers/index.ts b/webview-ui/src/components/settings/providers/index.ts index e8428eb66c..6028cd9919 100644 --- a/webview-ui/src/components/settings/providers/index.ts +++ b/webview-ui/src/components/settings/providers/index.ts @@ -18,6 +18,7 @@ export { OpenAICompatible } from "./OpenAICompatible" export { OpenRouter } from "./OpenRouter" export { Requesty } from "./Requesty" export { SambaNova } from "./SambaNova" +export { Tars } from "./Tars" export { Unbound } from "./Unbound" export { Vertex } from "./Vertex" export { VSCodeLM } from "./VSCodeLM"