Files
HomeAgent/internal/agent/core/toolbatch_test.go
JianFeeeee 28b42bcf0f refactor(toolcall): 批内消息改为「一个 assistant 带全部 tool_calls」(阶段 2a)
问题:现状每个工具各自 append 一对(assistant[tool_calls=[tc]] + tool),
既不表达「这是一批」,也无法支撑并行:
· 产生 N 条 assistant 消息,同一段 assistant 文本语义上只该出现一次
· 并行下完成顺序不确定,若等结果回来再落消息,assistant 就必须等所有
  结果齐了才能写——而 OpenAI 协议要求 assistant(tool_calls) 在结果**之前**

改动(task.go):
· TaskFrame 增 assistantMsgIdx
· 新增 ensureBatchAssistant:惰性写入,全批只写**一条** assistant,
  携带 f.PendingTools 全部 tool_calls;后续工具只补 tool 消息
· stepToolBegin 的 denied / unhealthy 分支与 stepToolAfter 统一改用它
· stepLLM 设 PendingTools 时清零 assistantMsgIdx
· msgContent 的 ContentOnce 归位移入 ensureBatchAssistant(仍是只挂第一条)

⚠️ 依赖:before_toolcall 阶段**不得**改写工具参数——已核实全仓无此用法
(grep ToolCalls[0].Arguments 赋值无结果)。若将来某插件要改写 args,
需改为「回填后重写该条 assistant」。该前提已写入代码注释。

判据(toolbatch_test.go 追加 TestBatchLayoutSingleAssistantCarriesAllToolCalls):
· 带 tool_calls 的 assistant **恰好一条**且携带 2 个 tool_calls
· 其后紧跟 2 条 tool 消息且按声明顺序(c1、c2)

变异验证:让 ensureBatchAssistant 退化为「每工具一条」⇒ 判据 FAIL
「批内 assistant 应带 2 个 tool_calls,实际 1」。

stage 0.5 补的三条判据(配对完整性 / ContentOnce / denied 后继续)
在本改动后**仍然全绿**——它们正是为这种改动准备的保护网。

回归:internal/agent/... internal/sdk/... internal/plugins/... 全绿(18 包)。
2026-09-27 11:09:02 +08:00

