fix: 修复 CoStrict provider 提前停止与 token override 问题
This commit is contained in:
parent
5c2e1918ca
commit
62f4cb9daf
275
src/costrict/provider/index.test.ts
Normal file
275
src/costrict/provider/index.test.ts
Normal file
|
|
@ -0,0 +1,275 @@
|
|||
import { afterEach, beforeEach, describe, expect, mock, test } from 'bun:test'
|
||||
import type { BetaRawMessageStreamEvent } from '@anthropic-ai/sdk/resources/beta/messages/messages.mjs'
|
||||
import type {
|
||||
AssistantMessage,
|
||||
StreamEvent,
|
||||
} from '../../types/message.js'
|
||||
|
||||
function makeMessageStart(
|
||||
overrides: Record<string, any> = {},
|
||||
): BetaRawMessageStreamEvent {
|
||||
return {
|
||||
type: 'message_start',
|
||||
message: {
|
||||
id: 'msg_test',
|
||||
type: 'message',
|
||||
role: 'assistant',
|
||||
content: [],
|
||||
model: 'test-model',
|
||||
stop_reason: null,
|
||||
stop_sequence: null,
|
||||
usage: {
|
||||
input_tokens: 0,
|
||||
output_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
cache_read_input_tokens: 0,
|
||||
},
|
||||
...overrides,
|
||||
},
|
||||
} as any
|
||||
}
|
||||
|
||||
function makeContentBlockStart(
|
||||
index: number,
|
||||
type: 'text' | 'thinking',
|
||||
): BetaRawMessageStreamEvent {
|
||||
return {
|
||||
type: 'content_block_start',
|
||||
index,
|
||||
content_block:
|
||||
type === 'text'
|
||||
? { type: 'text', text: '' }
|
||||
: { type: 'thinking', thinking: '', signature: '' },
|
||||
} as any
|
||||
}
|
||||
|
||||
function makeTextDelta(index: number, text: string): BetaRawMessageStreamEvent {
|
||||
return {
|
||||
type: 'content_block_delta',
|
||||
index,
|
||||
delta: { type: 'text_delta', text },
|
||||
} as any
|
||||
}
|
||||
|
||||
function makeThinkingDelta(
|
||||
index: number,
|
||||
thinking: string,
|
||||
): BetaRawMessageStreamEvent {
|
||||
return {
|
||||
type: 'content_block_delta',
|
||||
index,
|
||||
delta: { type: 'thinking_delta', thinking },
|
||||
} as any
|
||||
}
|
||||
|
||||
function makeContentBlockStop(index: number): BetaRawMessageStreamEvent {
|
||||
return { type: 'content_block_stop', index } as any
|
||||
}
|
||||
|
||||
function makeMessageDelta(
|
||||
stopReason: string,
|
||||
outputTokens: number,
|
||||
): BetaRawMessageStreamEvent {
|
||||
return {
|
||||
type: 'message_delta',
|
||||
delta: { stop_reason: stopReason, stop_sequence: null },
|
||||
usage: { output_tokens: outputTokens },
|
||||
} as any
|
||||
}
|
||||
|
||||
function makeMessageStop(): BetaRawMessageStreamEvent {
|
||||
return { type: 'message_stop' } as any
|
||||
}
|
||||
|
||||
async function* eventStream(events: BetaRawMessageStreamEvent[]) {
|
||||
for (const event of events) yield event
|
||||
}
|
||||
|
||||
let _nextEvents: BetaRawMessageStreamEvent[] = []
|
||||
let _lastCreateArgs: Record<string, any> | null = null
|
||||
let _mockModelMaxTokens: number | undefined
|
||||
|
||||
mock.module('openai', () => ({
|
||||
default: class OpenAI {
|
||||
chat = {
|
||||
completions: {
|
||||
create: async (args: Record<string, any>) => {
|
||||
_lastCreateArgs = args
|
||||
return { [Symbol.asyncIterator]: async function* () {} }
|
||||
},
|
||||
},
|
||||
}
|
||||
},
|
||||
}))
|
||||
|
||||
mock.module('@ant/model-provider', () => ({
|
||||
anthropicMessagesToOpenAI: () => [],
|
||||
anthropicToolsToOpenAI: () => [],
|
||||
anthropicToolChoiceToOpenAI: () => undefined,
|
||||
adaptOpenAIStreamToAnthropic: () => eventStream(_nextEvents),
|
||||
}))
|
||||
|
||||
mock.module('../../utils/messages.js', () => ({
|
||||
normalizeMessagesForAPI: (msgs: any) => msgs,
|
||||
normalizeContentFromAPI: (blocks: any[]) => blocks,
|
||||
createAssistantAPIErrorMessage: (opts: any) => ({
|
||||
type: 'assistant',
|
||||
message: {
|
||||
content: [{ type: 'text', text: opts.content }],
|
||||
apiError: opts.apiError,
|
||||
},
|
||||
uuid: 'error-uuid',
|
||||
timestamp: new Date().toISOString(),
|
||||
}),
|
||||
}))
|
||||
|
||||
mock.module('../../utils/api.js', () => ({
|
||||
toolToAPISchema: async (tool: any) => tool,
|
||||
}))
|
||||
|
||||
mock.module('../../utils/debug.js', () => ({
|
||||
logForDebugging: () => {},
|
||||
}))
|
||||
|
||||
mock.module('../../cost-tracker.js', () => ({
|
||||
addToTotalSessionCost: () => {},
|
||||
}))
|
||||
|
||||
mock.module('../../utils/modelCost.js', () => ({
|
||||
calculateUSDCost: () => 0,
|
||||
}))
|
||||
|
||||
mock.module('../../utils/proxy.js', () => ({
|
||||
getProxyFetchOptions: () => ({}),
|
||||
}))
|
||||
|
||||
mock.module('../../utils/http.js', () => ({
|
||||
getUserAgent: () => 'test-agent',
|
||||
}))
|
||||
|
||||
mock.module('./fetch.js', () => ({
|
||||
createCoStrictFetch: () => fetch,
|
||||
}))
|
||||
|
||||
mock.module('./modelMapping.js', () => ({
|
||||
resolveCoStrictModel: (model: string) => model,
|
||||
}))
|
||||
|
||||
mock.module('./auth.js', () => ({
|
||||
getCoStrictBaseURL: () => 'https://example.test',
|
||||
}))
|
||||
|
||||
mock.module('./credentials.js', () => ({
|
||||
loadCoStrictCredentials: async () => ({
|
||||
access_token: 'token',
|
||||
base_url: 'https://example.test',
|
||||
}),
|
||||
}))
|
||||
|
||||
mock.module('./models.js', () => ({
|
||||
fetchCoStrictModels: async () => [
|
||||
{
|
||||
id: 'test-model',
|
||||
maxTokens: _mockModelMaxTokens,
|
||||
maxTokensKey: 'max_completion_tokens',
|
||||
},
|
||||
],
|
||||
}))
|
||||
|
||||
mock.module('../../services/api/openai/requestBody.js', () => ({
|
||||
isOpenAIThinkingEnabled: (model: string) => model.includes('deepseek'),
|
||||
resolveOpenAIMaxTokens: (
|
||||
upperLimit: number,
|
||||
maxOutputTokensOverride?: number,
|
||||
) => maxOutputTokensOverride ?? upperLimit,
|
||||
}))
|
||||
|
||||
mock.module('../../bootstrap/state.js', () => ({
|
||||
getMainThreadAgentType: () => null,
|
||||
getActiveSkillName: () => null,
|
||||
}))
|
||||
|
||||
mock.module('../../utils/context.js', () => ({
|
||||
getModelMaxOutputTokens: () => ({ upperLimit: 8192, default: 8192 }),
|
||||
}))
|
||||
|
||||
async function runQueryModel(
|
||||
events: BetaRawMessageStreamEvent[],
|
||||
optionsOverrides: Record<string, unknown> = {},
|
||||
) {
|
||||
_nextEvents = events
|
||||
const { queryModelCoStrict } = await import('./index.js')
|
||||
const assistantMessages: AssistantMessage[] = []
|
||||
const streamEvents: StreamEvent[] = []
|
||||
|
||||
const options: any = {
|
||||
model: 'test-model',
|
||||
tools: [],
|
||||
agents: [],
|
||||
querySource: 'main_loop',
|
||||
...optionsOverrides,
|
||||
}
|
||||
|
||||
for await (const item of queryModelCoStrict(
|
||||
[],
|
||||
{ type: 'text', text: '' } as any,
|
||||
[],
|
||||
new AbortController().signal,
|
||||
options,
|
||||
)) {
|
||||
if (item.type === 'assistant') {
|
||||
assistantMessages.push(item as AssistantMessage)
|
||||
} else if (item.type === 'stream_event') {
|
||||
streamEvents.push(item as StreamEvent)
|
||||
}
|
||||
}
|
||||
|
||||
return { assistantMessages, streamEvents }
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
_nextEvents = []
|
||||
_lastCreateArgs = null
|
||||
_mockModelMaxTokens = undefined
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
_nextEvents = []
|
||||
_mockModelMaxTokens = undefined
|
||||
})
|
||||
|
||||
describe('queryModelCoStrict', () => {
|
||||
test('yields exactly one AssistantMessage for thinking + text content', async () => {
|
||||
const events = [
|
||||
makeMessageStart(),
|
||||
makeContentBlockStart(0, 'thinking'),
|
||||
makeThinkingDelta(0, 'let me think'),
|
||||
makeContentBlockStop(0),
|
||||
makeContentBlockStart(1, 'text'),
|
||||
makeTextDelta(1, 'answer'),
|
||||
makeContentBlockStop(1),
|
||||
makeMessageDelta('end_turn', 12),
|
||||
makeMessageStop(),
|
||||
]
|
||||
|
||||
const { assistantMessages } = await runQueryModel(events)
|
||||
|
||||
expect(assistantMessages).toHaveLength(1)
|
||||
expect(assistantMessages[0]!.message.stop_reason).toBe('end_turn')
|
||||
expect(
|
||||
(assistantMessages[0]!.message.content as any[]).map(
|
||||
block => block.type,
|
||||
),
|
||||
).toEqual(['thinking', 'text'])
|
||||
})
|
||||
|
||||
test('preserves explicit max token override over model metadata default', async () => {
|
||||
_mockModelMaxTokens = 16384
|
||||
const events = [makeMessageStart(), makeMessageStop()]
|
||||
|
||||
await runQueryModel(events, { maxOutputTokensOverride: 2048 })
|
||||
|
||||
expect(_lastCreateArgs).not.toBeNull()
|
||||
expect(_lastCreateArgs!.max_completion_tokens).toBe(2048)
|
||||
})
|
||||
})
|
||||
|
|
@ -17,7 +17,12 @@ import type { Options } from '../../services/api/claude.js'
|
|||
import OpenAI from 'openai'
|
||||
import { getProxyFetchOptions } from '../../utils/proxy.js'
|
||||
import { getUserAgent } from '../../utils/http.js'
|
||||
import { anthropicMessagesToOpenAI, anthropicToolsToOpenAI, anthropicToolChoiceToOpenAI, adaptOpenAIStreamToAnthropic } from '@ant/model-provider'
|
||||
import {
|
||||
anthropicMessagesToOpenAI,
|
||||
anthropicToolsToOpenAI,
|
||||
anthropicToolChoiceToOpenAI,
|
||||
adaptOpenAIStreamToAnthropic,
|
||||
} from '@ant/model-provider'
|
||||
import { normalizeMessagesForAPI } from '../../utils/messages.js'
|
||||
import { toolToAPISchema } from '../../utils/api.js'
|
||||
import { logForDebugging } from '../../utils/debug.js'
|
||||
|
|
@ -32,9 +37,16 @@ import { createCoStrictFetch } from './fetch.js'
|
|||
import { resolveCoStrictModel } from './modelMapping.js'
|
||||
import { getCoStrictBaseURL } from './auth.js'
|
||||
import { loadCoStrictCredentials } from './credentials.js'
|
||||
import { isOpenAIThinkingEnabled } from '../../services/api/openai/requestBody.js'
|
||||
import {
|
||||
isOpenAIThinkingEnabled,
|
||||
resolveOpenAIMaxTokens,
|
||||
} from '../../services/api/openai/requestBody.js'
|
||||
import { fetchCoStrictModels } from './models.js'
|
||||
import { getMainThreadAgentType, getActiveSkillName } from '../../bootstrap/state.js'
|
||||
import {
|
||||
getMainThreadAgentType,
|
||||
getActiveSkillName,
|
||||
} from '../../bootstrap/state.js'
|
||||
import { getModelMaxOutputTokens } from '../../utils/context.js'
|
||||
|
||||
/**
|
||||
* CoStrict 查询路径
|
||||
|
|
@ -60,8 +72,8 @@ export async function* queryModelCoStrict(
|
|||
const chatBaseURL = `${baseUrl}/chat-rag/api/v1`
|
||||
|
||||
// 3. 从模型列表获取 maxTokens 相关参数
|
||||
let defaultMaxTokens = getModelMaxOutputTokens(costrictModel).upperLimit
|
||||
let maxTokensParamKey: string = 'max_tokens'
|
||||
let maxTokensValue: number | undefined
|
||||
if (creds?.access_token) {
|
||||
try {
|
||||
const modelList = await fetchCoStrictModels(baseUrl, creds.access_token)
|
||||
|
|
@ -69,13 +81,17 @@ export async function* queryModelCoStrict(
|
|||
if (modelInfo) {
|
||||
maxTokensParamKey = modelInfo.maxTokensKey || 'max_tokens'
|
||||
if (modelInfo.maxTokens != null) {
|
||||
maxTokensValue = modelInfo.maxTokens
|
||||
defaultMaxTokens = modelInfo.maxTokens
|
||||
}
|
||||
}
|
||||
} catch {
|
||||
// 获取模型列表失败,使用默认值
|
||||
}
|
||||
}
|
||||
const maxTokensValue = resolveOpenAIMaxTokens(
|
||||
defaultMaxTokens,
|
||||
options.maxOutputTokensOverride,
|
||||
)
|
||||
|
||||
// 4. 规范化消息
|
||||
const messagesForAPI = normalizeMessagesForAPI(messages, tools)
|
||||
|
|
@ -107,7 +123,7 @@ export async function* queryModelCoStrict(
|
|||
const openaiMessages = anthropicMessagesToOpenAI(
|
||||
messagesForAPI,
|
||||
systemPrompt,
|
||||
{ enableThinking }
|
||||
{ enableThinking },
|
||||
)
|
||||
const openaiTools = anthropicToolsToOpenAI(standardTools)
|
||||
const openaiToolChoice = anthropicToolChoiceToOpenAI(options.toolChoice)
|
||||
|
|
@ -146,12 +162,16 @@ export async function* queryModelCoStrict(
|
|||
}),
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
...(options.temperatureOverride !== undefined && {
|
||||
temperature: options.temperatureOverride,
|
||||
}),
|
||||
...(maxTokensValue != null && {
|
||||
[maxTokensParamKey]: maxTokensValue,
|
||||
[maxTokensParamKey]: maxTokensValue,
|
||||
...(enableThinking && {
|
||||
thinking: { type: 'enabled' },
|
||||
enable_thinking: true,
|
||||
chat_template_kwargs: { thinking: true },
|
||||
}),
|
||||
...(!enableThinking &&
|
||||
options.temperatureOverride !== undefined && {
|
||||
temperature: options.temperatureOverride,
|
||||
}),
|
||||
}
|
||||
const stream = await client.chat.completions.create(
|
||||
requestBody as unknown as OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming,
|
||||
|
|
@ -162,8 +182,6 @@ export async function* queryModelCoStrict(
|
|||
const adaptedStream = adaptOpenAIStreamToAnthropic(stream, costrictModel)
|
||||
|
||||
const contentBlocks: Record<number, any> = {}
|
||||
// 跟踪已 yield 的 assistant messages,用于 message_delta 时回写 usage
|
||||
const yieldedMessages: AssistantMessage[] = []
|
||||
let partialMessage: any
|
||||
let stopReason: string | null = null
|
||||
let usage = {
|
||||
|
|
@ -174,6 +192,48 @@ export async function* queryModelCoStrict(
|
|||
}
|
||||
let ttftMs = 0
|
||||
const start = Date.now()
|
||||
const assembleFinalAssistantOutputs = (): (
|
||||
| AssistantMessage
|
||||
| SystemAPIErrorMessage
|
||||
)[] => {
|
||||
const outputs: (AssistantMessage | SystemAPIErrorMessage)[] = []
|
||||
if (!partialMessage) return outputs
|
||||
|
||||
const allBlocks = Object.keys(contentBlocks)
|
||||
.sort((a, b) => Number(a) - Number(b))
|
||||
.map(k => contentBlocks[Number(k)])
|
||||
.filter(Boolean)
|
||||
|
||||
if (allBlocks.length > 0) {
|
||||
outputs.push({
|
||||
message: {
|
||||
...partialMessage,
|
||||
content: normalizeContentFromAPI(allBlocks, tools, options.agentId),
|
||||
usage,
|
||||
stop_reason: stopReason,
|
||||
stop_sequence: null,
|
||||
},
|
||||
requestId: undefined,
|
||||
type: 'assistant',
|
||||
uuid: randomUUID(),
|
||||
timestamp: new Date().toISOString(),
|
||||
} as AssistantMessage)
|
||||
}
|
||||
|
||||
if (stopReason === 'max_tokens') {
|
||||
outputs.push(
|
||||
createAssistantAPIErrorMessage({
|
||||
content:
|
||||
`Output truncated: response exceeded the ${maxTokensValue} token limit. ` +
|
||||
`Set OPENAI_MAX_TOKENS or CLAUDE_CODE_MAX_OUTPUT_TOKENS to override.`,
|
||||
apiError: 'max_output_tokens',
|
||||
error: 'max_output_tokens',
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
return outputs
|
||||
}
|
||||
|
||||
for await (const event of adaptedStream) {
|
||||
switch (event.type) {
|
||||
|
|
@ -216,44 +276,27 @@ export async function* queryModelCoStrict(
|
|||
break
|
||||
}
|
||||
case 'content_block_stop': {
|
||||
const idx = (event as any).index
|
||||
const block = contentBlocks[idx]
|
||||
if (!block || !partialMessage) break
|
||||
const m: AssistantMessage = {
|
||||
message: {
|
||||
...partialMessage,
|
||||
content: normalizeContentFromAPI([block], tools, options.agentId),
|
||||
usage,
|
||||
},
|
||||
requestId: undefined,
|
||||
type: 'assistant',
|
||||
uuid: randomUUID(),
|
||||
timestamp: new Date().toISOString(),
|
||||
}
|
||||
yieldedMessages.push(m)
|
||||
yield m
|
||||
// Block accumulation is complete; emit one AssistantMessage at
|
||||
// message_stop so reasoning/text/tool blocks stay in a single turn.
|
||||
break
|
||||
}
|
||||
case 'message_delta': {
|
||||
const deltaUsage = (event as any).usage
|
||||
if (deltaUsage) usage = { ...usage, ...deltaUsage }
|
||||
// 回写 usage 到已 yield 的 assistant messages
|
||||
// 与 Anthropic 原生路径 claude.ts:2298 保持一致
|
||||
for (const msg of yieldedMessages) {
|
||||
msg.message.usage = usage
|
||||
}
|
||||
// 记录 stop_reason,回写到最后的 message
|
||||
if ((event as any).delta?.stop_reason != null) {
|
||||
stopReason = (event as any).delta.stop_reason
|
||||
const lastMsg = yieldedMessages[yieldedMessages.length - 1]
|
||||
if (lastMsg) {
|
||||
lastMsg.message.stop_reason = stopReason
|
||||
}
|
||||
}
|
||||
break
|
||||
}
|
||||
case 'message_stop':
|
||||
case 'message_stop': {
|
||||
if (partialMessage) {
|
||||
for (const output of assembleFinalAssistantOutputs()) {
|
||||
yield output
|
||||
}
|
||||
partialMessage = null
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if (
|
||||
|
|
@ -270,6 +313,12 @@ export async function* queryModelCoStrict(
|
|||
...(event.type === 'message_start' ? { ttftMs } : undefined),
|
||||
} as StreamEvent
|
||||
}
|
||||
|
||||
if (partialMessage) {
|
||||
for (const output of assembleFinalAssistantOutputs()) {
|
||||
yield output
|
||||
}
|
||||
}
|
||||
} catch (error) {
|
||||
const errorMsg = error instanceof Error ? error.message : String(error)
|
||||
logForDebugging(`[CoStrict] Error: ${errorMsg}`, { level: 'error' })
|
||||
|
|
|
|||
Loading…
Reference in New Issue
Block a user