mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-08-28 05:27:24 +00:00
roofactor: Migrate Roo provider to AI SDK (#11383)
This commit is contained in:
parent
08a96af22c
commit
8d57da8bc8
5 changed files with 590 additions and 1720 deletions
|
|
@ -1,119 +0,0 @@
|
|||
// npx vitest run api/providers/__tests__/base-openai-compatible-provider-timeout.spec.ts
|
||||
|
||||
import type { ModelInfo } from "@roo-code/types"
|
||||
|
||||
import { BaseOpenAiCompatibleProvider } from "../base-openai-compatible-provider"
|
||||
|
||||
// Mock the timeout config utility
|
||||
vitest.mock("../utils/timeout-config", () => ({
|
||||
getApiRequestTimeout: vitest.fn(),
|
||||
}))
|
||||
|
||||
import { getApiRequestTimeout } from "../utils/timeout-config"
|
||||
|
||||
// Mock OpenAI and capture constructor calls
|
||||
const mockOpenAIConstructor = vitest.fn()
|
||||
|
||||
vitest.mock("openai", () => {
|
||||
return {
|
||||
__esModule: true,
|
||||
default: vitest.fn().mockImplementation((config) => {
|
||||
mockOpenAIConstructor(config)
|
||||
return {
|
||||
chat: {
|
||||
completions: {
|
||||
create: vitest.fn(),
|
||||
},
|
||||
},
|
||||
}
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
// Create a concrete test implementation of the abstract base class
|
||||
class TestOpenAiCompatibleProvider extends BaseOpenAiCompatibleProvider<"test-model"> {
|
||||
constructor(apiKey: string) {
|
||||
const testModels: Record<"test-model", ModelInfo> = {
|
||||
"test-model": {
|
||||
maxTokens: 4096,
|
||||
contextWindow: 128000,
|
||||
supportsImages: false,
|
||||
supportsPromptCache: false,
|
||||
inputPrice: 0.5,
|
||||
outputPrice: 1.5,
|
||||
},
|
||||
}
|
||||
|
||||
super({
|
||||
providerName: "TestProvider",
|
||||
baseURL: "https://test.example.com/v1",
|
||||
defaultProviderModelId: "test-model",
|
||||
providerModels: testModels,
|
||||
apiKey,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
describe("BaseOpenAiCompatibleProvider Timeout Configuration", () => {
|
||||
beforeEach(() => {
|
||||
vitest.clearAllMocks()
|
||||
})
|
||||
|
||||
it("should call getApiRequestTimeout when creating the provider", () => {
|
||||
;(getApiRequestTimeout as any).mockReturnValue(600000)
|
||||
|
||||
new TestOpenAiCompatibleProvider("test-api-key")
|
||||
|
||||
expect(getApiRequestTimeout).toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("should pass the default timeout to the OpenAI client constructor", () => {
|
||||
;(getApiRequestTimeout as any).mockReturnValue(600000) // 600 seconds in ms
|
||||
|
||||
new TestOpenAiCompatibleProvider("test-api-key")
|
||||
|
||||
expect(mockOpenAIConstructor).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
baseURL: "https://test.example.com/v1",
|
||||
apiKey: "test-api-key",
|
||||
timeout: 600000,
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
it("should use custom timeout value from getApiRequestTimeout", () => {
|
||||
;(getApiRequestTimeout as any).mockReturnValue(1800000) // 30 minutes in ms
|
||||
|
||||
new TestOpenAiCompatibleProvider("test-api-key")
|
||||
|
||||
expect(mockOpenAIConstructor).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
timeout: 1800000,
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
it("should handle zero timeout (no timeout)", () => {
|
||||
;(getApiRequestTimeout as any).mockReturnValue(0)
|
||||
|
||||
new TestOpenAiCompatibleProvider("test-api-key")
|
||||
|
||||
expect(mockOpenAIConstructor).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
timeout: 0,
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
it("should pass DEFAULT_HEADERS to the OpenAI client constructor", () => {
|
||||
;(getApiRequestTimeout as any).mockReturnValue(600000)
|
||||
|
||||
new TestOpenAiCompatibleProvider("test-api-key")
|
||||
|
||||
expect(mockOpenAIConstructor).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
defaultHeaders: expect.any(Object),
|
||||
}),
|
||||
)
|
||||
})
|
||||
})
|
||||
|
|
@ -1,548 +0,0 @@
|
|||
// npx vitest run api/providers/__tests__/base-openai-compatible-provider.spec.ts
|
||||
|
||||
import { Anthropic } from "@anthropic-ai/sdk"
|
||||
import OpenAI from "openai"
|
||||
|
||||
import type { ModelInfo } from "@roo-code/types"
|
||||
|
||||
import { BaseOpenAiCompatibleProvider } from "../base-openai-compatible-provider"
|
||||
|
||||
// Create mock functions
|
||||
const mockCreate = vi.fn()
|
||||
|
||||
// Mock OpenAI module
|
||||
vi.mock("openai", () => ({
|
||||
default: vi.fn(() => ({
|
||||
chat: {
|
||||
completions: {
|
||||
create: mockCreate,
|
||||
},
|
||||
},
|
||||
})),
|
||||
}))
|
||||
|
||||
// Create a concrete test implementation of the abstract base class
|
||||
class TestOpenAiCompatibleProvider extends BaseOpenAiCompatibleProvider<"test-model"> {
|
||||
constructor(apiKey: string) {
|
||||
const testModels: Record<"test-model", ModelInfo> = {
|
||||
"test-model": {
|
||||
maxTokens: 4096,
|
||||
contextWindow: 128000,
|
||||
supportsImages: false,
|
||||
supportsPromptCache: false,
|
||||
inputPrice: 0.5,
|
||||
outputPrice: 1.5,
|
||||
},
|
||||
}
|
||||
|
||||
super({
|
||||
providerName: "TestProvider",
|
||||
baseURL: "https://test.example.com/v1",
|
||||
defaultProviderModelId: "test-model",
|
||||
providerModels: testModels,
|
||||
apiKey,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
describe("BaseOpenAiCompatibleProvider", () => {
|
||||
let handler: TestOpenAiCompatibleProvider
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
handler = new TestOpenAiCompatibleProvider("test-api-key")
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks()
|
||||
})
|
||||
|
||||
describe("TagMatcher reasoning tags", () => {
|
||||
it("should handle reasoning tags (<think>) from stream", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "<think>Let me think" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: " about this</think>" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "The answer is 42" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// TagMatcher yields chunks as they're processed
|
||||
expect(chunks).toEqual([
|
||||
{ type: "reasoning", text: "Let me think" },
|
||||
{ type: "reasoning", text: " about this" },
|
||||
{ type: "text", text: "The answer is 42" },
|
||||
])
|
||||
})
|
||||
|
||||
it("should handle complete <think> tag in a single chunk", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "Regular text before " } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "<think>Complete thought</think>" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: " regular text after" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// When a complete tag arrives in one chunk, TagMatcher may not parse it
|
||||
// This test documents the actual behavior
|
||||
expect(chunks.length).toBeGreaterThan(0)
|
||||
expect(chunks[0]).toEqual({ type: "text", text: "Regular text before " })
|
||||
})
|
||||
|
||||
it("should handle incomplete <think> tag at end of stream", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "<think>Incomplete thought" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// TagMatcher should handle incomplete tags and flush remaining content
|
||||
expect(chunks.length).toBeGreaterThan(0)
|
||||
expect(
|
||||
chunks.some(
|
||||
(c) => (c.type === "text" || c.type === "reasoning") && c.text.includes("Incomplete thought"),
|
||||
),
|
||||
).toBe(true)
|
||||
})
|
||||
|
||||
it("should handle text without any <think> tags", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "Just regular text" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: " without reasoning" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
expect(chunks).toEqual([
|
||||
{ type: "text", text: "Just regular text" },
|
||||
{ type: "text", text: " without reasoning" },
|
||||
])
|
||||
})
|
||||
|
||||
it("should handle <think> tags that start at beginning of stream", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "<think>reasoning" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: " content</think>" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: " normal text" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
expect(chunks).toEqual([
|
||||
{ type: "reasoning", text: "reasoning" },
|
||||
{ type: "reasoning", text: " content" },
|
||||
{ type: "text", text: " normal text" },
|
||||
])
|
||||
})
|
||||
})
|
||||
|
||||
describe("reasoning_content field", () => {
|
||||
it("should filter out whitespace-only reasoning_content", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { reasoning_content: "\n" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { reasoning_content: " " } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { reasoning_content: "\t\n " } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { content: "Regular content" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// Should only have the regular content, not the whitespace-only reasoning
|
||||
expect(chunks).toEqual([{ type: "text", text: "Regular content" }])
|
||||
})
|
||||
|
||||
it("should yield non-empty reasoning_content", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { reasoning_content: "Thinking step 1" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { reasoning_content: "\n" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { reasoning_content: "Thinking step 2" } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// Should only yield the non-empty reasoning content
|
||||
expect(chunks).toEqual([
|
||||
{ type: "reasoning", text: "Thinking step 1" },
|
||||
{ type: "reasoning", text: "Thinking step 2" },
|
||||
])
|
||||
})
|
||||
|
||||
it("should handle reasoning_content with leading/trailing whitespace", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: { choices: [{ delta: { reasoning_content: " content with spaces " } }] },
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// Should yield reasoning with spaces (only pure whitespace is filtered)
|
||||
expect(chunks).toEqual([{ type: "reasoning", text: " content with spaces " }])
|
||||
})
|
||||
})
|
||||
|
||||
describe("Basic functionality", () => {
|
||||
it("should create stream with correct parameters", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
async next() {
|
||||
return { done: true }
|
||||
},
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const systemPrompt = "Test system prompt"
|
||||
const messages: Anthropic.Messages.MessageParam[] = [{ role: "user", content: "Test message" }]
|
||||
|
||||
const messageGenerator = handler.createMessage(systemPrompt, messages)
|
||||
await messageGenerator.next()
|
||||
|
||||
expect(mockCreate).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
model: "test-model",
|
||||
temperature: 0,
|
||||
messages: expect.arrayContaining([{ role: "system", content: systemPrompt }]),
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
}),
|
||||
undefined,
|
||||
)
|
||||
})
|
||||
|
||||
it("should yield usage data from stream", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: {
|
||||
choices: [{ delta: {} }],
|
||||
usage: { prompt_tokens: 100, completion_tokens: 50 },
|
||||
},
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const firstChunk = await stream.next()
|
||||
|
||||
expect(firstChunk.done).toBe(false)
|
||||
expect(firstChunk.value).toMatchObject({ type: "usage", inputTokens: 100, outputTokens: 50 })
|
||||
})
|
||||
})
|
||||
|
||||
describe("Tool call handling", () => {
|
||||
it("should yield tool_call_end events when finish_reason is tool_calls", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: {
|
||||
choices: [
|
||||
{
|
||||
delta: {
|
||||
tool_calls: [
|
||||
{
|
||||
index: 0,
|
||||
id: "call_123",
|
||||
function: { name: "test_tool", arguments: '{"arg":' },
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: {
|
||||
choices: [
|
||||
{
|
||||
delta: {
|
||||
tool_calls: [
|
||||
{
|
||||
index: 0,
|
||||
function: { arguments: '"value"}' },
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: {
|
||||
choices: [
|
||||
{
|
||||
delta: {},
|
||||
finish_reason: "tool_calls",
|
||||
},
|
||||
],
|
||||
},
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// Should have tool_call_partial and tool_call_end
|
||||
const partialChunks = chunks.filter((chunk) => chunk.type === "tool_call_partial")
|
||||
const endChunks = chunks.filter((chunk) => chunk.type === "tool_call_end")
|
||||
|
||||
expect(partialChunks).toHaveLength(2)
|
||||
expect(endChunks).toHaveLength(1)
|
||||
expect(endChunks[0]).toEqual({ type: "tool_call_end", id: "call_123" })
|
||||
})
|
||||
|
||||
it("should yield multiple tool_call_end events for parallel tool calls", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: {
|
||||
choices: [
|
||||
{
|
||||
delta: {
|
||||
tool_calls: [
|
||||
{
|
||||
index: 0,
|
||||
id: "call_001",
|
||||
function: { name: "tool_a", arguments: "{}" },
|
||||
},
|
||||
{
|
||||
index: 1,
|
||||
id: "call_002",
|
||||
function: { name: "tool_b", arguments: "{}" },
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
})
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: {
|
||||
choices: [
|
||||
{
|
||||
delta: {},
|
||||
finish_reason: "tool_calls",
|
||||
},
|
||||
],
|
||||
},
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const endChunks = chunks.filter((chunk) => chunk.type === "tool_call_end")
|
||||
expect(endChunks).toHaveLength(2)
|
||||
expect(endChunks.map((c: any) => c.id).sort()).toEqual(["call_001", "call_002"])
|
||||
})
|
||||
|
||||
it("should not yield tool_call_end when finish_reason is not tool_calls", async () => {
|
||||
mockCreate.mockImplementationOnce(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: () => ({
|
||||
next: vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce({
|
||||
done: false,
|
||||
value: {
|
||||
choices: [
|
||||
{
|
||||
delta: { content: "Some text response" },
|
||||
finish_reason: "stop",
|
||||
},
|
||||
],
|
||||
},
|
||||
})
|
||||
.mockResolvedValueOnce({ done: true }),
|
||||
}),
|
||||
}
|
||||
})
|
||||
|
||||
const stream = handler.createMessage("system prompt", [])
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const endChunks = chunks.filter((chunk) => chunk.type === "tool_call_end")
|
||||
expect(endChunks).toHaveLength(0)
|
||||
})
|
||||
})
|
||||
})
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,260 +0,0 @@
|
|||
import { Anthropic } from "@anthropic-ai/sdk"
|
||||
import OpenAI from "openai"
|
||||
|
||||
import type { ModelInfo } from "@roo-code/types"
|
||||
|
||||
import { type ApiHandlerOptions, getModelMaxOutputTokens } from "../../shared/api"
|
||||
import { TagMatcher } from "../../utils/tag-matcher"
|
||||
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
|
||||
import { convertToOpenAiMessages } from "../transform/openai-format"
|
||||
|
||||
import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index"
|
||||
import { DEFAULT_HEADERS } from "./constants"
|
||||
import { BaseProvider } from "./base-provider"
|
||||
import { handleOpenAIError } from "./utils/openai-error-handler"
|
||||
import { calculateApiCostOpenAI } from "../../shared/cost"
|
||||
import { getApiRequestTimeout } from "./utils/timeout-config"
|
||||
|
||||
type BaseOpenAiCompatibleProviderOptions<ModelName extends string> = ApiHandlerOptions & {
|
||||
providerName: string
|
||||
baseURL: string
|
||||
defaultProviderModelId: ModelName
|
||||
providerModels: Record<ModelName, ModelInfo>
|
||||
defaultTemperature?: number
|
||||
}
|
||||
|
||||
export abstract class BaseOpenAiCompatibleProvider<ModelName extends string>
|
||||
extends BaseProvider
|
||||
implements SingleCompletionHandler
|
||||
{
|
||||
protected readonly providerName: string
|
||||
protected readonly baseURL: string
|
||||
protected readonly defaultTemperature: number
|
||||
protected readonly defaultProviderModelId: ModelName
|
||||
protected readonly providerModels: Record<ModelName, ModelInfo>
|
||||
|
||||
protected readonly options: ApiHandlerOptions
|
||||
|
||||
protected client: OpenAI
|
||||
|
||||
constructor({
|
||||
providerName,
|
||||
baseURL,
|
||||
defaultProviderModelId,
|
||||
providerModels,
|
||||
defaultTemperature,
|
||||
...options
|
||||
}: BaseOpenAiCompatibleProviderOptions<ModelName>) {
|
||||
super()
|
||||
|
||||
this.providerName = providerName
|
||||
this.baseURL = baseURL
|
||||
this.defaultProviderModelId = defaultProviderModelId
|
||||
this.providerModels = providerModels
|
||||
this.defaultTemperature = defaultTemperature ?? 0
|
||||
|
||||
this.options = options
|
||||
|
||||
if (!this.options.apiKey) {
|
||||
throw new Error("API key is required")
|
||||
}
|
||||
|
||||
this.client = new OpenAI({
|
||||
baseURL,
|
||||
apiKey: this.options.apiKey,
|
||||
defaultHeaders: DEFAULT_HEADERS,
|
||||
timeout: getApiRequestTimeout(),
|
||||
})
|
||||
}
|
||||
|
||||
protected createStream(
|
||||
systemPrompt: string,
|
||||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
requestOptions?: OpenAI.RequestOptions,
|
||||
) {
|
||||
const { id: model, info } = this.getModel()
|
||||
|
||||
// Centralized cap: clamp to 20% of the context window (unless provider-specific exceptions apply)
|
||||
const max_tokens =
|
||||
getModelMaxOutputTokens({
|
||||
modelId: model,
|
||||
model: info,
|
||||
settings: this.options,
|
||||
format: "openai",
|
||||
}) ?? undefined
|
||||
|
||||
const temperature = this.options.modelTemperature ?? info.defaultTemperature ?? this.defaultTemperature
|
||||
|
||||
const params: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = {
|
||||
model,
|
||||
max_tokens,
|
||||
temperature,
|
||||
messages: [{ role: "system", content: systemPrompt }, ...convertToOpenAiMessages(messages)],
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
tools: this.convertToolsForOpenAI(metadata?.tools),
|
||||
tool_choice: metadata?.tool_choice,
|
||||
parallel_tool_calls: metadata?.parallelToolCalls ?? true,
|
||||
}
|
||||
|
||||
// Add thinking parameter if reasoning is enabled and model supports it
|
||||
if (this.options.enableReasoningEffort && info.supportsReasoningBinary) {
|
||||
;(params as any).thinking = { type: "enabled" }
|
||||
}
|
||||
|
||||
try {
|
||||
return this.client.chat.completions.create(params, requestOptions)
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
}
|
||||
|
||||
override async *createMessage(
|
||||
systemPrompt: string,
|
||||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
): ApiStream {
|
||||
const stream = await this.createStream(systemPrompt, messages, metadata)
|
||||
|
||||
const matcher = new TagMatcher(
|
||||
"think",
|
||||
(chunk) =>
|
||||
({
|
||||
type: chunk.matched ? "reasoning" : "text",
|
||||
text: chunk.data,
|
||||
}) as const,
|
||||
)
|
||||
|
||||
let lastUsage: OpenAI.CompletionUsage | undefined
|
||||
const activeToolCallIds = new Set<string>()
|
||||
|
||||
for await (const chunk of stream) {
|
||||
// Check for provider-specific error responses (e.g., MiniMax base_resp)
|
||||
const chunkAny = chunk as any
|
||||
if (chunkAny.base_resp?.status_code && chunkAny.base_resp.status_code !== 0) {
|
||||
throw new Error(
|
||||
`${this.providerName} API Error (${chunkAny.base_resp.status_code}): ${chunkAny.base_resp.status_msg || "Unknown error"}`,
|
||||
)
|
||||
}
|
||||
|
||||
const delta = chunk.choices?.[0]?.delta
|
||||
const finishReason = chunk.choices?.[0]?.finish_reason
|
||||
|
||||
if (delta?.content) {
|
||||
for (const processedChunk of matcher.update(delta.content)) {
|
||||
yield processedChunk
|
||||
}
|
||||
}
|
||||
|
||||
if (delta) {
|
||||
for (const key of ["reasoning_content", "reasoning"] as const) {
|
||||
if (key in delta) {
|
||||
const reasoning_content = ((delta as any)[key] as string | undefined) || ""
|
||||
if (reasoning_content?.trim()) {
|
||||
yield { type: "reasoning", text: reasoning_content }
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Emit raw tool call chunks - NativeToolCallParser handles state management
|
||||
if (delta?.tool_calls) {
|
||||
for (const toolCall of delta.tool_calls) {
|
||||
if (toolCall.id) {
|
||||
activeToolCallIds.add(toolCall.id)
|
||||
}
|
||||
yield {
|
||||
type: "tool_call_partial",
|
||||
index: toolCall.index,
|
||||
id: toolCall.id,
|
||||
name: toolCall.function?.name,
|
||||
arguments: toolCall.function?.arguments,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Emit tool_call_end events when finish_reason is "tool_calls"
|
||||
// This ensures tool calls are finalized even if the stream doesn't properly close
|
||||
if (finishReason === "tool_calls" && activeToolCallIds.size > 0) {
|
||||
for (const id of activeToolCallIds) {
|
||||
yield { type: "tool_call_end", id }
|
||||
}
|
||||
activeToolCallIds.clear()
|
||||
}
|
||||
|
||||
if (chunk.usage) {
|
||||
lastUsage = chunk.usage
|
||||
}
|
||||
}
|
||||
|
||||
if (lastUsage) {
|
||||
yield this.processUsageMetrics(lastUsage, this.getModel().info)
|
||||
}
|
||||
|
||||
// Process any remaining content
|
||||
for (const processedChunk of matcher.final()) {
|
||||
yield processedChunk
|
||||
}
|
||||
}
|
||||
|
||||
protected processUsageMetrics(usage: any, modelInfo?: any): ApiStreamUsageChunk {
|
||||
const inputTokens = usage?.prompt_tokens || 0
|
||||
const outputTokens = usage?.completion_tokens || 0
|
||||
const cacheWriteTokens = usage?.prompt_tokens_details?.cache_write_tokens || 0
|
||||
const cacheReadTokens = usage?.prompt_tokens_details?.cached_tokens || 0
|
||||
|
||||
const { totalCost } = modelInfo
|
||||
? calculateApiCostOpenAI(modelInfo, inputTokens, outputTokens, cacheWriteTokens, cacheReadTokens)
|
||||
: { totalCost: 0 }
|
||||
|
||||
return {
|
||||
type: "usage",
|
||||
inputTokens,
|
||||
outputTokens,
|
||||
cacheWriteTokens: cacheWriteTokens || undefined,
|
||||
cacheReadTokens: cacheReadTokens || undefined,
|
||||
totalCost,
|
||||
}
|
||||
}
|
||||
|
||||
async completePrompt(prompt: string): Promise<string> {
|
||||
const { id: modelId, info: modelInfo } = this.getModel()
|
||||
|
||||
const params: OpenAI.Chat.Completions.ChatCompletionCreateParams = {
|
||||
model: modelId,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
}
|
||||
|
||||
// Add thinking parameter if reasoning is enabled and model supports it
|
||||
if (this.options.enableReasoningEffort && modelInfo.supportsReasoningBinary) {
|
||||
;(params as any).thinking = { type: "enabled" }
|
||||
}
|
||||
|
||||
try {
|
||||
const response = await this.client.chat.completions.create(params)
|
||||
|
||||
// Check for provider-specific error responses (e.g., MiniMax base_resp)
|
||||
const responseAny = response as any
|
||||
if (responseAny.base_resp?.status_code && responseAny.base_resp.status_code !== 0) {
|
||||
throw new Error(
|
||||
`${this.providerName} API Error (${responseAny.base_resp.status_code}): ${responseAny.base_resp.status_msg || "Unknown error"}`,
|
||||
)
|
||||
}
|
||||
|
||||
return response.choices?.[0]?.message.content || ""
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
}
|
||||
|
||||
override getModel() {
|
||||
const id =
|
||||
this.options.apiModelId && this.options.apiModelId in this.providerModels
|
||||
? (this.options.apiModelId as ModelName)
|
||||
: this.defaultProviderModelId
|
||||
|
||||
return { id, info: this.providerModels[id] }
|
||||
}
|
||||
}
|
||||
|
|
@ -1,90 +1,116 @@
|
|||
import { Anthropic } from "@anthropic-ai/sdk"
|
||||
import OpenAI from "openai"
|
||||
import { createOpenAICompatible } from "@ai-sdk/openai-compatible"
|
||||
import { streamText, generateText } from "ai"
|
||||
|
||||
import { rooDefaultModelId, getApiProtocol, type ImageGenerationApiMethod } from "@roo-code/types"
|
||||
import { CloudService } from "@roo-code/cloud"
|
||||
|
||||
import { NativeToolCallParser } from "../../core/assistant-message/NativeToolCallParser"
|
||||
|
||||
import { Package } from "../../shared/package"
|
||||
import type { ApiHandlerOptions } from "../../shared/api"
|
||||
import { calculateApiCostOpenAI } from "../../shared/cost"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
import { getModelParams } from "../transform/model-params"
|
||||
import { convertToOpenAiMessages } from "../transform/openai-format"
|
||||
import {
|
||||
convertToAiSdkMessages,
|
||||
convertToolsForAiSdk,
|
||||
processAiSdkStreamPart,
|
||||
handleAiSdkError,
|
||||
mapToolChoice,
|
||||
} from "../transform/ai-sdk"
|
||||
import { type ReasoningDetail } from "../transform/openai-format"
|
||||
import type { RooReasoningParams } from "../transform/reasoning"
|
||||
import { getRooReasoning } from "../transform/reasoning"
|
||||
|
||||
import type { ApiHandlerCreateMessageMetadata } from "../index"
|
||||
import { BaseOpenAiCompatibleProvider } from "./base-openai-compatible-provider"
|
||||
import { getModels, getModelsFromCache } from "../providers/fetchers/modelCache"
|
||||
import { handleOpenAIError } from "./utils/openai-error-handler"
|
||||
import type { ApiHandlerCreateMessageMetadata, SingleCompletionHandler } from "../index"
|
||||
import { BaseProvider } from "./base-provider"
|
||||
import { getModels, getModelsFromCache } from "./fetchers/modelCache"
|
||||
import { generateImageWithProvider, generateImageWithImagesApi, ImageGenerationResult } from "./utils/image-generation"
|
||||
import { t } from "../../i18n"
|
||||
|
||||
// Extend OpenAI's CompletionUsage to include Roo specific fields
|
||||
interface RooUsage extends OpenAI.CompletionUsage {
|
||||
cache_creation_input_tokens?: number
|
||||
cost?: number
|
||||
}
|
||||
|
||||
// Add custom interface for Roo params to support reasoning
|
||||
type RooChatCompletionParams = OpenAI.Chat.ChatCompletionCreateParamsStreaming & {
|
||||
reasoning?: RooReasoningParams
|
||||
}
|
||||
|
||||
function getSessionToken(): string {
|
||||
const token = CloudService.hasInstance() ? CloudService.instance.authService?.getSessionToken() : undefined
|
||||
return token ?? "unauthenticated"
|
||||
}
|
||||
|
||||
export class RooHandler extends BaseOpenAiCompatibleProvider<string> {
|
||||
export class RooHandler extends BaseProvider implements SingleCompletionHandler {
|
||||
protected options: ApiHandlerOptions
|
||||
private fetcherBaseURL: string
|
||||
private currentReasoningDetails: any[] = []
|
||||
private currentReasoningDetails: ReasoningDetail[] = []
|
||||
|
||||
constructor(options: ApiHandlerOptions) {
|
||||
const sessionToken = options.rooApiKey ?? getSessionToken()
|
||||
super()
|
||||
this.options = options
|
||||
|
||||
let baseURL = process.env.ROO_CODE_PROVIDER_URL ?? "https://api.roocode.com/proxy"
|
||||
|
||||
// Ensure baseURL ends with /v1 for OpenAI client, but don't duplicate it
|
||||
// Ensure baseURL ends with /v1 for API calls, but don't duplicate it
|
||||
if (!baseURL.endsWith("/v1")) {
|
||||
baseURL = `${baseURL}/v1`
|
||||
}
|
||||
|
||||
// Always construct the handler, even without a valid token.
|
||||
// The provider-proxy server will return 401 if authentication fails.
|
||||
super({
|
||||
...options,
|
||||
providerName: "Roo Code Cloud",
|
||||
baseURL, // Already has /v1 suffix
|
||||
apiKey: sessionToken,
|
||||
defaultProviderModelId: rooDefaultModelId,
|
||||
providerModels: {},
|
||||
})
|
||||
|
||||
// Load dynamic models asynchronously - strip /v1 from baseURL for fetcher
|
||||
// Strip /v1 from baseURL for fetcher
|
||||
this.fetcherBaseURL = baseURL.endsWith("/v1") ? baseURL.slice(0, -3) : baseURL
|
||||
|
||||
const sessionToken = options.rooApiKey ?? getSessionToken()
|
||||
|
||||
this.loadDynamicModels(this.fetcherBaseURL, sessionToken).catch((error) => {
|
||||
console.error("[RooHandler] Failed to load dynamic models:", error)
|
||||
})
|
||||
}
|
||||
|
||||
protected override createStream(
|
||||
/**
|
||||
* Per-request provider factory. Creates a fresh provider instance
|
||||
* to ensure the latest session token is used for each request.
|
||||
*/
|
||||
private createRooProvider(options?: { reasoning?: RooReasoningParams; taskId?: string }) {
|
||||
const token = this.options.rooApiKey ?? getSessionToken()
|
||||
const headers: Record<string, string> = {
|
||||
"X-Roo-App-Version": Package.version,
|
||||
}
|
||||
if (options?.taskId) {
|
||||
headers["X-Roo-Task-ID"] = options.taskId
|
||||
}
|
||||
const reasoning = options?.reasoning
|
||||
return createOpenAICompatible({
|
||||
name: "roo",
|
||||
apiKey: token || "not-provided",
|
||||
baseURL: `${this.fetcherBaseURL}/v1`,
|
||||
headers,
|
||||
...(reasoning && {
|
||||
transformRequestBody: (body: Record<string, unknown>) => ({
|
||||
...body,
|
||||
reasoning,
|
||||
}),
|
||||
}),
|
||||
})
|
||||
}
|
||||
|
||||
override isAiSdkProvider() {
|
||||
return true as const
|
||||
}
|
||||
|
||||
getReasoningDetails(): ReasoningDetail[] | undefined {
|
||||
return this.currentReasoningDetails.length > 0 ? this.currentReasoningDetails : undefined
|
||||
}
|
||||
|
||||
override async *createMessage(
|
||||
systemPrompt: string,
|
||||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
requestOptions?: OpenAI.RequestOptions,
|
||||
) {
|
||||
const { id: model, info } = this.getModel()
|
||||
): ApiStream {
|
||||
// Reset reasoning_details accumulator for this request
|
||||
this.currentReasoningDetails = []
|
||||
|
||||
// Get model parameters including reasoning
|
||||
const model = this.getModel()
|
||||
const { id: modelId, info } = model
|
||||
|
||||
// Get model parameters including reasoning budget/effort
|
||||
const params = getModelParams({
|
||||
format: "openai",
|
||||
modelId: model,
|
||||
modelId,
|
||||
model: info,
|
||||
settings: this.options,
|
||||
defaultTemperature: this.defaultTemperature,
|
||||
defaultTemperature: 0,
|
||||
})
|
||||
|
||||
// Get Roo-specific reasoning parameters
|
||||
|
|
@ -95,231 +121,102 @@ export class RooHandler extends BaseOpenAiCompatibleProvider<string> {
|
|||
settings: this.options,
|
||||
})
|
||||
|
||||
const max_tokens = params.maxTokens ?? undefined
|
||||
const temperature = params.temperature ?? this.defaultTemperature
|
||||
const maxTokens = params.maxTokens ?? undefined
|
||||
const temperature = params.temperature ?? 0
|
||||
|
||||
const rooParams: RooChatCompletionParams = {
|
||||
model,
|
||||
max_tokens,
|
||||
temperature,
|
||||
messages: [{ role: "system", content: systemPrompt }, ...convertToOpenAiMessages(messages)],
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
...(reasoning && { reasoning }),
|
||||
tools: this.convertToolsForOpenAI(metadata?.tools),
|
||||
tool_choice: metadata?.tool_choice,
|
||||
}
|
||||
// Create per-request provider with fresh session token
|
||||
const provider = this.createRooProvider({ reasoning, taskId: metadata?.taskId })
|
||||
|
||||
// Convert messages and tools to AI SDK format
|
||||
const aiSdkMessages = convertToAiSdkMessages(messages)
|
||||
const tools = convertToolsForAiSdk(this.convertToolsForOpenAI(metadata?.tools))
|
||||
|
||||
let accumulatedReasoningText = ""
|
||||
let lastStreamError: string | undefined
|
||||
|
||||
try {
|
||||
this.client.apiKey = this.options.rooApiKey ?? getSessionToken()
|
||||
return this.client.chat.completions.create(rooParams, requestOptions)
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
}
|
||||
const result = streamText({
|
||||
model: provider(modelId),
|
||||
system: systemPrompt,
|
||||
messages: aiSdkMessages,
|
||||
maxOutputTokens: maxTokens && maxTokens > 0 ? maxTokens : undefined,
|
||||
temperature,
|
||||
tools,
|
||||
toolChoice: mapToolChoice(metadata?.tool_choice),
|
||||
})
|
||||
|
||||
getReasoningDetails(): any[] | undefined {
|
||||
return this.currentReasoningDetails.length > 0 ? this.currentReasoningDetails : undefined
|
||||
}
|
||||
|
||||
override async *createMessage(
|
||||
systemPrompt: string,
|
||||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
): ApiStream {
|
||||
try {
|
||||
// Reset reasoning_details accumulator for this request
|
||||
this.currentReasoningDetails = []
|
||||
|
||||
const headers: Record<string, string> = {
|
||||
"X-Roo-App-Version": Package.version,
|
||||
}
|
||||
|
||||
if (metadata?.taskId) {
|
||||
headers["X-Roo-Task-ID"] = metadata.taskId
|
||||
}
|
||||
|
||||
const stream = await this.createStream(systemPrompt, messages, metadata, { headers })
|
||||
|
||||
let lastUsage: RooUsage | undefined = undefined
|
||||
// Accumulator for reasoning_details FROM the API.
|
||||
// We preserve the original shape of reasoning_details to prevent malformed responses.
|
||||
const reasoningDetailsAccumulator = new Map<
|
||||
string,
|
||||
{
|
||||
type: string
|
||||
text?: string
|
||||
summary?: string
|
||||
data?: string
|
||||
id?: string | null
|
||||
format?: string
|
||||
signature?: string
|
||||
index: number
|
||||
for await (const part of result.fullStream) {
|
||||
if (part.type === "reasoning-delta" && part.text !== "[REDACTED]") {
|
||||
accumulatedReasoningText += part.text
|
||||
}
|
||||
>()
|
||||
|
||||
// Track whether we've yielded displayable text from reasoning_details.
|
||||
// When reasoning_details has displayable content (reasoning.text or reasoning.summary),
|
||||
// we skip yielding the top-level reasoning field to avoid duplicate display.
|
||||
let hasYieldedReasoningFromDetails = false
|
||||
|
||||
for await (const chunk of stream) {
|
||||
const delta = chunk.choices[0]?.delta
|
||||
const finishReason = chunk.choices[0]?.finish_reason
|
||||
|
||||
if (delta) {
|
||||
// Handle reasoning_details array format (used by Gemini 3, Claude, OpenAI o-series, etc.)
|
||||
// See: https://openrouter.ai/docs/use-cases/reasoning-tokens#preserving-reasoning-blocks
|
||||
// Priority: Check for reasoning_details first, as it's the newer format
|
||||
const deltaWithReasoning = delta as typeof delta & {
|
||||
reasoning_details?: Array<{
|
||||
type: string
|
||||
text?: string
|
||||
summary?: string
|
||||
data?: string
|
||||
id?: string | null
|
||||
format?: string
|
||||
signature?: string
|
||||
index?: number
|
||||
}>
|
||||
for (const chunk of processAiSdkStreamPart(part)) {
|
||||
if (chunk.type === "error") {
|
||||
lastStreamError = chunk.message
|
||||
}
|
||||
|
||||
if (deltaWithReasoning.reasoning_details && Array.isArray(deltaWithReasoning.reasoning_details)) {
|
||||
for (const detail of deltaWithReasoning.reasoning_details) {
|
||||
const index = detail.index ?? 0
|
||||
// Use id as key when available to merge chunks that share the same reasoning block id
|
||||
// This ensures that reasoning.summary and reasoning.encrypted chunks with the same id
|
||||
// are merged into a single object, matching the provider's expected format
|
||||
const key = detail.id ?? `${detail.type}-${index}`
|
||||
const existing = reasoningDetailsAccumulator.get(key)
|
||||
|
||||
if (existing) {
|
||||
// Accumulate text/summary/data for existing reasoning detail
|
||||
if (detail.text !== undefined) {
|
||||
existing.text = (existing.text || "") + detail.text
|
||||
}
|
||||
if (detail.summary !== undefined) {
|
||||
existing.summary = (existing.summary || "") + detail.summary
|
||||
}
|
||||
if (detail.data !== undefined) {
|
||||
existing.data = (existing.data || "") + detail.data
|
||||
}
|
||||
// Update other fields if provided
|
||||
// Note: Don't update type - keep original type (e.g., reasoning.summary)
|
||||
// even when encrypted data chunks arrive with type reasoning.encrypted
|
||||
if (detail.id !== undefined) existing.id = detail.id
|
||||
if (detail.format !== undefined) existing.format = detail.format
|
||||
if (detail.signature !== undefined) existing.signature = detail.signature
|
||||
} else {
|
||||
// Start new reasoning detail accumulation
|
||||
reasoningDetailsAccumulator.set(key, {
|
||||
type: detail.type,
|
||||
text: detail.text,
|
||||
summary: detail.summary,
|
||||
data: detail.data,
|
||||
id: detail.id,
|
||||
format: detail.format,
|
||||
signature: detail.signature,
|
||||
index,
|
||||
})
|
||||
}
|
||||
|
||||
// Yield text for display (still fragmented for live streaming)
|
||||
// Only reasoning.text and reasoning.summary have displayable content
|
||||
// reasoning.encrypted is intentionally skipped as it contains redacted content
|
||||
let reasoningText: string | undefined
|
||||
if (detail.type === "reasoning.text" && typeof detail.text === "string") {
|
||||
reasoningText = detail.text
|
||||
} else if (detail.type === "reasoning.summary" && typeof detail.summary === "string") {
|
||||
reasoningText = detail.summary
|
||||
}
|
||||
|
||||
if (reasoningText) {
|
||||
hasYieldedReasoningFromDetails = true
|
||||
yield { type: "reasoning", text: reasoningText }
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Handle top-level reasoning field for UI display.
|
||||
// Skip if we've already yielded from reasoning_details to avoid duplicate display.
|
||||
if ("reasoning" in delta && delta.reasoning && typeof delta.reasoning === "string") {
|
||||
if (!hasYieldedReasoningFromDetails) {
|
||||
yield { type: "reasoning", text: delta.reasoning }
|
||||
}
|
||||
} else if ("reasoning_content" in delta && typeof delta.reasoning_content === "string") {
|
||||
// Also check for reasoning_content for backward compatibility
|
||||
if (!hasYieldedReasoningFromDetails) {
|
||||
yield { type: "reasoning", text: delta.reasoning_content }
|
||||
}
|
||||
}
|
||||
|
||||
// Emit raw tool call chunks - NativeToolCallParser handles state management
|
||||
if ("tool_calls" in delta && Array.isArray(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 (delta.content) {
|
||||
yield {
|
||||
type: "text",
|
||||
text: delta.content,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (finishReason) {
|
||||
const endEvents = NativeToolCallParser.processFinishReason(finishReason)
|
||||
for (const event of endEvents) {
|
||||
yield event
|
||||
}
|
||||
}
|
||||
|
||||
if (chunk.usage) {
|
||||
lastUsage = chunk.usage as RooUsage
|
||||
yield chunk
|
||||
}
|
||||
}
|
||||
|
||||
// After streaming completes, store ONLY the reasoning_details we received from the API.
|
||||
if (reasoningDetailsAccumulator.size > 0) {
|
||||
this.currentReasoningDetails = Array.from(reasoningDetailsAccumulator.values())
|
||||
// Build reasoning details from accumulated text
|
||||
if (accumulatedReasoningText) {
|
||||
this.currentReasoningDetails.push({
|
||||
type: "reasoning.text",
|
||||
text: accumulatedReasoningText,
|
||||
index: 0,
|
||||
})
|
||||
}
|
||||
|
||||
if (lastUsage) {
|
||||
// Check if the current model is marked as free
|
||||
const model = this.getModel()
|
||||
const isFreeModel = model.info.isFree ?? false
|
||||
// Check provider metadata for reasoning_details (override if present)
|
||||
const providerMetadata =
|
||||
(await result.providerMetadata) ?? (await (result as any).experimental_providerMetadata)
|
||||
const rooMeta = providerMetadata?.roo as Record<string, any> | undefined
|
||||
|
||||
// Normalize input tokens based on protocol expectations:
|
||||
// - OpenAI protocol expects TOTAL input tokens (cached + non-cached)
|
||||
// - Anthropic protocol expects NON-CACHED input tokens (caches passed separately)
|
||||
const modelId = model.id
|
||||
const apiProtocol = getApiProtocol("roo", modelId)
|
||||
const providerReasoningDetails = rooMeta?.reasoning_details as ReasoningDetail[] | undefined
|
||||
if (providerReasoningDetails && providerReasoningDetails.length > 0) {
|
||||
this.currentReasoningDetails = providerReasoningDetails
|
||||
}
|
||||
|
||||
const promptTokens = lastUsage.prompt_tokens || 0
|
||||
const cacheWrite = lastUsage.cache_creation_input_tokens || 0
|
||||
const cacheRead = lastUsage.prompt_tokens_details?.cached_tokens || 0
|
||||
const nonCached = Math.max(0, promptTokens - cacheWrite - cacheRead)
|
||||
// Process usage with protocol-aware normalization
|
||||
const usage = await result.usage
|
||||
const promptTokens = usage.inputTokens ?? 0
|
||||
const completionTokens = usage.outputTokens ?? 0
|
||||
|
||||
const inputTokensForDownstream = apiProtocol === "anthropic" ? nonCached : promptTokens
|
||||
// Extract cache tokens from provider metadata
|
||||
const cacheCreation = (rooMeta?.cache_creation_input_tokens as number) ?? 0
|
||||
const cacheRead = (rooMeta?.cache_read_input_tokens as number) ?? (rooMeta?.cached_tokens as number) ?? 0
|
||||
|
||||
yield {
|
||||
type: "usage",
|
||||
inputTokens: inputTokensForDownstream,
|
||||
outputTokens: lastUsage.completion_tokens || 0,
|
||||
cacheWriteTokens: cacheWrite,
|
||||
cacheReadTokens: cacheRead,
|
||||
totalCost: isFreeModel ? 0 : (lastUsage.cost ?? 0),
|
||||
}
|
||||
// Protocol-aware token normalization:
|
||||
// - OpenAI protocol expects TOTAL input tokens (cached + non-cached)
|
||||
// - Anthropic protocol expects NON-CACHED input tokens (caches passed separately)
|
||||
const apiProtocol = getApiProtocol("roo", modelId)
|
||||
const nonCached = Math.max(0, promptTokens - cacheCreation - cacheRead)
|
||||
const inputTokens = apiProtocol === "anthropic" ? nonCached : promptTokens
|
||||
|
||||
// Cost: prefer server-side cost, fall back to client-side calculation
|
||||
const isFreeModel = info.isFree === true
|
||||
const serverCost = rooMeta?.cost as number | undefined
|
||||
const { totalCost: calculatedCost } = calculateApiCostOpenAI(
|
||||
info,
|
||||
promptTokens,
|
||||
completionTokens,
|
||||
cacheCreation,
|
||||
cacheRead,
|
||||
)
|
||||
const totalCost = isFreeModel ? 0 : (serverCost ?? calculatedCost)
|
||||
|
||||
yield {
|
||||
type: "usage" as const,
|
||||
inputTokens,
|
||||
outputTokens: completionTokens,
|
||||
cacheWriteTokens: cacheCreation,
|
||||
cacheReadTokens: cacheRead,
|
||||
totalCost,
|
||||
}
|
||||
} catch (error) {
|
||||
if (lastStreamError) {
|
||||
throw new Error(lastStreamError)
|
||||
}
|
||||
|
||||
const errorContext = {
|
||||
error: error instanceof Error ? error.message : String(error),
|
||||
stack: error instanceof Error ? error.stack : undefined,
|
||||
|
|
@ -329,13 +226,24 @@ export class RooHandler extends BaseOpenAiCompatibleProvider<string> {
|
|||
|
||||
console.error(`[RooHandler] Error during message streaming: ${JSON.stringify(errorContext)}`)
|
||||
|
||||
throw error
|
||||
throw handleAiSdkError(error, "Roo Code Cloud")
|
||||
}
|
||||
}
|
||||
override async completePrompt(prompt: string): Promise<string> {
|
||||
// Update API key before making request to ensure we use the latest session token
|
||||
this.client.apiKey = this.options.rooApiKey ?? getSessionToken()
|
||||
return super.completePrompt(prompt)
|
||||
|
||||
async completePrompt(prompt: string): Promise<string> {
|
||||
const { id: modelId } = this.getModel()
|
||||
const provider = this.createRooProvider()
|
||||
|
||||
try {
|
||||
const result = await generateText({
|
||||
model: provider(modelId),
|
||||
prompt,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
})
|
||||
return result.text
|
||||
} catch (error) {
|
||||
throw handleAiSdkError(error, "Roo Code Cloud")
|
||||
}
|
||||
}
|
||||
|
||||
private async loadDynamicModels(baseURL: string, apiKey?: string): Promise<void> {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue