From 2e489d5bc31e1ba6394f7e136cb2e5dc85c68cb6 Mon Sep 17 00:00:00 2001 From: lofyer Date: Wed, 12 Aug 2026 23:29:37 +0800 Subject: [PATCH] fix: stream direct-model reasoning with tools --- src/main/agent/model-runtime.test.ts | 159 ++++++++++++++- src/main/agent/model-runtime.ts | 290 +++++++++++++++++++++++++-- src/renderer/src/App.test.tsx | 41 ++++ src/renderer/src/App.tsx | 1 + 4 files changed, 469 insertions(+), 22 deletions(-) diff --git a/src/main/agent/model-runtime.test.ts b/src/main/agent/model-runtime.test.ts index 82cdaed..cbb6f8c 100644 --- a/src/main/agent/model-runtime.test.ts +++ b/src/main/agent/model-runtime.test.ts @@ -759,7 +759,10 @@ describe('ModelAgentRuntime', () => { fetcher.mock.calls[0]?.[1]?.body as string ) as Record expect(firstBody).toMatchObject({ - stream: false, + stream: true, + stream_options: { + include_usage: true + }, tools: [ { type: 'function', @@ -840,6 +843,160 @@ describe('ModelAgentRuntime', () => { expect(toolProvider.dispose).toHaveBeenCalledOnce() }) + it('streams reasoning while using OpenAI-compatible tools', async () => { + const streams = [ + [ + `data: ${JSON.stringify({ + choices: [ + { + delta: { + reasoning_content: '先读取文件', + tool_calls: [ + { + index: 0, + id: 'call-streamed', + type: 'function', + function: { + name: 'workspace_read_text', + arguments: '{"path":' + } + } + ] + } + } + ] + })}`, + '', + `data: ${JSON.stringify({ + choices: [ + { + delta: { + tool_calls: [ + { + index: 0, + function: { + arguments: '"README.md"}' + } + } + ] + } + } + ] + })}`, + '', + 'data: [DONE]', + '', + '' + ].join('\n'), + [ + `data: ${JSON.stringify({ + choices: [ + { + delta: { + reasoning_content: '再整理结果' + } + } + ] + })}`, + '', + `data: ${JSON.stringify({ + choices: [ + { + delta: { + content: '文件内容已读取。' + } + } + ] + })}`, + '', + 'data: [DONE]', + '', + '' + ].join('\n') + ] + const fetcher = vi.fn(async () => + new Response(streams.shift(), { + status: 200, + headers: { 'content-type': 'text/event-stream' } + }) + ) + const toolProvider = createToolProvider() + const runtime = new ModelAgentRuntime({ + apiKey: 'test-key', + baseUrl: 'https://api.deepseek.com', + model: 'deepseek-v4-flash', + protocol: 'openai-chat-completions', + authentication: 'api-key', + fetcher, + toolProvider + }) + const events = [] + + for await (const event of runtime.run( + { + requestId: 'a431666e-5ec8-45e6-beb4-654132eed140', + conversationId: 'conversation-streamed-tools', + prompt: '读取 README', + workMode: 'execute' + }, + new AbortController().signal, + async () => 'once' + )) { + events.push(event) + } + + expect( + events + .filter( + (event) => + event.type === 'reasoning' || + event.type === 'tool' || + event.type === 'text' + ) + .map((event) => + event.type === 'tool' + ? `${event.type}:${event.state}` + : `${event.type}:${event.delta}` + ) + ).toEqual([ + 'reasoning:先读取文件', + 'tool:pending', + 'tool:running', + 'tool:completed', + 'reasoning:再整理结果', + 'text:文件内容已读取。' + ]) + expect(toolProvider.callTool).toHaveBeenCalledWith( + 'workspace_read_text', + { path: 'README.md' }, + expect.any(AbortSignal), + expect.objectContaining({ + conversationId: 'conversation-streamed-tools', + workMode: 'execute' + }) + ) + const secondBody = JSON.parse( + fetcher.mock.calls[1]?.[1]?.body as string + ) as { messages: Array> } + expect(secondBody.messages).toContainEqual( + expect.objectContaining({ + role: 'assistant', + content: null, + reasoning_content: '先读取文件', + tool_calls: [ + expect.objectContaining({ + id: 'call-streamed', + function: { + name: 'workspace_read_text', + arguments: '{"path":"README.md"}' + } + }) + ] + }) + ) + expect(events.at(-1)).toMatchObject({ type: 'done' }) + }) + it('uses refreshed tool definitions in subsequent model rounds', async () => { const loadTool: ModelToolDefinition = { name: 'mcp_load_tools', diff --git a/src/main/agent/model-runtime.ts b/src/main/agent/model-runtime.ts index dbb6750..7a0ddee 100644 --- a/src/main/agent/model-runtime.ts +++ b/src/main/agent/model-runtime.ts @@ -92,6 +92,7 @@ type ModelToolResponse = { assistantMessage?: Record responsesOutput?: Array> usage: ModelUsageUpdate + streamed?: boolean } const maxGeneratedImageBytes = 3_900_000 @@ -886,6 +887,9 @@ function parseModelToolResponse( assistantMessage: { role: 'assistant', content: message.content ?? null, + ...(reasoning + ? { reasoning_content: reasoning } + : {}), ...(toolCalls.length > 0 ? { tool_calls: message.tool_calls } : {}) @@ -978,6 +982,27 @@ function parseStreamBlock( } } +function parseSseData( + block: string +): { event?: unknown; stopped: boolean } { + const data = block + .split('\n') + .filter((line) => line.startsWith('data:')) + .map((line) => line.slice(5).trimStart()) + .join('\n') + if (!data) { + return { stopped: false } + } + if (data === '[DONE]') { + return { stopped: true } + } + try { + return { event: JSON.parse(data), stopped: false } + } catch { + return { stopped: false } + } +} + export class ModelAgentRuntime implements AgentRuntime { readonly runtimeId = 'model' readonly requiresToolApproval = false @@ -1331,14 +1356,16 @@ export class ModelAgentRuntime implements AgentRuntime { } } - private async requestToolModel( + private async *requestToolModel( messages: Array>, tools: ModelToolDefinition[], system: string, anthropic: boolean, - signal: AbortSignal - ): Promise { + signal: AbortSignal, + requestId: string + ): AsyncGenerator { const responses = this.options.protocol === 'openai-responses' + const streamOpenAIChat = !responses && !anthropic const providerTools = responses ? tools.map((tool) => ({ type: 'function', @@ -1383,7 +1410,10 @@ export class ModelAgentRuntime implements AgentRuntime { : { model: this.options.model, max_tokens: 4096, - stream: false, + stream: true, + stream_options: { + include_usage: true + }, messages, tools: providerTools } @@ -1397,9 +1427,222 @@ export class ModelAgentRuntime implements AgentRuntime { body, signal }) + if (!response.ok) { + const responseText = await readBoundedText( + response, + 128 * 1024 + ) + let detail: string | undefined + try { + detail = getErrorMessage( + responseText.trim() + ? JSON.parse(responseText) + : undefined + ) + } catch { + detail = undefined + } + throw new Error( + detail ?? + `模型接口请求失败(HTTP ${response.status})` + ) + } + if ( + streamOpenAIChat && + response.headers + .get('content-type') + ?.toLocaleLowerCase() + .includes('text/event-stream') + ) { + if (!response.body) { + throw new Error('模型接口未返回流式响应') + } + const reader = response.body.getReader() + const decoder = new TextDecoder() + const streamedToolCalls = new Map< + number, + { arguments: string; id: string; name: string } + >() + const usage: ModelUsageAccumulator = { + reported: false + } + let answer = '' + let reasoning = '' + let buffer = '' + let receivedStop = false + let receivedBytes = 0 + let streamEnded = false + + try { + while (!receivedStop) { + const { done, value } = await reader.read() + streamEnded = done + receivedBytes += value?.byteLength ?? 0 + if (receivedBytes > maxChatResponseBytes) { + throw new Error('模型接口流式响应超过安全限制') + } + buffer += decoder.decode(value, { stream: !done }).replaceAll( + '\r\n', + '\n' + ) + if (Buffer.byteLength(buffer) > maxChatResponseBytes) { + throw new Error('模型接口流式响应块超过安全限制') + } + const blocks = buffer.split('\n\n') + buffer = blocks.pop() ?? '' + if (done && buffer.trim()) { + blocks.push(buffer) + buffer = '' + } + for (const block of blocks) { + const parsed = parseSseData(block) + if (parsed.stopped) { + receivedStop = true + break + } + if (parsed.event === undefined) { + continue + } + const providerError = getErrorMessage(parsed.event) + if (providerError) { + throw new Error(providerError) + } + applyUsageUpdate( + usage, + getUsageUpdate(parsed.event, 'openai') + ) + const reasoningDelta = getOpenAIReasoningDelta( + parsed.event + ) + if (reasoningDelta) { + reasoning += reasoningDelta + yield { + requestId, + type: 'reasoning', + delta: reasoningDelta + } + } + const textDelta = getOpenAITextDelta(parsed.event) + if (textDelta) { + answer += textDelta + yield { + requestId, + type: 'text', + delta: textDelta + } + } + const event = getRecord(parsed.event) + const firstChoice = Array.isArray(event?.choices) + ? getRecord(event.choices[0]) + : undefined + const delta = getRecord(firstChoice?.delta) + if (delta?.tool_calls === undefined) { + continue + } + if (!Array.isArray(delta.tool_calls)) { + throw new Error( + 'OpenAI 模型接口返回了无效流式工具调用' + ) + } + for (const item of delta.tool_calls) { + const toolDelta = getRecord(item) + const index = toolDelta?.index + if ( + !Number.isSafeInteger(index) || + (index as number) < 0 || + (index as number) >= maxToolCallsPerRun + ) { + throw new Error( + 'OpenAI 模型接口返回了无效流式工具调用序号' + ) + } + const functionDelta = getRecord(toolDelta?.function) + const current = streamedToolCalls.get(index as number) ?? { + arguments: '', + id: '', + name: '' + } + const next = { + arguments: + current.arguments + + (typeof functionDelta?.arguments === 'string' + ? functionDelta.arguments + : ''), + id: + typeof toolDelta?.id === 'string' + ? toolDelta.id + : current.id, + name: + typeof functionDelta?.name === 'string' + ? functionDelta.name + : current.name + } + if ( + next.id.length > 256 || + next.name.length > 128 || + Buffer.byteLength(next.arguments) > + maxToolArgumentBytes + ) { + throw new Error( + 'OpenAI 模型接口返回的流式工具调用超过安全限制' + ) + } + streamedToolCalls.set(index as number, next) + } + } + if (done) { + break + } + } + } finally { + if (!streamEnded) { + await reader.cancel().catch(() => undefined) + } + reader.releaseLock() + } + if (!receivedStop) { + throw new Error('模型接口流式响应意外中断') + } + const rawToolCalls = [...streamedToolCalls.entries()] + .sort(([left], [right]) => left - right) + .map(([, call]) => { + const identity = parseToolCallIdentity(call.id, call.name) + return { + parsed: { + ...identity, + arguments: parseToolArguments(call.arguments) + }, + raw: { + id: identity.id, + type: 'function', + function: { + name: identity.name, + arguments: call.arguments + } + } + } + }) + return { + text: answer, + reasoning, + toolCalls: rawToolCalls.map((call) => call.parsed), + assistantMessage: { + role: 'assistant', + content: answer || null, + ...(reasoning + ? { reasoning_content: reasoning } + : {}), + ...(rawToolCalls.length > 0 + ? { tool_calls: rawToolCalls.map((call) => call.raw) } + : {}) + }, + usage, + streamed: true + } + } const responseText = await readBoundedText( response, - response.ok ? maxChatResponseBytes : 128 * 1024 + maxChatResponseBytes ) let payload: unknown try { @@ -1409,12 +1652,6 @@ export class ModelAgentRuntime implements AgentRuntime { } catch (error) { throw new Error('模型接口返回了无效 JSON', { cause: error }) } - if (!response.ok) { - throw new Error( - getErrorMessage(payload) ?? - `模型接口请求失败(HTTP ${response.status})` - ) - } const providerError = getErrorMessage(payload) if (providerError) { throw new Error(providerError) @@ -1496,13 +1733,20 @@ export class ModelAgentRuntime implements AgentRuntime { if (round > 0) { toolSnapshot = await loadToolSnapshot() } - const response = await this.requestToolModel( + const responseStream = this.requestToolModel( messages, toolSnapshot.tools, system, anthropic, - signal + signal, + request.requestId ) + let responseStep = await responseStream.next() + while (!responseStep.done) { + yield responseStep.value + responseStep = await responseStream.next() + } + const response = responseStep.value const usage = { reported: false } satisfies ModelUsageAccumulator @@ -1517,10 +1761,12 @@ export class ModelAgentRuntime implements AgentRuntime { yield usageEvent } if (response.reasoning) { - yield { - requestId: request.requestId, - type: 'reasoning', - delta: response.reasoning + if (!response.streamed) { + yield { + requestId: request.requestId, + type: 'reasoning', + delta: response.reasoning + } } } if (response.text) { @@ -1528,10 +1774,12 @@ export class ModelAgentRuntime implements AgentRuntime { if (Buffer.byteLength(answer) > 1024 * 1024) { throw new Error('直连模型回答超过 1MB 安全限制') } - yield { - requestId: request.requestId, - type: 'text', - delta: response.text + if (!response.streamed) { + yield { + requestId: request.requestId, + type: 'text', + delta: response.text + } } } if (response.toolCalls.length === 0) { diff --git a/src/renderer/src/App.test.tsx b/src/renderer/src/App.test.tsx index 3e3eed3..dec34de 100644 --- a/src/renderer/src/App.test.tsx +++ b/src/renderer/src/App.test.tsx @@ -1423,6 +1423,47 @@ describe('App', () => { expect(screen.getByText('项目:默认项目')).toHaveClass('scope-badge') }) + it('replaces the direct-model thinking status with real reasoning', async () => { + render() + + fireEvent.change(screen.getByLabelText('向 GoodBuddy 提问'), { + target: { value: '分析这个问题' } + }) + fireEvent.click(await screen.findByLabelText('发送')) + await waitFor(() => expect(run).toHaveBeenCalledOnce()) + const request = run.mock.calls[0]?.[0] + if (!request) { + throw new Error('Missing request') + } + + act(() => { + agentListener?.({ + requestId: request.requestId, + type: 'status', + message: 'deepseek-v4-flash 正在思考' + }) + }) + expect( + screen.getByText('deepseek-v4-flash 正在思考') + ).toBeInTheDocument() + + act(() => { + agentListener?.({ + requestId: request.requestId, + type: 'reasoning', + delta: '正在分析真实推理内容' + }) + }) + + expect( + screen.queryByText('deepseek-v4-flash 正在思考') + ).not.toBeInTheDocument() + expect(screen.getByText('正在推理').closest('details')).toHaveAttribute( + 'open' + ) + expect(screen.getByText('正在分析真实推理内容')).toBeVisible() + }) + it('keeps a tool failure in details and hides retry after continuing', async () => { render() diff --git a/src/renderer/src/App.tsx b/src/renderer/src/App.tsx index e9522c0..4e93fbb 100644 --- a/src/renderer/src/App.tsx +++ b/src/renderer/src/App.tsx @@ -2511,6 +2511,7 @@ function App(): React.JSX.Element { return { ...message, reasoning: `${currentReasoning}${acceptedDelta}`, + status: undefined, blocks: appendMessageContentBlock( message.blocks, 'reasoning',