diff --git a/internal/agent/api/provider.go b/internal/agent/api/provider.go index ca89b06..dccbd46 100644 --- a/internal/agent/api/provider.go +++ b/internal/agent/api/provider.go @@ -186,10 +186,11 @@ type TokenUsage struct { } type ToolCall struct { - ID string `json:"id"` - Type string `json:"type"` - Name string `json:"name"` - Arguments map[string]interface{} `json:"arguments"` + ID string `json:"id"` + Type string `json:"type"` + Name string `json:"name"` + Arguments map[string]interface{} `json:"arguments"` + RawArguments string `json:"raw_arguments,omitempty"` // 流式分片原始 JSON 字符串 } type apiToolCall struct { @@ -490,10 +491,52 @@ func normalizeOpenAIToolCalls(raw []openAIToolCall) []ToolCall { typ = "function" } out = append(out, ToolCall{ - ID: tc.ID, - Type: typ, - Name: name, - Arguments: parseToolArguments(argsRaw), + ID: tc.ID, + Type: typ, + Name: name, + Arguments: parseToolArguments(argsRaw), + RawArguments: rawArgsString(argsRaw), + }) + } + return out +} + +// rawArgsString 将 arguments 字段转为字符串形式(用于流式分片拼接)。 +func rawArgsString(v interface{}) string { + switch x := v.(type) { + case nil: + return "" + case string: + return x + default: + b, _ := json.Marshal(x) + return string(b) + } +} + +// normalizeStreamToolCalls 流式专用:保留无 name 的分片(后续 arguments +// 分片 name 为空,但携带 RawArguments 需要拼接),由调用方按 index 累积。 +func normalizeStreamToolCalls(raw []openAIToolCall) []ToolCall { + if len(raw) == 0 { + return nil + } + out := make([]ToolCall, 0, len(raw)) + for _, tc := range raw { + name := tc.Function.Name + argsRaw := tc.Function.Arguments + if name == "" { + name = tc.Name + argsRaw = tc.Arguments + } + typ := tc.Type + if typ == "" && (tc.ID != "" || name != "" || argsRaw != nil) { + typ = "function" + } + out = append(out, ToolCall{ + ID: tc.ID, + Type: typ, + Name: name, + RawArguments: rawArgsString(argsRaw), }) } return out @@ -603,7 +646,7 @@ func parseOpenAICompatibleStreamChunkFull(data string) (StreamChunk, bool) { ck := StreamChunk{ Content: stringifyContent(choice.Delta.Content), ReasoningContent: choice.Delta.ReasoningContent, - ToolCalls: normalizeOpenAIToolCalls(choice.Delta.ToolCalls), + ToolCalls: normalizeStreamToolCalls(choice.Delta.ToolCalls), Usage: usage, } // finish reason 为空字符串不算终止信号(sensenova 每块都发 "") diff --git a/internal/agent/core/process.go b/internal/agent/core/process.go index 749f89d..c43a5d5 100644 --- a/internal/agent/core/process.go +++ b/internal/agent/core/process.go @@ -2,6 +2,7 @@ package core import ( "context" + "encoding/json" "errors" "fmt" "log" @@ -120,7 +121,7 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri a.cancelLLM = fCancel a.llmMu.Unlock() - resp, llmErr = fbProvider.Chat(fCtx, req) + resp, llmErr = chatStreamWithFallback(fCtx, fbProvider, req, a) a.llmMu.Lock() a.cancelLLM = nil @@ -294,6 +295,155 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri } } +// chatStreamWithFallback 优先流式调用 provider,失败时回退非流式 Chat()。 +// +// 流式路径:ChatStream 拿到 chunk channel,逐块累积 content/reasoning_content, +// 并发布 EventReasoningDelta / EventContentDelta 增量事件(新订阅者可选订, +// 旧订阅者不认识自然忽略)。流结束后拼出与 Chat() 等价的 CompletionResponse +// 返回——process() 的后续逻辑(stageCtx/聚合事件/工具循环)完全不变。 +// +// 回退条件:ChatStream 返回错误(连接失败、provider 不支持流式)。 +// 已收到部分 chunk 后出错则不回退(避免重复生成),直接返回已累积内容。 +// +// 超时收益:首包 ~1-3s 到达即建立活性,后续只要 token 在流动就不会触发 +// 空闲超时;总生成时长不再受限於 180s 整体超时。 +func chatStreamWithFallback(ctx context.Context, p agentAPI.Provider, req *agentAPI.CompletionRequest, a *Agent) (*agentAPI.CompletionResponse, error) { + ch, err := p.ChatStream(ctx, req) + if err != nil { + log.Printf("[agent] stream connect failed (%v), falling back to non-stream chat", err) + return p.Chat(ctx, req) + } + + resp, accErr := accumulateStream(ctx, ch, a) + if accErr == nil { + return resp, nil + } + + // 流中途错误:若已累积到内容则返回部分结果,否则回退非流式 + if resp != nil && (resp.Content != "" || len(resp.ToolCalls) > 0) { + log.Printf("[agent] stream interrupted mid-way (%v), returning partial result", accErr) + return resp, nil + } + log.Printf("[agent] stream failed before content (%v), falling back to non-stream chat", accErr) + return p.Chat(ctx, req) +} + +// toolCallAcc 累积流式 tool call 的各个分片。OpenAI 风格:每个 index 的 +// id/name/arguments 跨多个 chunk 增量到达,arguments 是 JSON 字符串分片。 +type toolCallAcc struct { + id string + name string + argsRaw strings.Builder +} + +// accumulateStream 消费 chunk channel,累积为完整 CompletionResponse, +// 同时发布增量事件。返回的 response 与非流式 Chat() 的返回等价。 +func accumulateStream(ctx context.Context, ch <-chan agentAPI.StreamChunk, a *Agent) (*agentAPI.CompletionResponse, error) { + resp := &agentAPI.CompletionResponse{ + ToolCalls: make([]agentAPI.ToolCall, 0), + } + accs := make(map[int]*toolCallAcc) // index → 累积中的 tool call + var lastFinish string + + flushToolCall := func(idx int) { + acc := accs[idx] + if acc == nil { + return + } + if acc.name == "" { + delete(accs, idx) + return + } + tc := agentAPI.ToolCall{ + ID: acc.id, + Name: acc.name, + Arguments: parseToolArgsJSON(acc.argsRaw.String()), + } + resp.ToolCalls = append(resp.ToolCalls, tc) + delete(accs, idx) + } + + for { + select { + case ck, ok := <-ch: + if !ok { + for idx := range accs { + flushToolCall(idx) + } + if lastFinish != "" { + resp.FinishReason = lastFinish + } + return resp, nil + } + + if ck.ReasoningContent != "" { + resp.ReasoningContent += ck.ReasoningContent + if a != nil { + a.publishEvent(events.EventReasoningDelta, map[string]interface{}{ + "content": ck.ReasoningContent, + "channel": a.currentOutputChannel, + }) + } + } + if ck.Content != "" { + resp.Content += ck.Content + if a != nil { + a.publishEvent(events.EventContentDelta, map[string]interface{}{ + "content": ck.Content, + "channel": a.currentOutputChannel, + }) + } + } + + // 增量 tool call 分片:OpenAI 风格按 index 拼接 id/name/arguments + for i, tc := range ck.ToolCalls { + idx := i + acc := accs[idx] + if acc == nil { + acc = &toolCallAcc{} + accs[idx] = acc + } + if tc.ID != "" { + acc.id = tc.ID + } + if tc.Name != "" { + acc.name = tc.Name + } + // arguments 以 JSON 字符串分片到达(OpenAI 标准),拼接后最终解析 + if tc.RawArguments != "" { + acc.argsRaw.WriteString(tc.RawArguments) + } + } + + if ck.Done && ck.FinishReason != "" { + lastFinish = ck.FinishReason + } + if ck.Usage != nil { + resp.TokenUsage = *ck.Usage + } + + case <-ctx.Done(): + for idx := range accs { + flushToolCall(idx) + } + return resp, ctx.Err() + } + } +} + +// parseToolArgsJSON 将经过完整拼接的 tool call arguments JSON 字符串解析为 map。 +// 空字符串返回空 map。 +func parseToolArgsJSON(s string) map[string]interface{} { + if s == "" { + return map[string]interface{}{} + } + var m map[string]interface{} + if err := json.Unmarshal([]byte(s), &m); err == nil && m != nil { + return m + } + return map[string]interface{}{} +} + func convertToolCalls(tcs []agentAPI.ToolCall) []sdk.ToolCall { if tcs == nil { return nil diff --git a/internal/events/bus.go b/internal/events/bus.go index 7e902fd..c400de2 100644 --- a/internal/events/bus.go +++ b/internal/events/bus.go @@ -17,7 +17,15 @@ const ( EventStage EventType = "stage" EventSystem EventType = "system" EventTerminalOutput EventType = "terminal_output" - EventAll EventType = "*" + + // 流式增量事件(LLM token 级):核心改为流式后每收到一个增量块发布。 + // 订阅者可选订;不认识的旧订阅者自然忽略(Bus 按 EventType 精确匹配分发)。 + // 聚合事件 EventReasoning / EventAgentLLMChain 仍照常在每轮结束时全文发布, + // 插件体系行为不变。 + EventReasoningDelta EventType = "reasoning_delta" + EventContentDelta EventType = "content_delta" + + EventAll EventType = "*" ) type Event struct {