302 lines
11 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package core
import (
"context"
"sync"
"testing"
agentAPI "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api"
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
)
// 本文件钉死「同一批多个 tool_call」这条路径的现状行为。
//
// 为何必须先有它(已核实):仓内此前**没有任何测试直接驱动**
// PendingTools / ToolIdx —— 即批内循环(StepToolBegin → Exec → After →
// StepToolBegin…)**无判据可依**。而阶段 2(并行执行层)要改的正是这段。
// 没有判据就改,等于在无保护的核心路径上动手。
//
// 这些断言全部是**确定性**的:阶段 0 已消除 map 迭代随机性,
// 同一批工具的落序与消息配对可稳定断言。
// batchProvider 依次返回预设响应;ChatStream 不支持流式(驱动回退到 Chat)。
type batchProvider struct {
mu sync.Mutex
responses []*agentAPI.CompletionResponse
calls int
}
func (p *batchProvider) Name() string { return "batch" }
func (p *batchProvider) Chat(ctx context.Context, req *agentAPI.CompletionRequest) (*agentAPI.CompletionResponse, error) {
p.mu.Lock()
defer p.mu.Unlock()
p.calls++
if p.calls-1 < len(p.responses) {
return p.responses[p.calls-1], nil
}
return &agentAPI.CompletionResponse{Content: "done"}, nil
}
func (p *batchProvider) ChatStream(ctx context.Context, req *agentAPI.CompletionRequest) (<-chan agentAPI.StreamChunk, error) {
return nil, context.Canceled
}
func (p *batchProvider) MaxContextTokens() int { return 8192 }
// newBatchAgent 建一个带两枚「记录型」工具的 agent。
// 返回的 *[]string 按调用顺序记录 tool 名,便于断言批内顺序。
func newBatchAgent(t *testing.T, sp agentAPI.Provider) (*Agent, *[]string) {
t.Helper()
a := New(AgentConfig{
ID: "batchagent",
Provider: sp,
ProviderManager: agentAPI.NewProviderManager(),
IO: agentIO.NewIOManager(),
StageHost: NewStageHost(),
})
var mu sync.Mutex
var called []string
// 注意:阶段 2 起组内会**并发**执行,届时 toolFn 可能被多 goroutine 同时调用,
// 故 mu 必须一直持有(不能为“只在串行时用”而省略)。
dev := &mockOutputDevice{
name: "batchdev",
caps: agentIO.CapText,
tools: []agentIO.ToolDef{
{Name: "tool_alpha", Description: "a"},
{Name: "tool_beta", Description: "b"},
},
toolFn: func(tool string, args map[string]interface{}) (interface{}, error) {
mu.Lock()
called = append(called, tool)
mu.Unlock()
return "ran:" + tool, nil
},
}
if err := a.io.RegisterDevice(dev); err != nil {
t.Fatalf("注册测试设备失败: %v", err)
}
return a, &called
}
// 回归:同一批多个 tool_call 必须**全部**执行,且都进入 ToolResults。
//
// 现状:StepToolBegin 用 f.ToolIdx 遍历 f.PendingTools,逐一执行。
// 批内若有工具被漏掉(如索引推进错误),本判据立即失败。
func TestBatchExecutesEveryToolCall(t *testing.T) {
tcA := agentAPI.ToolCall{ID: "c1", Name: "tool_alpha", Arguments: map[string]interface{}{}}
tcB := agentAPI.ToolCall{ID: "c2", Name: "tool_beta", Arguments: map[string]interface{}{}}
sp := &batchProvider{responses: []*agentAPI.CompletionResponse{
{Content: "batch text", ToolCalls: []agentAPI.ToolCall{tcA, tcB}},
{Content: "final"},
}}
a, called := newBatchAgent(t, sp)
resp, toolsUsed, results, err := a.process("go", a.stageCtxFromInput("go", "", ""))
if err != nil {
t.Fatalf("process: %v", err)
}
if resp != "final" {
t.Errorf("最终响应 = %q,期望 %q", resp, "final")
}
if len(toolsUsed) != 2 || toolsUsed[0] != "tool_alpha" || toolsUsed[1] != "tool_beta" {
t.Errorf("ToolsUsed = %v,期望 [tool_alpha tool_beta]", toolsUsed)
}
if len(*called) != 2 {
t.Fatalf("实际执行 %v,期望两个都执行", *called)
}
// 批内顺序必须等于模型给出的顺序(阶段 0 修复的落序问题在批内的体现)。
if (*called)[0] != "tool_alpha" || (*called)[1] != "tool_beta" {
t.Errorf("批内执行顺序 = %v,期望 [tool_alpha tool_beta]", *called)
}
if len(results) != 2 {
t.Fatalf("ToolResults 有 %d 条,期望 2", len(results))
}
for _, r := range results {
if r.Output == "" {
t.Errorf("工具 %s 的结果为空", r.Name)
}
}
}
// 回归:批内每个 tool_call_id 都必须有且仅有一条 role=tool 消息配对。
//
// 这是并行化(阶段 2 把落法改成「一个 assistant 带全部 tool_calls + N 条 tool」)
// 的**安全网**:配对一旦断裂,上游会因 tool_call_id 找不到结果而报错。
// 本判据只看配对完整性,不断言消息的物理排列(那正是 2a 要改的部分)。
func TestBatchToolCallIDsAllPaired(t *testing.T) {
tcA := agentAPI.ToolCall{ID: "c1", Name: "tool_alpha", Arguments: map[string]interface{}{}}
tcB := agentAPI.ToolCall{ID: "c2", Name: "tool_beta", Arguments: map[string]interface{}{}}
sp := &batchProvider{responses: []*agentAPI.CompletionResponse{
{ToolCalls: []agentAPI.ToolCall{tcA, tcB}},
{Content: "final"},
}}
a, _ := newBatchAgent(t, sp)
// 直接驱动状态机并保留帧,以便检查 msgs。
f := a.newTaskFrame("go", a.stageCtxFromInput("go", "", ""))
if out := a.runTaskSteps(f); out != outcomeDone {
t.Fatalf("runTaskSteps = %v, err=%v", out, f.Err)
}
// 收集 assistant 声明的 tool_call_id 与 tool 消息回填的 id。
declared := map[string]int{}
answered := map[string]int{}
for _, m := range f.Msgs {
for _, tc := range m.ToolCalls {
declared[tc.ID]++
}
if m.Role == "tool" {
answered[m.ToolCallID]++
}
}
for _, id := range []string{"c1", "c2"} {
if declared[id] != 1 {
t.Errorf("tool_call_id %q 被声明 %d 次,期望 1 次", id, declared[id])
}
if answered[id] != 1 {
t.Errorf("tool_call_id %q 被回填 %d 次,期望 1 次", id, answered[id])
}
}
// 不得有悬空的 tool 消息。
for id, n := range answered {
if declared[id] == 0 {
t.Errorf("存在无对应 tool_call 声明的 tool 消息: id=%q n=%d", id, n)
}
}
}
// 回归:批内 assistant 的文本只应出现**一次**(ContentOnce 语义)。
//
// 现状:f.ContentOnce 保证 f.Resp.Content 只挂在第一个工具的 assistant 消息上,
// 避免同一段文本在批内被重复 N 次、撑爆上下文。
func TestBatchAssistantTextAppearsOnce(t *testing.T) {
tcA := agentAPI.ToolCall{ID: "c1", Name: "tool_alpha", Arguments: map[string]interface{}{}}
tcB := agentAPI.ToolCall{ID: "c2", Name: "tool_beta", Arguments: map[string]interface{}{}}
sp := &batchProvider{responses: []*agentAPI.CompletionResponse{
{Content: "UNIQUE_BATCH_TEXT", ToolCalls: []agentAPI.ToolCall{tcA, tcB}},
{Content: "final"},
}}
a, _ := newBatchAgent(t, sp)
f := a.newTaskFrame("go", a.stageCtxFromInput("go", "", ""))
if out := a.runTaskSteps(f); out != outcomeDone {
t.Fatalf("runTaskSteps = %v, err=%v", out, f.Err)
}
n := 0
for _, m := range f.Msgs {
if m.Role == "assistant" && m.Content == "UNIQUE_BATCH_TEXT" {
n++
}
}
if n != 1 {
t.Errorf("批内 assistant 文本出现 %d 次,期望恰好 1 次(ContentOnce 语义)", n)
}
}
// 回归:批内某个工具**被拒绝**时,批内其余工具仍应继续执行。
//
// 现状:stepToolBegin 的 denied 分支只 f.ToolIdx++ 并 continue,不中断整批。
// 若改成 abort,模型会丢掉本可执行的后续调用。
func TestBatchContinuesAfterDeniedTool(t *testing.T) {
reason := "策略拒绝"
sp := &batchProvider{responses: []*agentAPI.CompletionResponse{
{ToolCalls: []agentAPI.ToolCall{
{ID: "c1", Name: "tool_alpha", Arguments: map[string]interface{}{}},
{ID: "c2", Name: "tool_beta", Arguments: map[string]interface{}{}},
}},
{Content: "final"},
}}
a, _ := newBatchAgent(t, sp)
// before_toolcall 只拒绝 tool_alpha;tool_beta 放行。
a.stageHost.RegisterStage(sdk.StageBeforeToolcall, func(ctx *sdk.StageContext) error {
if len(ctx.ToolCalls) > 0 && ctx.ToolCalls[0].Name == "tool_alpha" {
r := reason
ctx.Response = &r
}
return nil
})
f := a.newTaskFrame("go", a.stageCtxFromInput("go", "", ""))
if out := a.runTaskSteps(f); out != outcomeDone {
t.Fatalf("runTaskSteps = %v, err=%v", out, f.Err)
}
// 批内两个工具都要被「尝试」(拒绝也计入 ToolsUsed)。
if len(f.ToolsUsed) != 2 {
t.Errorf("被拒后整批应继续,ToolsUsed = %v 期望两个", f.ToolsUsed)
}
// 拒绝理由必须回填给模型,且不阻断 tool_beta 的结果。
var sawReason, sawBeta bool
for _, m := range f.Msgs {
if m.Role == "tool" {
if m.ToolCallID == "c1" && m.Content == reason {
sawReason = true
}
if m.ToolCallID == "c2" && m.Content != "" {
sawBeta = true
}
}
}
if !sawReason {
t.Errorf("拒绝理由未回填给模型")
}
if !sawBeta {
t.Errorf("tool_beta 的结果未回填(被前一工具的拒绝连带丢弃)")
}
}
// 阶段 2a:消息落法改为「**一个** assistant 带全部 tool_calls + N 条 tool」。
//
// 现状:每个工具各自 append 一对(assistant[tool_calls=[tc]] + tool),
// 不表达「这是一批」。阶段 2 的批内并发要求消息形态与之对应,且并行下
// 多个 tool message 的相对顺序必须**按 index 确定**,否则模型读到的
// 上下文顺序 ≠ 执行顺序,会诱导出错误的因果推断。
//
// 本判据钉死:新布局下(a)配对仍完整、(b)assistant 只出现一条且带全部
// tool_calls、(c)tool 消息按 index 升序、(d)多模态 user 消息仍紧跟
// 各自的 tool 消息。
func TestBatchLayoutSingleAssistantCarriesAllToolCalls(t *testing.T) {
tcA := agentAPI.ToolCall{ID: "c1", Name: "tool_alpha", Arguments: map[string]interface{}{}}
tcB := agentAPI.ToolCall{ID: "c2", Name: "tool_beta", Arguments: map[string]interface{}{}}
sp := &batchProvider{responses: []*agentAPI.CompletionResponse{
{Content: "BATCHTEXT", ToolCalls: []agentAPI.ToolCall{tcA, tcB}},
{Content: "final"},
}}
a, _ := newBatchAgent(t, sp)
f := a.newTaskFrame("go", a.stageCtxFromInput("go", "", ""))
if out := a.runTaskSteps(f); out != outcomeDone {
t.Fatalf("runTaskSteps=%v err=%v", out, f.Err)
}
// ① 找带 tool_calls 的 assistant 消息,必须**恰好一条**且带 2 个。
var assistants []int
for i, m := range f.Msgs {
if m.Role == "assistant" && len(m.ToolCalls) > 0 {
assistants = append(assistants, i)
if len(m.ToolCalls) != 2 {
t.Errorf("批内 assistant 应带 2 个 tool_calls,实际 %d", len(m.ToolCalls))
}
}
}
if len(assistants) != 1 {
t.Fatalf("带 tool_calls 的 assistant 应恰好 1 条,实际 %d 条(索引 %v)", len(assistants), assistants)
}
// ② 该 assistant 之后应紧跟 2 条 tool 消息,且按声明顺序。
idx := assistants[0]
var gotIDs []string
for i := idx + 1; i < len(f.Msgs); i++ {
if f.Msgs[i].Role == "tool" {
gotIDs = append(gotIDs, f.Msgs[i].ToolCallID)
}
}
if len(gotIDs) != 2 {
t.Fatalf("assistant 之后应有 2 条 tool 消息,实际 %d(%v)", len(gotIDs), gotIDs)
}
if gotIDs[0] != "c1" || gotIDs[1] != "c2" {
t.Errorf("tool 消息应按 index 升序,实际 %v", gotIDs)
}
}