From 62f4cb9daf11be4eb09f4b918f327c8e0976d8a3 Mon Sep 17 00:00:00 2001 From: IronRookieCoder Date: Wed, 13 May 2026 16:24:34 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=20CoStrict=20provider?= =?UTF-8?q?=20=E6=8F=90=E5=89=8D=E5=81=9C=E6=AD=A2=E4=B8=8E=20token=20over?= =?UTF-8?q?ride=20=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/costrict/provider/index.test.ts | 275 ++++++++++++++++++++++++++++ src/costrict/provider/index.ts | 129 +++++++++---- 2 files changed, 364 insertions(+), 40 deletions(-) create mode 100644 src/costrict/provider/index.test.ts diff --git a/src/costrict/provider/index.test.ts b/src/costrict/provider/index.test.ts new file mode 100644 index 000000000..eba953416 --- /dev/null +++ b/src/costrict/provider/index.test.ts @@ -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 = {}, +): 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 | null = null +let _mockModelMaxTokens: number | undefined + +mock.module('openai', () => ({ + default: class OpenAI { + chat = { + completions: { + create: async (args: Record) => { + _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 = {}, +) { + _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) + }) +}) diff --git a/src/costrict/provider/index.ts b/src/costrict/provider/index.ts index 8eb2bbe5f..924a387c7 100644 --- a/src/costrict/provider/index.ts +++ b/src/costrict/provider/index.ts @@ -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 = {} - // 跟踪已 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' })