fix: 修复 CoStrict provider 提前停止与 token override 问题

This commit is contained in:
IronRookieCoder 2026-05-13 16:24:34 +08:00
parent 5c2e1918ca
commit 62f4cb9daf
2 changed files with 364 additions and 40 deletions

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

View File

@ -17,7 +17,12 @@ import type { Options } from '../../services/api/claude.js'
import OpenAI from 'openai' import OpenAI from 'openai'
import { getProxyFetchOptions } from '../../utils/proxy.js' import { getProxyFetchOptions } from '../../utils/proxy.js'
import { getUserAgent } from '../../utils/http.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 { normalizeMessagesForAPI } from '../../utils/messages.js'
import { toolToAPISchema } from '../../utils/api.js' import { toolToAPISchema } from '../../utils/api.js'
import { logForDebugging } from '../../utils/debug.js' import { logForDebugging } from '../../utils/debug.js'
@ -32,9 +37,16 @@ import { createCoStrictFetch } from './fetch.js'
import { resolveCoStrictModel } from './modelMapping.js' import { resolveCoStrictModel } from './modelMapping.js'
import { getCoStrictBaseURL } from './auth.js' import { getCoStrictBaseURL } from './auth.js'
import { loadCoStrictCredentials } from './credentials.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 { 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 * CoStrict
@ -60,8 +72,8 @@ export async function* queryModelCoStrict(
const chatBaseURL = `${baseUrl}/chat-rag/api/v1` const chatBaseURL = `${baseUrl}/chat-rag/api/v1`
// 3. 从模型列表获取 maxTokens 相关参数 // 3. 从模型列表获取 maxTokens 相关参数
let defaultMaxTokens = getModelMaxOutputTokens(costrictModel).upperLimit
let maxTokensParamKey: string = 'max_tokens' let maxTokensParamKey: string = 'max_tokens'
let maxTokensValue: number | undefined
if (creds?.access_token) { if (creds?.access_token) {
try { try {
const modelList = await fetchCoStrictModels(baseUrl, creds.access_token) const modelList = await fetchCoStrictModels(baseUrl, creds.access_token)
@ -69,13 +81,17 @@ export async function* queryModelCoStrict(
if (modelInfo) { if (modelInfo) {
maxTokensParamKey = modelInfo.maxTokensKey || 'max_tokens' maxTokensParamKey = modelInfo.maxTokensKey || 'max_tokens'
if (modelInfo.maxTokens != null) { if (modelInfo.maxTokens != null) {
maxTokensValue = modelInfo.maxTokens defaultMaxTokens = modelInfo.maxTokens
} }
} }
} catch { } catch {
// 获取模型列表失败,使用默认值 // 获取模型列表失败,使用默认值
} }
} }
const maxTokensValue = resolveOpenAIMaxTokens(
defaultMaxTokens,
options.maxOutputTokensOverride,
)
// 4. 规范化消息 // 4. 规范化消息
const messagesForAPI = normalizeMessagesForAPI(messages, tools) const messagesForAPI = normalizeMessagesForAPI(messages, tools)
@ -107,7 +123,7 @@ export async function* queryModelCoStrict(
const openaiMessages = anthropicMessagesToOpenAI( const openaiMessages = anthropicMessagesToOpenAI(
messagesForAPI, messagesForAPI,
systemPrompt, systemPrompt,
{ enableThinking } { enableThinking },
) )
const openaiTools = anthropicToolsToOpenAI(standardTools) const openaiTools = anthropicToolsToOpenAI(standardTools)
const openaiToolChoice = anthropicToolChoiceToOpenAI(options.toolChoice) const openaiToolChoice = anthropicToolChoiceToOpenAI(options.toolChoice)
@ -146,12 +162,16 @@ export async function* queryModelCoStrict(
}), }),
stream: true, stream: true,
stream_options: { include_usage: true }, stream_options: { include_usage: true },
...(options.temperatureOverride !== undefined && { [maxTokensParamKey]: maxTokensValue,
temperature: options.temperatureOverride, ...(enableThinking && {
}), thinking: { type: 'enabled' },
...(maxTokensValue != null && { enable_thinking: true,
[maxTokensParamKey]: maxTokensValue, chat_template_kwargs: { thinking: true },
}), }),
...(!enableThinking &&
options.temperatureOverride !== undefined && {
temperature: options.temperatureOverride,
}),
} }
const stream = await client.chat.completions.create( const stream = await client.chat.completions.create(
requestBody as unknown as OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming, requestBody as unknown as OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming,
@ -162,8 +182,6 @@ export async function* queryModelCoStrict(
const adaptedStream = adaptOpenAIStreamToAnthropic(stream, costrictModel) const adaptedStream = adaptOpenAIStreamToAnthropic(stream, costrictModel)
const contentBlocks: Record<number, any> = {} const contentBlocks: Record<number, any> = {}
// 跟踪已 yield 的 assistant messages用于 message_delta 时回写 usage
const yieldedMessages: AssistantMessage[] = []
let partialMessage: any let partialMessage: any
let stopReason: string | null = null let stopReason: string | null = null
let usage = { let usage = {
@ -174,6 +192,48 @@ export async function* queryModelCoStrict(
} }
let ttftMs = 0 let ttftMs = 0
const start = Date.now() 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) { for await (const event of adaptedStream) {
switch (event.type) { switch (event.type) {
@ -216,44 +276,27 @@ export async function* queryModelCoStrict(
break break
} }
case 'content_block_stop': { case 'content_block_stop': {
const idx = (event as any).index // Block accumulation is complete; emit one AssistantMessage at
const block = contentBlocks[idx] // message_stop so reasoning/text/tool blocks stay in a single turn.
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
break break
} }
case 'message_delta': { case 'message_delta': {
const deltaUsage = (event as any).usage const deltaUsage = (event as any).usage
if (deltaUsage) usage = { ...usage, ...deltaUsage } 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) { if ((event as any).delta?.stop_reason != null) {
stopReason = (event as any).delta.stop_reason stopReason = (event as any).delta.stop_reason
const lastMsg = yieldedMessages[yieldedMessages.length - 1]
if (lastMsg) {
lastMsg.message.stop_reason = stopReason
}
} }
break break
} }
case 'message_stop': case 'message_stop': {
if (partialMessage) {
for (const output of assembleFinalAssistantOutputs()) {
yield output
}
partialMessage = null
}
break break
}
} }
if ( if (
@ -270,6 +313,12 @@ export async function* queryModelCoStrict(
...(event.type === 'message_start' ? { ttftMs } : undefined), ...(event.type === 'message_start' ? { ttftMs } : undefined),
} as StreamEvent } as StreamEvent
} }
if (partialMessage) {
for (const output of assembleFinalAssistantOutputs()) {
yield output
}
}
} catch (error) { } catch (error) {
const errorMsg = error instanceof Error ? error.message : String(error) const errorMsg = error instanceof Error ? error.message : String(error)
logForDebugging(`[CoStrict] Error: ${errorMsg}`, { level: 'error' }) logForDebugging(`[CoStrict] Error: ${errorMsg}`, { level: 'error' })