mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 07:43:58 +00:00
feat(agent): token-level streaming in core process loop
Replace the blocking Chat() call in process() with
chatStreamWithFallback: ChatStream first, accumulate chunks, fall back
to non-stream Chat on connect failure or empty-stream failure.
Why: the non-streaming path blocked for the ENTIRE LLM generation (up
to the 180s HTTP timeout). Reasoning models thinking 60-120s plus AUTO
chain failover regularly exceeded it -> context canceled -> full turn
wasted. With streaming the first chunk arrives in ~1-3s and any
flowing token keeps the connection alive; total generation time is no
longer bounded by an overall timeout.
Compatibility (external behavior unchanged):
- process() signature/return values unchanged
- Aggregated events (EventReasoning / EventAgentLLMChain) still fire
once per turn with full text after stream completion - existing
plugin subscribers see identical payloads as before
- New incremental events EventReasoningDelta / EventContentDelta are
additive; old subscribers ignore unknown event types
- Tool execution loop, memory pipeline, stage pipeline untouched
Streaming details:
- Tool call fragments accumulated per OpenAI streaming convention:
id/name arrive on the first fragment, arguments as raw JSON string
shards across fragments; merged and parsed once at stream end
- normalizeStreamToolCalls keeps nameless argument shards (the
non-stream normalizer drops them); ToolCall gains RawArguments to
carry shard text
- Interrupt mid-stream returns partial content instead of discarding
the whole generation
Verified end-to-end against llmsproxy: plain chat streams correctly;
curl confirms tool-call shard wire format ({" + command" + :"date"}
-> {"command":"date"}); unit tests cover shard merging and
content/reasoning accumulation.
This commit is contained in:
@ -190,6 +190,7 @@ type ToolCall struct {
|
|||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Arguments map[string]interface{} `json:"arguments"`
|
Arguments map[string]interface{} `json:"arguments"`
|
||||||
|
RawArguments string `json:"raw_arguments,omitempty"` // 流式分片原始 JSON 字符串
|
||||||
}
|
}
|
||||||
|
|
||||||
type apiToolCall struct {
|
type apiToolCall struct {
|
||||||
@ -494,6 +495,48 @@ func normalizeOpenAIToolCalls(raw []openAIToolCall) []ToolCall {
|
|||||||
Type: typ,
|
Type: typ,
|
||||||
Name: name,
|
Name: name,
|
||||||
Arguments: parseToolArguments(argsRaw),
|
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
|
return out
|
||||||
@ -603,7 +646,7 @@ func parseOpenAICompatibleStreamChunkFull(data string) (StreamChunk, bool) {
|
|||||||
ck := StreamChunk{
|
ck := StreamChunk{
|
||||||
Content: stringifyContent(choice.Delta.Content),
|
Content: stringifyContent(choice.Delta.Content),
|
||||||
ReasoningContent: choice.Delta.ReasoningContent,
|
ReasoningContent: choice.Delta.ReasoningContent,
|
||||||
ToolCalls: normalizeOpenAIToolCalls(choice.Delta.ToolCalls),
|
ToolCalls: normalizeStreamToolCalls(choice.Delta.ToolCalls),
|
||||||
Usage: usage,
|
Usage: usage,
|
||||||
}
|
}
|
||||||
// finish reason 为空字符串不算终止信号(sensenova 每块都发 "")
|
// finish reason 为空字符串不算终止信号(sensenova 每块都发 "")
|
||||||
|
|||||||
@ -2,6 +2,7 @@ package core
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
@ -120,7 +121,7 @@ func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response stri
|
|||||||
a.cancelLLM = fCancel
|
a.cancelLLM = fCancel
|
||||||
a.llmMu.Unlock()
|
a.llmMu.Unlock()
|
||||||
|
|
||||||
resp, llmErr = fbProvider.Chat(fCtx, req)
|
resp, llmErr = chatStreamWithFallback(fCtx, fbProvider, req, a)
|
||||||
|
|
||||||
a.llmMu.Lock()
|
a.llmMu.Lock()
|
||||||
a.cancelLLM = nil
|
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 {
|
func convertToolCalls(tcs []agentAPI.ToolCall) []sdk.ToolCall {
|
||||||
if tcs == nil {
|
if tcs == nil {
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@ -17,6 +17,14 @@ const (
|
|||||||
EventStage EventType = "stage"
|
EventStage EventType = "stage"
|
||||||
EventSystem EventType = "system"
|
EventSystem EventType = "system"
|
||||||
EventTerminalOutput EventType = "terminal_output"
|
EventTerminalOutput EventType = "terminal_output"
|
||||||
|
|
||||||
|
// 流式增量事件(LLM token 级):核心改为流式后每收到一个增量块发布。
|
||||||
|
// 订阅者可选订;不认识的旧订阅者自然忽略(Bus 按 EventType 精确匹配分发)。
|
||||||
|
// 聚合事件 EventReasoning / EventAgentLLMChain 仍照常在每轮结束时全文发布,
|
||||||
|
// 插件体系行为不变。
|
||||||
|
EventReasoningDelta EventType = "reasoning_delta"
|
||||||
|
EventContentDelta EventType = "content_delta"
|
||||||
|
|
||||||
EventAll EventType = "*"
|
EventAll EventType = "*"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user