mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-09 03:17:58 +00:00
feat: migrate OpenAI Native provider to @ai-sdk/openai (#11330)
* feat: migrate OpenAI Native provider to @ai-sdk/openai
Replace the raw OpenAI SDK (openai) usage in OpenAiNativeHandler with
@ai-sdk/openai and AI SDK's streamText/generateText, following the same
pattern used by other migrated providers (Groq, xAI, Fireworks, etc.).
Key changes:
- Use createOpenAI from @ai-sdk/openai with provider.responses() for
the Responses API
- Use streamText/generateText from ai for streaming and completions
- Pass OpenAI-specific features via providerOptions.openai (store,
reasoningEffort, reasoningSummary, textVerbosity, serviceTier,
promptCacheRetention, parallelToolCalls, include)
- Capture responseId, serviceTier, and encrypted reasoning content
from providerMetadata after streaming
- Preserve getEncryptedContent() and getResponseId() for Task.ts
- Preserve service tier pricing adjustment in cost calculation
- Mark as isAiSdkProvider: true
- Eliminate ~1100 lines of manual SSE parsing, raw fetch fallback,
and event handling code
- Rewrite all 3 test files to use AI SDK mocking pattern
* fix: remove non-existent cacheWriteTokens from providerMetadata
The OpenAI Responses API does not report cache write tokens separately.
Remove the reference to providerMetadata?.openai?.cacheWriteTokens which
does not exist in the @ai-sdk/openai provider metadata schema.
* fix: filter standalone encrypted reasoning items from messages
Task.ts buildCleanConversationHistory injects standalone reasoning items
with { type: 'reasoning', encrypted_content: '...' } into the messages
array. These have no 'role' property and would be silently dropped by
convertToAiSdkMessages. Filter them explicitly to prevent confusion.
Note: Encrypted reasoning content round-tripping for stateless continuity
is a known limitation of the AI SDK migration. The @ai-sdk/openai
provider does not support injecting raw Responses API reasoning items.
Plain-text reasoning round-tripping works correctly via isAiSdkProvider().
* fix: restore reasoning round-trip for OpenAI Responses API via AI SDK
- Strip plain-text reasoning blocks from assistant messages before
convertToAiSdkMessages() to eliminate 'Non-OpenAI reasoning parts'
warnings from @ai-sdk/openai Responses provider
- Re-inject encrypted reasoning items as AI SDK reasoning parts with
providerOptions.openai.itemId and reasoningEncryptedContent, restoring
reasoning continuity that was silently broken after the migration
- Restructure createMessage() into a 5-step pipeline:
collect → filter → strip → convert → inject
- Add 21 new tests for both plain-text stripping and encrypted
reasoning injection
---------
Co-authored-by: Hannes Rudolph <hrudolph@gmail.com>
This commit is contained in:
parent
7c58f29975
commit
b8ef352808
7 changed files with 2216 additions and 3318 deletions
17
pnpm-lock.yaml
generated
17
pnpm-lock.yaml
generated
|
|
@ -770,6 +770,9 @@ importers:
|
|||
'@ai-sdk/mistral':
|
||||
specifier: ^3.0.19
|
||||
version: 3.0.19(zod@3.25.76)
|
||||
'@ai-sdk/openai':
|
||||
specifier: ^3.0.26
|
||||
version: 3.0.26(zod@3.25.76)
|
||||
'@ai-sdk/xai':
|
||||
specifier: ^3.0.48
|
||||
version: 3.0.48(zod@3.25.76)
|
||||
|
|
@ -1486,6 +1489,12 @@ packages:
|
|||
peerDependencies:
|
||||
zod: 3.25.76
|
||||
|
||||
'@ai-sdk/openai@3.0.26':
|
||||
resolution: {integrity: sha512-W/hiwxIfG29IO0Fob1HwWpFssMsNrxWoX8A7DwNGOtKArDBmJNuGzQeU/k0Fnh8WyvZEnfxkjO4oXkSXfVBayg==}
|
||||
engines: {node: '>=18'}
|
||||
peerDependencies:
|
||||
zod: 3.25.76
|
||||
|
||||
'@ai-sdk/provider-utils@3.0.5':
|
||||
resolution: {integrity: sha512-HliwB/yzufw3iwczbFVE2Fiwf1XqROB/I6ng8EKUsPM5+2wnIa8f4VbljZcDx+grhFrPV+PnRZH7zBqi8WZM7Q==}
|
||||
engines: {node: '>=18'}
|
||||
|
|
@ -11107,6 +11116,12 @@ snapshots:
|
|||
'@ai-sdk/provider-utils': 4.0.14(zod@3.25.76)
|
||||
zod: 3.25.76
|
||||
|
||||
'@ai-sdk/openai@3.0.26(zod@3.25.76)':
|
||||
dependencies:
|
||||
'@ai-sdk/provider': 3.0.8
|
||||
'@ai-sdk/provider-utils': 4.0.14(zod@3.25.76)
|
||||
zod: 3.25.76
|
||||
|
||||
'@ai-sdk/provider-utils@3.0.5(zod@3.25.76)':
|
||||
dependencies:
|
||||
'@ai-sdk/provider': 2.0.0
|
||||
|
|
@ -14930,7 +14945,7 @@ snapshots:
|
|||
sirv: 3.0.1
|
||||
tinyglobby: 0.2.14
|
||||
tinyrainbow: 2.0.0
|
||||
vitest: 3.2.4(@types/debug@4.1.12)(@types/node@24.2.1)(@vitest/ui@3.2.4)(jiti@2.4.2)(jsdom@26.1.0)(lightningcss@1.30.1)(tsx@4.19.4)(yaml@2.8.0)
|
||||
vitest: 3.2.4(@types/debug@4.1.12)(@types/node@20.17.50)(@vitest/ui@3.2.4)(jiti@2.4.2)(jsdom@26.1.0)(lightningcss@1.30.1)(tsx@4.19.4)(yaml@2.8.0)
|
||||
|
||||
'@vitest/utils@3.2.4':
|
||||
dependencies:
|
||||
|
|
|
|||
565
src/api/providers/__tests__/openai-native-reasoning.spec.ts
Normal file
565
src/api/providers/__tests__/openai-native-reasoning.spec.ts
Normal file
|
|
@ -0,0 +1,565 @@
|
|||
// npx vitest run api/providers/__tests__/openai-native-reasoning.spec.ts
|
||||
|
||||
import type { Anthropic } from "@anthropic-ai/sdk"
|
||||
import type { ModelMessage } from "ai"
|
||||
|
||||
import {
|
||||
stripPlainTextReasoningBlocks,
|
||||
collectEncryptedReasoningItems,
|
||||
injectEncryptedReasoning,
|
||||
type EncryptedReasoningItem,
|
||||
} from "../openai-native"
|
||||
|
||||
describe("OpenAI Native reasoning helpers", () => {
|
||||
// ───────────────────────────────────────────────────────────
|
||||
// stripPlainTextReasoningBlocks
|
||||
// ───────────────────────────────────────────────────────────
|
||||
describe("stripPlainTextReasoningBlocks", () => {
|
||||
it("passes through user messages unchanged", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hello" }] },
|
||||
]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toEqual(messages)
|
||||
})
|
||||
|
||||
it("passes through assistant messages with only text blocks", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{ role: "assistant", content: [{ type: "text", text: "Hi there" }] },
|
||||
]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toEqual(messages)
|
||||
})
|
||||
|
||||
it("passes through string-content assistant messages", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [{ role: "assistant", content: "Hello" }]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toEqual(messages)
|
||||
})
|
||||
|
||||
it("strips plain-text reasoning blocks from assistant content", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{
|
||||
type: "reasoning",
|
||||
text: "Let me think...",
|
||||
} as unknown as Anthropic.Messages.ContentBlockParam,
|
||||
{ type: "text", text: "The answer is 42" },
|
||||
],
|
||||
},
|
||||
]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toHaveLength(1)
|
||||
expect(result[0].content).toEqual([{ type: "text", text: "The answer is 42" }])
|
||||
})
|
||||
|
||||
it("removes assistant messages whose content becomes empty after filtering", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{
|
||||
type: "reasoning",
|
||||
text: "Thinking only...",
|
||||
} as unknown as Anthropic.Messages.ContentBlockParam,
|
||||
],
|
||||
},
|
||||
]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toHaveLength(0)
|
||||
})
|
||||
|
||||
it("preserves tool_use blocks alongside stripped reasoning", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{ type: "reasoning", text: "Thinking..." } as unknown as Anthropic.Messages.ContentBlockParam,
|
||||
{ type: "tool_use", id: "call_1", name: "read_file", input: { path: "a.ts" } },
|
||||
],
|
||||
},
|
||||
]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toHaveLength(1)
|
||||
expect(result[0].content).toEqual([
|
||||
{ type: "tool_use", id: "call_1", name: "read_file", input: { path: "a.ts" } },
|
||||
])
|
||||
})
|
||||
|
||||
it("does NOT strip blocks that have encrypted_content (those are not plain-text reasoning)", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{
|
||||
type: "reasoning",
|
||||
text: "summary",
|
||||
encrypted_content: "abc123",
|
||||
} as unknown as Anthropic.Messages.ContentBlockParam,
|
||||
{ type: "text", text: "Response" },
|
||||
],
|
||||
},
|
||||
]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toHaveLength(1)
|
||||
// Both blocks should remain
|
||||
expect(result[0].content).toHaveLength(2)
|
||||
})
|
||||
|
||||
it("handles multiple messages correctly", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{ role: "user", content: [{ type: "text", text: "Q1" }] },
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{ type: "reasoning", text: "Think1" } as unknown as Anthropic.Messages.ContentBlockParam,
|
||||
{ type: "text", text: "A1" },
|
||||
],
|
||||
},
|
||||
{ role: "user", content: [{ type: "text", text: "Q2" }] },
|
||||
{
|
||||
role: "assistant",
|
||||
content: [
|
||||
{ type: "reasoning", text: "Think2" } as unknown as Anthropic.Messages.ContentBlockParam,
|
||||
{ type: "text", text: "A2" },
|
||||
],
|
||||
},
|
||||
]
|
||||
const result = stripPlainTextReasoningBlocks(messages)
|
||||
expect(result).toHaveLength(4)
|
||||
expect(result[1].content).toEqual([{ type: "text", text: "A1" }])
|
||||
expect(result[3].content).toEqual([{ type: "text", text: "A2" }])
|
||||
})
|
||||
})
|
||||
|
||||
// ───────────────────────────────────────────────────────────
|
||||
// collectEncryptedReasoningItems
|
||||
// ───────────────────────────────────────────────────────────
|
||||
describe("collectEncryptedReasoningItems", () => {
|
||||
it("returns empty array when no encrypted reasoning items exist", () => {
|
||||
const messages: Anthropic.Messages.MessageParam[] = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hello" }] },
|
||||
{ role: "assistant", content: [{ type: "text", text: "Hi" }] },
|
||||
]
|
||||
const result = collectEncryptedReasoningItems(messages)
|
||||
expect(result).toEqual([])
|
||||
})
|
||||
|
||||
it("collects a single encrypted reasoning item", () => {
|
||||
const messages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hello" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_abc",
|
||||
encrypted_content: "encrypted_data_1",
|
||||
summary: [{ type: "summary_text", text: "I thought about it" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "Hi" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const result = collectEncryptedReasoningItems(messages)
|
||||
expect(result).toHaveLength(1)
|
||||
expect(result[0]).toEqual({
|
||||
id: "rs_abc",
|
||||
encrypted_content: "encrypted_data_1",
|
||||
summary: [{ type: "summary_text", text: "I thought about it" }],
|
||||
originalIndex: 1,
|
||||
})
|
||||
})
|
||||
|
||||
it("collects multiple encrypted reasoning items with correct indices", () => {
|
||||
const messages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Q1" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_1",
|
||||
encrypted_content: "enc_1",
|
||||
summary: [{ type: "summary_text", text: "Summary 1" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "A1" }] },
|
||||
{ role: "user", content: [{ type: "text", text: "Q2" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_2",
|
||||
encrypted_content: "enc_2",
|
||||
summary: [{ type: "summary_text", text: "Summary 2" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "A2" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const result = collectEncryptedReasoningItems(messages)
|
||||
expect(result).toHaveLength(2)
|
||||
expect(result[0].id).toBe("rs_1")
|
||||
expect(result[0].originalIndex).toBe(1)
|
||||
expect(result[1].id).toBe("rs_2")
|
||||
expect(result[1].originalIndex).toBe(4)
|
||||
})
|
||||
|
||||
it("ignores messages that have type 'reasoning' but no encrypted_content", () => {
|
||||
const messages = [
|
||||
{ type: "reasoning", id: "rs_x", text: "plain reasoning" },
|
||||
{ role: "user", content: [{ type: "text", text: "Hello" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const result = collectEncryptedReasoningItems(messages)
|
||||
expect(result).toEqual([])
|
||||
})
|
||||
|
||||
it("handles items without summary", () => {
|
||||
const messages = [
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_no_summary",
|
||||
encrypted_content: "enc_data",
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "Hi" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const result = collectEncryptedReasoningItems(messages)
|
||||
expect(result).toHaveLength(1)
|
||||
expect(result[0].summary).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
// ───────────────────────────────────────────────────────────
|
||||
// injectEncryptedReasoning
|
||||
// ───────────────────────────────────────────────────────────
|
||||
describe("injectEncryptedReasoning", () => {
|
||||
it("does nothing when encryptedItems is empty", () => {
|
||||
const aiSdkMessages: ModelMessage[] = [
|
||||
{ role: "user", content: "Hello" },
|
||||
{ role: "assistant", content: [{ type: "text", text: "Hi" }] },
|
||||
]
|
||||
const original = JSON.parse(JSON.stringify(aiSdkMessages))
|
||||
injectEncryptedReasoning(aiSdkMessages, [], [])
|
||||
expect(aiSdkMessages).toEqual(original)
|
||||
})
|
||||
|
||||
it("injects a single encrypted reasoning part into the next assistant message", () => {
|
||||
// Original messages: [user, encrypted_reasoning, assistant]
|
||||
const originalMessages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hello" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_abc",
|
||||
encrypted_content: "enc_123",
|
||||
summary: [{ type: "summary_text", text: "I considered the question" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "Hi there" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
// AI SDK messages (after filtering encrypted items + converting)
|
||||
const aiSdkMessages: ModelMessage[] = [
|
||||
{ role: "user", content: "Hello" },
|
||||
{ role: "assistant", content: [{ type: "text", text: "Hi there" }] },
|
||||
]
|
||||
|
||||
const encryptedItems: EncryptedReasoningItem[] = [
|
||||
{
|
||||
id: "rs_abc",
|
||||
encrypted_content: "enc_123",
|
||||
summary: [{ type: "summary_text", text: "I considered the question" }],
|
||||
originalIndex: 1,
|
||||
},
|
||||
]
|
||||
|
||||
injectEncryptedReasoning(aiSdkMessages, encryptedItems, originalMessages)
|
||||
|
||||
const assistantMsg = aiSdkMessages[1] as Record<string, unknown>
|
||||
const content = assistantMsg.content as unknown[]
|
||||
expect(content).toHaveLength(2)
|
||||
|
||||
// First part should be the injected reasoning
|
||||
const reasoningPart = content[0] as Record<string, unknown>
|
||||
expect(reasoningPart.type).toBe("reasoning")
|
||||
expect(reasoningPart.text).toBe("I considered the question")
|
||||
|
||||
const providerOptions = reasoningPart.providerOptions as Record<string, Record<string, unknown>>
|
||||
expect(providerOptions.openai.itemId).toBe("rs_abc")
|
||||
expect(providerOptions.openai.reasoningEncryptedContent).toBe("enc_123")
|
||||
|
||||
// Second part should be the original text
|
||||
const textPart = content[1] as Record<string, unknown>
|
||||
expect(textPart.type).toBe("text")
|
||||
expect(textPart.text).toBe("Hi there")
|
||||
})
|
||||
|
||||
it("handles multiple encrypted reasoning items across different assistant messages", () => {
|
||||
const originalMessages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Q1" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_1",
|
||||
encrypted_content: "enc_1",
|
||||
summary: [{ type: "summary_text", text: "Thought 1" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "A1" }] },
|
||||
{ role: "user", content: [{ type: "text", text: "Q2" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_2",
|
||||
encrypted_content: "enc_2",
|
||||
summary: [{ type: "summary_text", text: "Thought 2" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "A2" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const aiSdkMessages: ModelMessage[] = [
|
||||
{ role: "user", content: "Q1" },
|
||||
{ role: "assistant", content: [{ type: "text", text: "A1" }] },
|
||||
{ role: "user", content: "Q2" },
|
||||
{ role: "assistant", content: [{ type: "text", text: "A2" }] },
|
||||
]
|
||||
|
||||
const encryptedItems: EncryptedReasoningItem[] = [
|
||||
{
|
||||
id: "rs_1",
|
||||
encrypted_content: "enc_1",
|
||||
summary: [{ type: "summary_text", text: "Thought 1" }],
|
||||
originalIndex: 1,
|
||||
},
|
||||
{
|
||||
id: "rs_2",
|
||||
encrypted_content: "enc_2",
|
||||
summary: [{ type: "summary_text", text: "Thought 2" }],
|
||||
originalIndex: 4,
|
||||
},
|
||||
]
|
||||
|
||||
injectEncryptedReasoning(aiSdkMessages, encryptedItems, originalMessages)
|
||||
|
||||
// First assistant message
|
||||
const content1 = (aiSdkMessages[1] as Record<string, unknown>).content as unknown[]
|
||||
expect(content1).toHaveLength(2)
|
||||
expect((content1[0] as Record<string, unknown>).type).toBe("reasoning")
|
||||
expect(
|
||||
((content1[0] as Record<string, unknown>).providerOptions as Record<string, Record<string, unknown>>)
|
||||
.openai.itemId,
|
||||
).toBe("rs_1")
|
||||
|
||||
// Second assistant message
|
||||
const content2 = (aiSdkMessages[3] as Record<string, unknown>).content as unknown[]
|
||||
expect(content2).toHaveLength(2)
|
||||
expect((content2[0] as Record<string, unknown>).type).toBe("reasoning")
|
||||
expect(
|
||||
((content2[0] as Record<string, unknown>).providerOptions as Record<string, Record<string, unknown>>)
|
||||
.openai.itemId,
|
||||
).toBe("rs_2")
|
||||
})
|
||||
|
||||
it("joins multiple summary texts with newlines", () => {
|
||||
const originalMessages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hi" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_multi",
|
||||
encrypted_content: "enc_multi",
|
||||
summary: [
|
||||
{ type: "summary_text", text: "First thought" },
|
||||
{ type: "summary_text", text: "Second thought" },
|
||||
],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "Response" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const aiSdkMessages: ModelMessage[] = [
|
||||
{ role: "user", content: "Hi" },
|
||||
{ role: "assistant", content: [{ type: "text", text: "Response" }] },
|
||||
]
|
||||
|
||||
const encryptedItems: EncryptedReasoningItem[] = [
|
||||
{
|
||||
id: "rs_multi",
|
||||
encrypted_content: "enc_multi",
|
||||
summary: [
|
||||
{ type: "summary_text", text: "First thought" },
|
||||
{ type: "summary_text", text: "Second thought" },
|
||||
],
|
||||
originalIndex: 1,
|
||||
},
|
||||
]
|
||||
|
||||
injectEncryptedReasoning(aiSdkMessages, encryptedItems, originalMessages)
|
||||
|
||||
const content = (aiSdkMessages[1] as Record<string, unknown>).content as unknown[]
|
||||
const reasoningPart = content[0] as Record<string, unknown>
|
||||
expect(reasoningPart.text).toBe("First thought\nSecond thought")
|
||||
})
|
||||
|
||||
it("uses empty string when summary is undefined", () => {
|
||||
const originalMessages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hi" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_nosummary",
|
||||
encrypted_content: "enc_nosummary",
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "Response" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const aiSdkMessages: ModelMessage[] = [
|
||||
{ role: "user", content: "Hi" },
|
||||
{ role: "assistant", content: [{ type: "text", text: "Response" }] },
|
||||
]
|
||||
|
||||
const encryptedItems: EncryptedReasoningItem[] = [
|
||||
{
|
||||
id: "rs_nosummary",
|
||||
encrypted_content: "enc_nosummary",
|
||||
summary: undefined,
|
||||
originalIndex: 1,
|
||||
},
|
||||
]
|
||||
|
||||
injectEncryptedReasoning(aiSdkMessages, encryptedItems, originalMessages)
|
||||
|
||||
const content = (aiSdkMessages[1] as Record<string, unknown>).content as unknown[]
|
||||
const reasoningPart = content[0] as Record<string, unknown>
|
||||
expect(reasoningPart.text).toBe("")
|
||||
})
|
||||
|
||||
it("handles consecutive encrypted items before the same assistant message", () => {
|
||||
// Two encrypted reasoning items before one assistant message
|
||||
const originalMessages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hi" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_a",
|
||||
encrypted_content: "enc_a",
|
||||
summary: [{ type: "summary_text", text: "Step A" }],
|
||||
},
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_b",
|
||||
encrypted_content: "enc_b",
|
||||
summary: [{ type: "summary_text", text: "Step B" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "Done" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const aiSdkMessages: ModelMessage[] = [
|
||||
{ role: "user", content: "Hi" },
|
||||
{ role: "assistant", content: [{ type: "text", text: "Done" }] },
|
||||
]
|
||||
|
||||
const encryptedItems: EncryptedReasoningItem[] = [
|
||||
{
|
||||
id: "rs_a",
|
||||
encrypted_content: "enc_a",
|
||||
summary: [{ type: "summary_text", text: "Step A" }],
|
||||
originalIndex: 1,
|
||||
},
|
||||
{
|
||||
id: "rs_b",
|
||||
encrypted_content: "enc_b",
|
||||
summary: [{ type: "summary_text", text: "Step B" }],
|
||||
originalIndex: 2,
|
||||
},
|
||||
]
|
||||
|
||||
injectEncryptedReasoning(aiSdkMessages, encryptedItems, originalMessages)
|
||||
|
||||
const content = (aiSdkMessages[1] as Record<string, unknown>).content as unknown[]
|
||||
// Both reasoning parts should be injected before the text
|
||||
expect(content).toHaveLength(3)
|
||||
expect((content[0] as Record<string, unknown>).type).toBe("reasoning")
|
||||
expect(
|
||||
((content[0] as Record<string, unknown>).providerOptions as Record<string, Record<string, unknown>>)
|
||||
.openai.itemId,
|
||||
).toBe("rs_a")
|
||||
expect((content[1] as Record<string, unknown>).type).toBe("reasoning")
|
||||
expect(
|
||||
((content[1] as Record<string, unknown>).providerOptions as Record<string, Record<string, unknown>>)
|
||||
.openai.itemId,
|
||||
).toBe("rs_b")
|
||||
expect((content[2] as Record<string, unknown>).type).toBe("text")
|
||||
})
|
||||
|
||||
it("handles tool messages splitting (user messages with tool_results create extra tool-role messages)", () => {
|
||||
// Original: [user_with_tool_result, encrypted_reasoning, assistant]
|
||||
// After filtering: [user_with_tool_result, assistant]
|
||||
// AI SDK: [tool, user, assistant] (tool_result split into tool + user messages)
|
||||
const originalMessages = [
|
||||
{
|
||||
role: "user",
|
||||
content: [
|
||||
{ type: "tool_result", tool_use_id: "call_1", content: "result" },
|
||||
{ type: "text", text: "Continue" },
|
||||
],
|
||||
},
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_tool",
|
||||
encrypted_content: "enc_tool",
|
||||
summary: [{ type: "summary_text", text: "Thought after tool" }],
|
||||
},
|
||||
{ role: "assistant", content: [{ type: "text", text: "OK" }] },
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
// AI SDK messages after conversion (tool_result splits into tool + user)
|
||||
const aiSdkMessages: ModelMessage[] = [
|
||||
{
|
||||
role: "tool",
|
||||
content: [
|
||||
{ type: "tool-result", toolCallId: "call_1", toolName: "unknown_tool", result: "result" },
|
||||
],
|
||||
} as unknown as ModelMessage,
|
||||
{ role: "user", content: [{ type: "text", text: "Continue" }] },
|
||||
{ role: "assistant", content: [{ type: "text", text: "OK" }] },
|
||||
]
|
||||
|
||||
const encryptedItems: EncryptedReasoningItem[] = [
|
||||
{
|
||||
id: "rs_tool",
|
||||
encrypted_content: "enc_tool",
|
||||
summary: [{ type: "summary_text", text: "Thought after tool" }],
|
||||
originalIndex: 1,
|
||||
},
|
||||
]
|
||||
|
||||
injectEncryptedReasoning(aiSdkMessages, encryptedItems, originalMessages)
|
||||
|
||||
// The assistant message (index 2) should have the reasoning injected
|
||||
const content = (aiSdkMessages[2] as Record<string, unknown>).content as unknown[]
|
||||
expect(content).toHaveLength(2)
|
||||
expect((content[0] as Record<string, unknown>).type).toBe("reasoning")
|
||||
expect(
|
||||
((content[0] as Record<string, unknown>).providerOptions as Record<string, Record<string, unknown>>)
|
||||
.openai.itemId,
|
||||
).toBe("rs_tool")
|
||||
})
|
||||
|
||||
it("gracefully handles encrypted items with no following assistant message", () => {
|
||||
const originalMessages = [
|
||||
{ role: "user", content: [{ type: "text", text: "Hi" }] },
|
||||
{
|
||||
type: "reasoning",
|
||||
id: "rs_orphan",
|
||||
encrypted_content: "enc_orphan",
|
||||
},
|
||||
] as unknown as Anthropic.Messages.MessageParam[]
|
||||
|
||||
const aiSdkMessages: ModelMessage[] = [{ role: "user", content: "Hi" }]
|
||||
|
||||
const encryptedItems: EncryptedReasoningItem[] = [
|
||||
{
|
||||
id: "rs_orphan",
|
||||
encrypted_content: "enc_orphan",
|
||||
summary: undefined,
|
||||
originalIndex: 1,
|
||||
},
|
||||
]
|
||||
|
||||
// Should not throw
|
||||
expect(() => {
|
||||
injectEncryptedReasoning(aiSdkMessages, encryptedItems, originalMessages)
|
||||
}).not.toThrow()
|
||||
|
||||
// User message unchanged
|
||||
expect(aiSdkMessages).toHaveLength(1)
|
||||
expect(aiSdkMessages[0].role).toBe("user")
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
// npx vitest run api/providers/__tests__/openai-native-tools.spec.ts
|
||||
|
||||
import OpenAI from "openai"
|
||||
|
||||
import { OpenAiHandler } from "../openai"
|
||||
import { OpenAiNativeHandler } from "../openai-native"
|
||||
import type { ApiHandlerOptions } from "../../../shared/api"
|
||||
|
||||
describe("OpenAiHandler native tools", () => {
|
||||
it("includes tools in request when tools are provided via metadata (regression test)", async () => {
|
||||
|
|
@ -68,35 +68,103 @@ describe("OpenAiHandler native tools", () => {
|
|||
})
|
||||
})
|
||||
|
||||
describe("OpenAiNativeHandler MCP tool schema handling", () => {
|
||||
it("should add additionalProperties: false to MCP tools while keeping strict: false", async () => {
|
||||
let capturedRequestBody: any
|
||||
// Use vi.hoisted to define mock functions for AI SDK
|
||||
const { mockStreamText } = vi.hoisted(() => ({
|
||||
mockStreamText: vi.fn(),
|
||||
}))
|
||||
|
||||
vi.mock("ai", async (importOriginal) => {
|
||||
const actual = await importOriginal<typeof import("ai")>()
|
||||
return {
|
||||
...actual,
|
||||
streamText: mockStreamText,
|
||||
generateText: vi.fn(),
|
||||
}
|
||||
})
|
||||
|
||||
vi.mock("@ai-sdk/openai", () => ({
|
||||
createOpenAI: vi.fn(() => {
|
||||
const provider = vi.fn(() => ({
|
||||
modelId: "gpt-4o",
|
||||
provider: "openai",
|
||||
}))
|
||||
;(provider as any).responses = vi.fn(() => ({
|
||||
modelId: "gpt-4o",
|
||||
provider: "openai.responses",
|
||||
}))
|
||||
return provider
|
||||
}),
|
||||
}))
|
||||
|
||||
import { OpenAiNativeHandler } from "../openai-native"
|
||||
import type { ApiHandlerOptions } from "../../../shared/api"
|
||||
|
||||
describe("OpenAiNativeHandler tool handling with AI SDK", () => {
|
||||
function createMockStreamReturn() {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "test" }
|
||||
}
|
||||
|
||||
return {
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 10, outputTokens: 5 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
}
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
|
||||
it("should pass tools through convertToolsForOpenAI and convertToolsForAiSdk to streamText", async () => {
|
||||
mockStreamText.mockReturnValue(createMockStreamReturn())
|
||||
|
||||
const handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-4o",
|
||||
} as ApiHandlerOptions)
|
||||
|
||||
// Mock the responses API call
|
||||
const mockClient = {
|
||||
responses: {
|
||||
create: vi.fn().mockImplementation((body: any) => {
|
||||
capturedRequestBody = body
|
||||
return {
|
||||
[Symbol.asyncIterator]: async function* () {
|
||||
yield {
|
||||
type: "response.done",
|
||||
response: {
|
||||
output: [{ type: "message", content: [{ type: "output_text", text: "test" }] }],
|
||||
usage: { input_tokens: 10, output_tokens: 5 },
|
||||
},
|
||||
}
|
||||
const tools: OpenAI.Chat.ChatCompletionTool[] = [
|
||||
{
|
||||
type: "function",
|
||||
function: {
|
||||
name: "read_file",
|
||||
description: "Read a file from the filesystem",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
path: { type: "string", description: "File path" },
|
||||
},
|
||||
}
|
||||
}),
|
||||
},
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
const stream = handler.createMessage("system prompt", [], {
|
||||
taskId: "test-task-id",
|
||||
tools,
|
||||
})
|
||||
for await (const _ of stream) {
|
||||
// consume
|
||||
}
|
||||
;(handler as any).client = mockClient
|
||||
|
||||
expect(mockStreamText).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
tools: expect.objectContaining({
|
||||
read_file: expect.anything(),
|
||||
}),
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
it("should pass MCP tools to streamText", async () => {
|
||||
mockStreamText.mockReturnValue(createMockStreamReturn())
|
||||
|
||||
const handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-4o",
|
||||
} as ApiHandlerOptions)
|
||||
|
||||
const mcpTools: OpenAI.Chat.ChatCompletionTool[] = [
|
||||
{
|
||||
|
|
@ -119,120 +187,36 @@ describe("OpenAiNativeHandler MCP tool schema handling", () => {
|
|||
taskId: "test-task-id",
|
||||
tools: mcpTools,
|
||||
})
|
||||
|
||||
// Consume the stream
|
||||
for await (const _ of stream) {
|
||||
// Just consume
|
||||
// consume
|
||||
}
|
||||
|
||||
// Verify the request body
|
||||
expect(capturedRequestBody.tools).toBeDefined()
|
||||
expect(capturedRequestBody.tools.length).toBe(1)
|
||||
|
||||
const tool = capturedRequestBody.tools[0]
|
||||
expect(tool.name).toBe("mcp--github--get_me")
|
||||
expect(tool.strict).toBe(false) // MCP tools should have strict: false
|
||||
expect(tool.parameters.additionalProperties).toBe(false) // Should have additionalProperties: false
|
||||
expect(tool.parameters.required).toEqual(["token"]) // Should preserve original required array
|
||||
expect(mockStreamText).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
tools: expect.objectContaining({
|
||||
"mcp--github--get_me": expect.anything(),
|
||||
}),
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
it("should add additionalProperties: false and required array to non-MCP tools with strict: true", async () => {
|
||||
let capturedRequestBody: any
|
||||
it("should pass both regular and MCP tools together", async () => {
|
||||
mockStreamText.mockReturnValue(createMockStreamReturn())
|
||||
|
||||
const handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-4o",
|
||||
} as ApiHandlerOptions)
|
||||
|
||||
// Mock the responses API call
|
||||
const mockClient = {
|
||||
responses: {
|
||||
create: vi.fn().mockImplementation((body: any) => {
|
||||
capturedRequestBody = body
|
||||
return {
|
||||
[Symbol.asyncIterator]: async function* () {
|
||||
yield {
|
||||
type: "response.done",
|
||||
response: {
|
||||
output: [{ type: "message", content: [{ type: "output_text", text: "test" }] }],
|
||||
usage: { input_tokens: 10, output_tokens: 5 },
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
}),
|
||||
},
|
||||
}
|
||||
;(handler as any).client = mockClient
|
||||
|
||||
const regularTools: OpenAI.Chat.ChatCompletionTool[] = [
|
||||
const mixedTools: OpenAI.Chat.ChatCompletionTool[] = [
|
||||
{
|
||||
type: "function",
|
||||
function: {
|
||||
name: "read_file",
|
||||
description: "Read a file from the filesystem",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
path: { type: "string", description: "File path" },
|
||||
encoding: { type: "string", description: "File encoding" },
|
||||
},
|
||||
},
|
||||
description: "Read a file",
|
||||
parameters: { type: "object", properties: { path: { type: "string" } } },
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
const stream = handler.createMessage("system prompt", [], {
|
||||
taskId: "test-task-id",
|
||||
tools: regularTools,
|
||||
})
|
||||
|
||||
// Consume the stream
|
||||
for await (const _ of stream) {
|
||||
// Just consume
|
||||
}
|
||||
|
||||
// Verify the request body
|
||||
expect(capturedRequestBody.tools).toBeDefined()
|
||||
expect(capturedRequestBody.tools.length).toBe(1)
|
||||
|
||||
const tool = capturedRequestBody.tools[0]
|
||||
expect(tool.name).toBe("read_file")
|
||||
expect(tool.strict).toBe(true) // Non-MCP tools should have strict: true
|
||||
expect(tool.parameters.additionalProperties).toBe(false) // Should have additionalProperties: false
|
||||
expect(tool.parameters.required).toEqual(["path", "encoding"]) // Should have all properties as required
|
||||
})
|
||||
|
||||
it("should recursively add additionalProperties: false to nested objects in MCP tools", async () => {
|
||||
let capturedRequestBody: any
|
||||
|
||||
const handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-4o",
|
||||
} as ApiHandlerOptions)
|
||||
|
||||
// Mock the responses API call
|
||||
const mockClient = {
|
||||
responses: {
|
||||
create: vi.fn().mockImplementation((body: any) => {
|
||||
capturedRequestBody = body
|
||||
return {
|
||||
[Symbol.asyncIterator]: async function* () {
|
||||
yield {
|
||||
type: "response.done",
|
||||
response: {
|
||||
output: [{ type: "message", content: [{ type: "output_text", text: "test" }] }],
|
||||
usage: { input_tokens: 10, output_tokens: 5 },
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
}),
|
||||
},
|
||||
}
|
||||
;(handler as any).client = mockClient
|
||||
|
||||
const mcpToolsWithNestedObjects: OpenAI.Chat.ChatCompletionTool[] = [
|
||||
{
|
||||
type: "function",
|
||||
function: {
|
||||
|
|
@ -240,24 +224,8 @@ describe("OpenAiNativeHandler MCP tool schema handling", () => {
|
|||
description: "Create a Linear issue",
|
||||
parameters: {
|
||||
type: "object",
|
||||
properties: {
|
||||
title: { type: "string" },
|
||||
metadata: {
|
||||
type: "object",
|
||||
properties: {
|
||||
priority: { type: "number" },
|
||||
labels: {
|
||||
type: "array",
|
||||
items: {
|
||||
type: "object",
|
||||
properties: {
|
||||
name: { type: "string" },
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
properties: { title: { type: "string" } },
|
||||
required: ["title"],
|
||||
},
|
||||
},
|
||||
},
|
||||
|
|
@ -265,72 +233,64 @@ describe("OpenAiNativeHandler MCP tool schema handling", () => {
|
|||
|
||||
const stream = handler.createMessage("system prompt", [], {
|
||||
taskId: "test-task-id",
|
||||
tools: mcpToolsWithNestedObjects,
|
||||
tools: mixedTools,
|
||||
})
|
||||
|
||||
// Consume the stream
|
||||
for await (const _ of stream) {
|
||||
// Just consume
|
||||
// consume
|
||||
}
|
||||
|
||||
// Verify the request body
|
||||
const tool = capturedRequestBody.tools[0]
|
||||
expect(tool.strict).toBe(false) // MCP tool should have strict: false
|
||||
expect(tool.parameters.additionalProperties).toBe(false) // Root level
|
||||
expect(tool.parameters.properties.metadata.additionalProperties).toBe(false) // Nested object
|
||||
expect(tool.parameters.properties.metadata.properties.labels.items.additionalProperties).toBe(false) // Array items
|
||||
const callArgs = mockStreamText.mock.calls[0][0]
|
||||
expect(callArgs.tools).toBeDefined()
|
||||
expect(callArgs.tools.read_file).toBeDefined()
|
||||
expect(callArgs.tools["mcp--linear--create_issue"]).toBeDefined()
|
||||
})
|
||||
|
||||
it("should handle missing call_id and name in tool_call_arguments.delta by using pending tool identity", async () => {
|
||||
it("should pass parallelToolCalls in provider options", async () => {
|
||||
mockStreamText.mockReturnValue(createMockStreamReturn())
|
||||
|
||||
const handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-4o",
|
||||
} as ApiHandlerOptions)
|
||||
|
||||
const mockClient = {
|
||||
responses: {
|
||||
create: vi.fn().mockImplementation(() => {
|
||||
return {
|
||||
[Symbol.asyncIterator]: async function* () {
|
||||
// 1. Emit output_item.added with tool identity
|
||||
yield {
|
||||
type: "response.output_item.added",
|
||||
item: {
|
||||
type: "function_call",
|
||||
call_id: "call_123",
|
||||
name: "read_file",
|
||||
arguments: "",
|
||||
},
|
||||
}
|
||||
|
||||
// 2. Emit tool_call_arguments.delta WITHOUT identity (just args)
|
||||
yield {
|
||||
type: "response.function_call_arguments.delta",
|
||||
delta: '{"path":',
|
||||
}
|
||||
|
||||
// 3. Emit another delta
|
||||
yield {
|
||||
type: "response.function_call_arguments.delta",
|
||||
delta: '"/tmp/test.txt"}',
|
||||
}
|
||||
|
||||
// 4. Emit output_item.done
|
||||
yield {
|
||||
type: "response.output_item.done",
|
||||
item: {
|
||||
type: "function_call",
|
||||
call_id: "call_123",
|
||||
name: "read_file",
|
||||
arguments: '{"path":"/tmp/test.txt"}',
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
}),
|
||||
},
|
||||
const stream = handler.createMessage("system prompt", [], {
|
||||
taskId: "test-task-id",
|
||||
parallelToolCalls: false,
|
||||
})
|
||||
for await (const _ of stream) {
|
||||
// consume
|
||||
}
|
||||
;(handler as any).client = mockClient
|
||||
|
||||
expect(mockStreamText).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
providerOptions: expect.objectContaining({
|
||||
openai: expect.objectContaining({
|
||||
parallelToolCalls: false,
|
||||
}),
|
||||
}),
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
it("should handle tool call streaming events", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "tool-input-start", id: "call_abc", toolName: "read_file" }
|
||||
yield { type: "tool-input-delta", id: "call_abc", delta: '{"path":' }
|
||||
yield { type: "tool-input-delta", id: "call_abc", delta: '"/tmp/test.txt"}' }
|
||||
yield { type: "tool-input-end", id: "call_abc" }
|
||||
}
|
||||
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 10, outputTokens: 5 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
const handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-4o",
|
||||
} as ApiHandlerOptions)
|
||||
|
||||
const stream = handler.createMessage("system prompt", [], {
|
||||
taskId: "test-task-id",
|
||||
|
|
@ -338,25 +298,19 @@ describe("OpenAiNativeHandler MCP tool schema handling", () => {
|
|||
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
if (chunk.type === "tool_call_partial") {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
expect(chunks.length).toBe(2)
|
||||
expect(chunks[0]).toEqual({
|
||||
type: "tool_call_partial",
|
||||
index: 0,
|
||||
id: "call_123", // Should be filled from pendingToolCallId
|
||||
name: "read_file", // Should be filled from pendingToolCallName
|
||||
arguments: '{"path":',
|
||||
})
|
||||
expect(chunks[1]).toEqual({
|
||||
type: "tool_call_partial",
|
||||
index: 0,
|
||||
id: "call_123",
|
||||
name: "read_file",
|
||||
arguments: '"/tmp/test.txt"}',
|
||||
})
|
||||
const toolStart = chunks.filter((c) => c.type === "tool_call_start")
|
||||
expect(toolStart).toHaveLength(1)
|
||||
expect(toolStart[0].id).toBe("call_abc")
|
||||
expect(toolStart[0].name).toBe("read_file")
|
||||
|
||||
const toolDeltas = chunks.filter((c) => c.type === "tool_call_delta")
|
||||
expect(toolDeltas).toHaveLength(2)
|
||||
|
||||
const toolEnd = chunks.filter((c) => c.type === "tool_call_end")
|
||||
expect(toolEnd).toHaveLength(1)
|
||||
expect(toolEnd[0].id).toBe("call_abc")
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -1,422 +1,355 @@
|
|||
import { describe, it, expect, beforeEach } from "vitest"
|
||||
import { OpenAiNativeHandler } from "../openai-native"
|
||||
// npx vitest run api/providers/__tests__/openai-native-usage.spec.ts
|
||||
|
||||
const { mockStreamText, mockGenerateText } = vi.hoisted(() => ({
|
||||
mockStreamText: vi.fn(),
|
||||
mockGenerateText: vi.fn(),
|
||||
}))
|
||||
|
||||
vi.mock("ai", async (importOriginal) => {
|
||||
const actual = await importOriginal<typeof import("ai")>()
|
||||
return {
|
||||
...actual,
|
||||
streamText: mockStreamText,
|
||||
generateText: mockGenerateText,
|
||||
}
|
||||
})
|
||||
|
||||
vi.mock("@ai-sdk/openai", () => ({
|
||||
createOpenAI: vi.fn(() => {
|
||||
const provider = vi.fn(() => ({
|
||||
modelId: "gpt-4.1",
|
||||
provider: "openai",
|
||||
}))
|
||||
;(provider as any).responses = vi.fn(() => ({
|
||||
modelId: "gpt-4.1",
|
||||
provider: "openai.responses",
|
||||
}))
|
||||
return provider
|
||||
}),
|
||||
}))
|
||||
|
||||
import type { Anthropic } from "@anthropic-ai/sdk"
|
||||
|
||||
import { openAiNativeModels } from "@roo-code/types"
|
||||
|
||||
describe("OpenAiNativeHandler - normalizeUsage", () => {
|
||||
import { OpenAiNativeHandler } from "../openai-native"
|
||||
import type { ApiHandlerOptions } from "../../../shared/api"
|
||||
|
||||
describe("OpenAiNativeHandler - usage metrics", () => {
|
||||
let handler: OpenAiNativeHandler
|
||||
const mockModel = {
|
||||
id: "gpt-4o",
|
||||
info: openAiNativeModels["gpt-4o"],
|
||||
}
|
||||
const systemPrompt = "You are a helpful assistant."
|
||||
const messages: Anthropic.Messages.MessageParam[] = [{ role: "user", content: "Hello!" }]
|
||||
|
||||
beforeEach(() => {
|
||||
handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-4.1",
|
||||
})
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
|
||||
describe("basic token counts", () => {
|
||||
it("should handle basic input and output tokens", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 100, outputTokens: 50 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(usageChunks[0].inputTokens).toBe(100)
|
||||
expect(usageChunks[0].outputTokens).toBe(50)
|
||||
})
|
||||
|
||||
it("should handle zero tokens", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "" }
|
||||
}
|
||||
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 0, outputTokens: 0 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(usageChunks[0].inputTokens).toBe(0)
|
||||
expect(usageChunks[0].outputTokens).toBe(0)
|
||||
})
|
||||
})
|
||||
|
||||
describe("detailed token shapes (Responses API)", () => {
|
||||
it("should handle detailed shapes with cached and miss tokens", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
input_tokens_details: {
|
||||
cached_tokens: 30,
|
||||
cache_miss_tokens: 70,
|
||||
},
|
||||
describe("cache metrics", () => {
|
||||
it("should handle cached input tokens from usage details", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 0, // miss tokens are NOT cache writes
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
details: {
|
||||
cachedInputTokens: 30,
|
||||
},
|
||||
}),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(usageChunks[0].cacheReadTokens).toBe(30)
|
||||
})
|
||||
|
||||
it("should derive total input tokens from details when totals are missing", () => {
|
||||
const usage = {
|
||||
// No input_tokens or prompt_tokens
|
||||
output_tokens: 50,
|
||||
input_tokens_details: {
|
||||
cached_tokens: 30,
|
||||
cache_miss_tokens: 70,
|
||||
},
|
||||
it("should handle no cache information", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100, // Derived from 30 + 70
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 0, // miss tokens are NOT cache writes
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 50, outputTokens: 25 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle prompt_tokens_details variant", () => {
|
||||
const usage = {
|
||||
prompt_tokens: 100,
|
||||
completion_tokens: 50,
|
||||
prompt_tokens_details: {
|
||||
cached_tokens: 30,
|
||||
cache_miss_tokens: 70,
|
||||
},
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 0, // miss tokens are NOT cache writes
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle cache_creation_input_tokens for actual cache writes", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
cache_creation_input_tokens: 20,
|
||||
input_tokens_details: {
|
||||
cached_tokens: 30,
|
||||
cache_miss_tokens: 50, // 50 miss + 30 cached + 20 creation = 100 total
|
||||
},
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 20, // Actual cache writes from cache_creation_input_tokens
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle reasoning tokens in output details", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 150,
|
||||
output_tokens_details: {
|
||||
reasoning_tokens: 50,
|
||||
},
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 150,
|
||||
reasoningTokens: 50,
|
||||
})
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(usageChunks[0].cacheReadTokens).toBeUndefined()
|
||||
expect(usageChunks[0].cacheWriteTokens).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe("legacy field names", () => {
|
||||
it("should handle cache_creation_input_tokens and cache_read_input_tokens", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
cache_creation_input_tokens: 20,
|
||||
cache_read_input_tokens: 30,
|
||||
describe("reasoning tokens", () => {
|
||||
it("should handle reasoning tokens in usage details", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "reasoning-delta", text: "thinking..." }
|
||||
yield { type: "text-delta", text: "answer" }
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 20,
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
details: {
|
||||
reasoningTokens: 30,
|
||||
},
|
||||
}),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle cache_write_tokens and cache_read_tokens", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
cache_write_tokens: 20,
|
||||
cache_read_tokens: 30,
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 20,
|
||||
})
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(usageChunks[0].reasoningTokens).toBe(30)
|
||||
})
|
||||
|
||||
it("should handle cached_tokens field", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
cached_tokens: 30,
|
||||
it("should omit reasoning tokens when not present", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "answer" }
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
}),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle prompt_tokens and completion_tokens", () => {
|
||||
const usage = {
|
||||
prompt_tokens: 100,
|
||||
completion_tokens: 50,
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 0,
|
||||
cacheWriteTokens: 0,
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe("SSE-only events", () => {
|
||||
it("should handle SSE events with minimal usage data", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 0,
|
||||
cacheWriteTokens: 0,
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle SSE events with no cache information", () => {
|
||||
const usage = {
|
||||
prompt_tokens: 100,
|
||||
completion_tokens: 50,
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 0,
|
||||
cacheWriteTokens: 0,
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe("edge cases", () => {
|
||||
it("should handle undefined usage", () => {
|
||||
const result = (handler as any).normalizeUsage(undefined, mockModel)
|
||||
expect(result).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should handle null usage", () => {
|
||||
const result = (handler as any).normalizeUsage(null, mockModel)
|
||||
expect(result).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should handle empty usage object", () => {
|
||||
const usage = {}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 0,
|
||||
outputTokens: 0,
|
||||
cacheReadTokens: 0,
|
||||
cacheWriteTokens: 0,
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle missing details but with cache fields", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
cache_read_input_tokens: 30,
|
||||
// No input_tokens_details
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 0,
|
||||
})
|
||||
})
|
||||
|
||||
it("should use all available cache information with proper fallbacks", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
cached_tokens: 20, // Legacy field (will be used as fallback)
|
||||
input_tokens_details: {
|
||||
cached_tokens: 30, // Detailed shape
|
||||
cache_miss_tokens: 70,
|
||||
},
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
// The implementation uses nullish coalescing, so it will use the first non-nullish value:
|
||||
// cache_read_input_tokens ?? cache_read_tokens ?? cached_tokens ?? cachedFromDetails
|
||||
// Since none of the first two exist, it falls back to cached_tokens (20) before cachedFromDetails
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 20, // From cached_tokens (legacy field comes before details in fallback chain)
|
||||
cacheWriteTokens: 0, // miss tokens are NOT cache writes
|
||||
})
|
||||
})
|
||||
|
||||
it("should use detailed shapes when legacy fields are not present", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
// No cached_tokens legacy field
|
||||
input_tokens_details: {
|
||||
cached_tokens: 30,
|
||||
cache_miss_tokens: 70,
|
||||
},
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 100,
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30, // From details since no legacy field exists
|
||||
cacheWriteTokens: 0, // miss tokens are NOT cache writes
|
||||
})
|
||||
})
|
||||
|
||||
it("should handle totals missing with only partial details", () => {
|
||||
const usage = {
|
||||
// No input_tokens or prompt_tokens
|
||||
output_tokens: 50,
|
||||
input_tokens_details: {
|
||||
cached_tokens: 30,
|
||||
// No cache_miss_tokens
|
||||
},
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
|
||||
expect(result).toMatchObject({
|
||||
type: "usage",
|
||||
inputTokens: 30, // Derived from cached_tokens only
|
||||
outputTokens: 50,
|
||||
cacheReadTokens: 30,
|
||||
cacheWriteTokens: 0,
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe("OpenAiNativeHandler - prompt cache retention", () => {
|
||||
let handler: OpenAiNativeHandler
|
||||
|
||||
beforeEach(() => {
|
||||
handler = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
})
|
||||
})
|
||||
|
||||
const buildRequestBodyForModel = (modelId: string) => {
|
||||
// Force the handler to use the requested model ID
|
||||
;(handler as any).options.apiModelId = modelId
|
||||
const model = handler.getModel()
|
||||
// Minimal formatted input/systemPrompt/verbosity/metadata for building the body
|
||||
return (handler as any).buildRequestBody(model, [], "", model.verbosity, undefined, undefined)
|
||||
}
|
||||
|
||||
it("should set prompt_cache_retention=24h for gpt-5.1 models that support prompt caching", () => {
|
||||
const body = buildRequestBodyForModel("gpt-5.1")
|
||||
expect(body.prompt_cache_retention).toBe("24h")
|
||||
|
||||
const codexBody = buildRequestBodyForModel("gpt-5.1-codex")
|
||||
expect(codexBody.prompt_cache_retention).toBe("24h")
|
||||
|
||||
const codexMiniBody = buildRequestBodyForModel("gpt-5.1-codex-mini")
|
||||
expect(codexMiniBody.prompt_cache_retention).toBe("24h")
|
||||
})
|
||||
|
||||
it("should not set prompt_cache_retention for non-gpt-5.1 models even if they support prompt caching", () => {
|
||||
const body = buildRequestBodyForModel("gpt-5")
|
||||
expect(body.prompt_cache_retention).toBeUndefined()
|
||||
|
||||
const fourOBody = buildRequestBodyForModel("gpt-4o")
|
||||
expect(fourOBody.prompt_cache_retention).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should not set prompt_cache_retention when the model does not support prompt caching", () => {
|
||||
const modelId = "codex-mini-latest"
|
||||
expect(openAiNativeModels[modelId as keyof typeof openAiNativeModels].supportsPromptCache).toBe(false)
|
||||
|
||||
const body = buildRequestBodyForModel(modelId)
|
||||
expect(body.prompt_cache_retention).toBeUndefined()
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(usageChunks[0].reasoningTokens).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe("cost calculation", () => {
|
||||
it("should pass total input tokens to calculateApiCostOpenAI", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
cache_read_input_tokens: 30,
|
||||
cache_creation_input_tokens: 20,
|
||||
it("should include totalCost in usage metrics", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({
|
||||
inputTokens: 1000,
|
||||
outputTokens: 500,
|
||||
}),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
expect(result).toHaveProperty("totalCost")
|
||||
expect(result.totalCost).toBeGreaterThan(0)
|
||||
// calculateApiCostOpenAI handles subtracting cache tokens internally
|
||||
// It will compute: 100 - 30 - 20 = 50 uncached input tokens
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(typeof usageChunks[0].totalCost).toBe("number")
|
||||
expect(usageChunks[0].totalCost).toBeGreaterThanOrEqual(0)
|
||||
})
|
||||
|
||||
it("should handle cost calculation with no cache reads", () => {
|
||||
const usage = {
|
||||
input_tokens: 100,
|
||||
output_tokens: 50,
|
||||
it("should handle all details together", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
const result = (handler as any).normalizeUsage(usage, mockModel)
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({
|
||||
inputTokens: 200,
|
||||
outputTokens: 100,
|
||||
details: {
|
||||
cachedInputTokens: 50,
|
||||
reasoningTokens: 25,
|
||||
},
|
||||
}),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
expect(result).toHaveProperty("totalCost")
|
||||
expect(result.totalCost).toBeGreaterThan(0)
|
||||
// Cost should be calculated with full input tokens since no cache reads
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
const chunks: any[] = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
const usageChunks = chunks.filter((c) => c.type === "usage")
|
||||
expect(usageChunks).toHaveLength(1)
|
||||
expect(usageChunks[0].inputTokens).toBe(200)
|
||||
expect(usageChunks[0].outputTokens).toBe(100)
|
||||
expect(usageChunks[0].cacheReadTokens).toBe(50)
|
||||
expect(usageChunks[0].reasoningTokens).toBe(25)
|
||||
expect(typeof usageChunks[0].totalCost).toBe("number")
|
||||
})
|
||||
})
|
||||
|
||||
describe("prompt cache retention", () => {
|
||||
it("should set promptCacheRetention=24h for gpt-5.1 models that support prompt caching", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 10, outputTokens: 5 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
const h = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "gpt-5.1",
|
||||
})
|
||||
|
||||
const stream = h.createMessage(systemPrompt, messages)
|
||||
for await (const _ of stream) {
|
||||
// consume
|
||||
}
|
||||
|
||||
const callArgs = mockStreamText.mock.calls[0][0]
|
||||
const modelInfo = openAiNativeModels["gpt-5.1"]
|
||||
if (modelInfo.supportsPromptCache && modelInfo.promptCacheRetention === "24h") {
|
||||
expect(callArgs.providerOptions.openai.promptCacheRetention).toBe("24h")
|
||||
}
|
||||
})
|
||||
|
||||
it("should not set promptCacheRetention for non-gpt-5.1 models", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 10, outputTokens: 5 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
const stream = handler.createMessage(systemPrompt, messages)
|
||||
for await (const _ of stream) {
|
||||
// consume
|
||||
}
|
||||
|
||||
const callArgs = mockStreamText.mock.calls[0][0]
|
||||
expect(callArgs.providerOptions.openai.promptCacheRetention).toBeUndefined()
|
||||
})
|
||||
|
||||
it("should not set promptCacheRetention when the model does not support prompt caching", async () => {
|
||||
async function* mockFullStream() {
|
||||
yield { type: "text-delta", text: "Test" }
|
||||
}
|
||||
|
||||
mockStreamText.mockReturnValue({
|
||||
fullStream: mockFullStream(),
|
||||
usage: Promise.resolve({ inputTokens: 10, outputTokens: 5 }),
|
||||
providerMetadata: Promise.resolve({}),
|
||||
content: Promise.resolve([]),
|
||||
})
|
||||
|
||||
// o3-mini doesn't support prompt caching
|
||||
const h = new OpenAiNativeHandler({
|
||||
openAiNativeApiKey: "test-key",
|
||||
apiModelId: "o3-mini-high",
|
||||
})
|
||||
|
||||
const stream = h.createMessage(systemPrompt, messages)
|
||||
for await (const _ of stream) {
|
||||
// consume
|
||||
}
|
||||
|
||||
const callArgs = mockStreamText.mock.calls[0][0]
|
||||
expect(callArgs.providerOptions.openai.promptCacheRetention).toBeUndefined()
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
File diff suppressed because it is too large
Load diff
|
|
@ -458,6 +458,7 @@
|
|||
"@ai-sdk/google": "^3.0.22",
|
||||
"@ai-sdk/google-vertex": "^4.0.45",
|
||||
"@ai-sdk/mistral": "^3.0.19",
|
||||
"@ai-sdk/openai": "^3.0.26",
|
||||
"@ai-sdk/xai": "^3.0.48",
|
||||
"@anthropic-ai/sdk": "^0.37.0",
|
||||
"@anthropic-ai/vertex-sdk": "^0.7.0",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue