mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-28 13:23:03 +00:00
refactor(toolcall): StageContext 拆 per-tool,为并发执行消除共享槽(阶段 2c)
问题:f.StageCtx 是**单槽**,批内每个工具都覆写它(ToolCalls=[单元素]、 ToolResults 覆写、Results[0] 回读)。串行下看不出问题,但并发下 N 个 goroutine 同写一个 ctx = 数据竞争,且 after_toolcall 插件可能读到 **别的工具**的结果。 改动(task.go): · TaskFrame 增 toolCtxs []sdk.StageContext(每工具一份) · buildToolContexts 在 stepLLM 设 PendingTools 时建池; Extra **逐份浅拷贝**——共享同一 map 即竞争(stage handler 会写它) · toolCtxFor(i) 取第 i 份,越界/未建时回落 f.StageCtx(测试替身安全) · stepToolBegin(before_toolcall + 参数回填)、stepToolExec(写结果)、 stepToolAfter(读结果)三处全部切到 per-tool ctx 判据(toolbatch_test.go 追加两条): · 每个工具的 before_toolcall ctx 只带自己的 ToolCalls[0].Name, 且 Extra[output_channel] 逐份带过去(stage.go:18 依赖它) · after_toolcall 读到的 Result 必须属于当前工具,不能是批内另一个的 ★ 诚实记录:这两条判据在**串行**下**测不出与单槽的差别**——串行时 每工具跑完才进下一个,不存在交错。变体验证(toolCtxFor 退回单槽)后 判据仍全绿。故 2c 记为「实现已就位、判据未闭合」,真正判据必须与 2d (并发执行)一起写,并以 -race 确认无竞争。已在执行计划中标注。 过程中两次自伤: · 我的 harness 没设 Extra[output_channel](那是 prepareInputTask 才写的, task.go:380),判据一度报「产品缺陷」——核实后是我造的场景,已对齐生产; · 阶段 1 的 TestStageCtxSuccessIsHonestEndToEnd 读 f.StageCtx.ToolResults, 拆分后失效——判据跟随新结构改为按批索引取 toolCtxs[i], 断言的仍是内核产出的 Success 值本身。 顺带记录(非本次引入):TestResidualKeep/Drop 偶发失败,根因是 offload_test.go 的 SpawnResident 起了子调度器 goroutine,而测试 enqueue 后无同步就读同一队列。干净基线 3/3 全绿属运气。已在计划中 记为待修,避免后续误判为并行化引入的回归。
This commit is contained in:
@ -110,6 +110,16 @@ type TaskFrame struct {
|
||||
CurTool agentAPI.ToolCall
|
||||
CurToolPlugin string
|
||||
CurResult string
|
||||
// toolCtxs 是**每个工具各一份**的 StageContext(阶段 2c)。
|
||||
//
|
||||
// 为何必须拆:f.StageCtx 原本是单槽,批内每个工具都覆写
|
||||
// (ToolCalls=[单元素]、ToolResults 覆写、Results[0] 回读)。
|
||||
// 并行下 N 个 goroutine 同写一个 ctx = 数据竞争,且 after_toolcall
|
||||
// 插件可能读到**别的工具**的结果。
|
||||
// 拆分后每个工具只写自己那份,Extra 在构造时逐份复制
|
||||
//(output_channel / input_source / media_* —— stage.go:18 依赖前者)。
|
||||
toolCtxs []sdk.StageContext
|
||||
|
||||
// assistantMsgIdx 是本批 assistant(tool_calls) 消息在 Msgs 中的下标,
|
||||
// -1 表示尚未写入。阶段 2a:批内只写**一条** assistant 承载全部
|
||||
// tool_calls,工具结果各自作为 tool 消息追加在它之后。
|
||||
@ -682,6 +692,8 @@ func (a *Agent) stepLLM(f *TaskFrame) stepOutcome {
|
||||
}
|
||||
f.ContentOnce = true
|
||||
f.PendingTools = resp.ToolCalls
|
||||
// 阶段 2c:为批内每个工具预建独立的 StageContext(Extra 逐份复制)。
|
||||
f.buildToolContexts()
|
||||
// 阶段 2a:批内消息**预置**为「一个 assistant 带全部 tool_calls」。
|
||||
//
|
||||
// 为何预置而不是逐步 append:并行执行下多个工具的**完成顺序不确定**,
|
||||
@ -713,10 +725,12 @@ func (a *Agent) stepToolBegin(f *TaskFrame) stepOutcome {
|
||||
}
|
||||
|
||||
sdkTC := sdk.ToolCall{ID: tc.ID, Name: tc.Name, Plugin: pluginName, Arguments: tc.Arguments}
|
||||
f.StageCtx.ToolCalls = []sdk.ToolCall{sdkTC}
|
||||
f.StageCtx.ToolResults = nil
|
||||
if a.runStage(sdk.StageBeforeToolcall, f.StageCtx) {
|
||||
result := denialResultText(f.StageCtx, tc.Name)
|
||||
// 阶段 2c:写**本工具自己的** ctx,不再覆写共享的 f.StageCtx。
|
||||
tctx := f.toolCtxFor(f.ToolIdx)
|
||||
tctx.ToolCalls = []sdk.ToolCall{sdkTC}
|
||||
tctx.ToolResults = nil
|
||||
if a.runStage(sdk.StageBeforeToolcall, tctx) {
|
||||
result := denialResultText(tctx, tc.Name)
|
||||
f.ensureBatchAssistant()
|
||||
f.Msgs = append(f.Msgs, agentAPI.Message{Role: "tool", ToolCallID: tc.ID, Content: result})
|
||||
a.publishEvent(events.EventToolCall, map[string]interface{}{
|
||||
@ -730,7 +744,7 @@ func (a *Agent) stepToolBegin(f *TaskFrame) stepOutcome {
|
||||
f.ToolIdx++
|
||||
return outcomeContinue
|
||||
}
|
||||
tc.Arguments = f.StageCtx.ToolCalls[0].Arguments
|
||||
tc.Arguments = tctx.ToolCalls[0].Arguments
|
||||
|
||||
if pluginName != "" && !a.pluginHealth.isHealthy(pluginName) {
|
||||
result := fmt.Sprintf("插件 %s 处于崩溃状态,已跳过执行,等待自动恢复重载", pluginName)
|
||||
@ -777,7 +791,7 @@ func (a *Agent) stepToolExec(f *TaskFrame) stepOutcome {
|
||||
// false。工具失败是以 nil error + 错误**值**返回的,所以判据必须看返回值。
|
||||
// ⚠️ isToolError 必须同时覆盖存量插件的两种失败约定与「成功不误判」,
|
||||
// 否则升级会把存量插件的成功判成失败(见 toolerror_test.go)。
|
||||
f.StageCtx.ToolResults = []sdk.ToolResult{{
|
||||
f.toolCtxFor(f.ToolIdx).ToolResults = []sdk.ToolResult{{
|
||||
CallID: f.CurTool.ID, Name: f.CurTool.Name, Plugin: f.CurToolPlugin,
|
||||
Success: !isToolError(outcome.Raw), Result: outcome.Raw,
|
||||
}}
|
||||
@ -791,12 +805,13 @@ func (a *Agent) stepToolAfter(f *TaskFrame) stepOutcome {
|
||||
pluginName := f.CurToolPlugin
|
||||
result := f.CurResult
|
||||
|
||||
a.runStage(sdk.StageAfterToolcall, f.StageCtx)
|
||||
if len(f.StageCtx.ToolResults) > 0 {
|
||||
tctx := f.toolCtxFor(f.ToolIdx)
|
||||
a.runStage(sdk.StageAfterToolcall, tctx)
|
||||
if len(tctx.ToolResults) > 0 {
|
||||
// 此前是 `Result.(string)` 类型断言,而插件返回的多是 map ⇒ 断言几乎
|
||||
// 恒失败,after_toolcall 阶段对结构化结果的改写**静默失效**。
|
||||
// 改为:字符串就替换文本;结构化值则保留其原值并按契约渲染。
|
||||
switch r := f.StageCtx.ToolResults[0].Result.(type) {
|
||||
switch r := tctx.ToolResults[0].Result.(type) {
|
||||
case string:
|
||||
result = r
|
||||
case nil:
|
||||
@ -1050,3 +1065,42 @@ func (a *Agent) callLLMWithFallback(req *agentAPI.CompletionRequest, providers [
|
||||
|
||||
return resp, llmErr
|
||||
}
|
||||
|
||||
// buildToolContexts 为本批每个工具预建一份独立的 StageContext。
|
||||
//
|
||||
// Extra 必须**逐份复制**而不是共享同一个 map:Extra 会被 stage handler 写
|
||||
// (例如插件注入 media_blocks),共享即竞争。逐份浅拷贝即可——里面的值
|
||||
// (string / []agentAPI.ContentBlock)本身是只读的。
|
||||
func (f *TaskFrame) buildToolContexts() {
|
||||
base := f.StageCtx
|
||||
f.toolCtxs = make([]sdk.StageContext, len(f.PendingTools))
|
||||
for i := range f.PendingTools {
|
||||
f.toolCtxs[i] = sdk.StageContext{
|
||||
Extra: copyExtraMap(base.Extra),
|
||||
NoMemory: base.NoMemory,
|
||||
}
|
||||
// Phase 由 runStage 每次调用时设置,此处不预置。
|
||||
}
|
||||
}
|
||||
|
||||
// toolCtxFor 返回第 i 个工具的 StageContext;越界或未建时回落到 f.StageCtx,
|
||||
// 保证调用方不必判空(测试替身等未走 buildToolContexts 的路径)。
|
||||
func (f *TaskFrame) toolCtxFor(i int) *sdk.StageContext {
|
||||
if i >= 0 && i < len(f.toolCtxs) {
|
||||
return &f.toolCtxs[i]
|
||||
}
|
||||
return f.StageCtx
|
||||
}
|
||||
|
||||
// copyExtraMap 浅拷贝 Extra(值视为只读)。
|
||||
// nil 安全:base.Extra 为 nil 时返回 nil,写入方需自行判空。
|
||||
func copyExtraMap(src map[string]interface{}) map[string]interface{} {
|
||||
if src == nil {
|
||||
return nil
|
||||
}
|
||||
dst := make(map[string]interface{}, len(src))
|
||||
for k, v := range src {
|
||||
dst[k] = v
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
@ -2,6 +2,8 @@ package core
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
@ -299,3 +301,104 @@ func TestBatchLayoutSingleAssistantCarriesAllToolCalls(t *testing.T) {
|
||||
t.Errorf("tool 消息应按 index 升序,实际 %v", gotIDs)
|
||||
}
|
||||
}
|
||||
|
||||
// 阶段 2c:`StageContext` 拆 per-tool。
|
||||
//
|
||||
// 现状:f.StageCtx 是**单槽**,批内每个工具都覆写它
|
||||
// (ToolCalls=[单元素]、ToolResults 覆写、Results[0] 回读)。并行下
|
||||
// N 个 goroutine 同写一个 ctx = 数据竞争,且 after_toolcall 插件读到的
|
||||
// 可能是**别的工具**的结果。
|
||||
//
|
||||
// 本判据钉死:每个工具的 before/after stage 必须各自看到**自己的**
|
||||
// ToolCalls[0].Name 与自己的结果,输出通道等 Extra 也要逐份带过去。
|
||||
func TestBatchEachToolSeesItsOwnStageContext(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)
|
||||
|
||||
var mu sync.Mutex
|
||||
beforeSeen := map[string]string{}
|
||||
var missingChannel int
|
||||
|
||||
a.stageHost.RegisterStage(sdk.StageBeforeToolcall, func(ctx *sdk.StageContext) error {
|
||||
name := ""
|
||||
if len(ctx.ToolCalls) > 0 {
|
||||
name = ctx.ToolCalls[0].Name
|
||||
}
|
||||
// 每个工具的 ctx 必须只带它自己(长度恒为 1),否则就是单槽串味。
|
||||
if len(ctx.ToolCalls) != 1 {
|
||||
t.Errorf("before_toolcall 的 ctx 应只带 1 个 ToolCall,实际 %d", len(ctx.ToolCalls))
|
||||
}
|
||||
// Extra 里的 output_channel 必须逐份复制过来(stage.go:18 依赖它)。
|
||||
if _, ok := ctx.Extra["output_channel"]; !ok {
|
||||
mu.Lock()
|
||||
missingChannel++
|
||||
mu.Unlock()
|
||||
}
|
||||
mu.Lock()
|
||||
beforeSeen[name] = name
|
||||
mu.Unlock()
|
||||
return nil
|
||||
})
|
||||
|
||||
// 复刻生产:prepareInputTask 会往 Extra 写 input_source/output_channel
|
||||
//(task.go:380-381)。我的 harness 若不设它,测的就是「Extra 缺失」
|
||||
// 这一**自己造的**场景,而不是「per-tool 复制」——先修正 harness。
|
||||
ctx := a.stageCtxFromInput("go", "cli", "")
|
||||
ctx.Extra["output_channel"] = "cli"
|
||||
ctx.Extra["input_source"] = "cli"
|
||||
if out := a.runTaskSteps(a.newTaskFrame("go", ctx)); out != outcomeDone {
|
||||
t.Fatalf("runTaskSteps 未收敛: %v", out)
|
||||
}
|
||||
if len(beforeSeen) != 2 {
|
||||
t.Errorf("before_toolcall 应被两个工具各触发一次且名字不同,实际 %v", beforeSeen)
|
||||
}
|
||||
for _, want := range []string{"tool_alpha", "tool_beta"} {
|
||||
if beforeSeen[want] != want {
|
||||
t.Errorf("工具 %s 的 before_toolcall 未看到自己(看到 %q)", want, beforeSeen[want])
|
||||
}
|
||||
}
|
||||
if missingChannel > 0 {
|
||||
t.Errorf("有 %d 个工具的 ctx 缺少 Extra[output_channel]", missingChannel)
|
||||
}
|
||||
}
|
||||
|
||||
// 反向断言:after_toolcall 读到的结果必须属于**当前**工具,
|
||||
// 不能是批内另一个工具的(单槽下极易串味)。
|
||||
func TestBatchAfterToolcallSeesOwnResult(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)
|
||||
|
||||
var mu sync.Mutex
|
||||
bad := map[string]string{}
|
||||
a.stageHost.RegisterStage(sdk.StageAfterToolcall, func(ctx *sdk.StageContext) error {
|
||||
if len(ctx.ToolResults) == 0 {
|
||||
return nil
|
||||
}
|
||||
name := ctx.ToolResults[0].Name
|
||||
res := fmt.Sprint(ctx.ToolResults[0].Result)
|
||||
// 结果文案必须含自己的工具名("ran:tool_alpha"),否则就是串味。
|
||||
if !strings.Contains(res, name) {
|
||||
mu.Lock()
|
||||
bad[name] = res
|
||||
mu.Unlock()
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
if out := a.runTaskSteps(a.newTaskFrame("go", a.stageCtxFromInput("go", "", ""))); out != outcomeDone {
|
||||
t.Fatalf("runTaskSteps 未收敛: %v", out)
|
||||
}
|
||||
if len(bad) > 0 {
|
||||
t.Errorf("after_toolcall 读到别的工具的结果: %v", bad)
|
||||
}
|
||||
}
|
||||
|
||||
@ -6,6 +6,7 @@ import (
|
||||
|
||||
agentAPI "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api"
|
||||
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
|
||||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||||
)
|
||||
|
||||
// 阶段 1b:诚实化 Success。
|
||||
@ -121,10 +122,24 @@ func TestStageCtxSuccessIsHonestEndToEnd(t *testing.T) {
|
||||
if out := a.runTaskSteps(f); out != outcomeDone {
|
||||
t.Fatalf("runTaskSteps=%v err=%v", out, f.Err)
|
||||
}
|
||||
if len(f.StageCtx.ToolResults) == 0 {
|
||||
t.Fatal("StageCtx.ToolResults 为空")
|
||||
// 阶段 2c 起结果写在**每个工具自己的** ctx 上(不再回写 f.StageCtx),
|
||||
// 因此这里按批索引取对应那份——判据跟着结构走,但断言的仍是
|
||||
// **内核产出的 Success 值本身**。
|
||||
if len(f.toolCtxs) == 0 {
|
||||
t.Fatal("toolCtxs 为空(per-tool ctx 未建立)")
|
||||
}
|
||||
var tr sdk.ToolResult
|
||||
found := false
|
||||
for i := range f.toolCtxs {
|
||||
if len(f.toolCtxs[i].ToolResults) > 0 {
|
||||
tr = f.toolCtxs[i].ToolResults[0]
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("各工具 ctx 的 ToolResults 全为空")
|
||||
}
|
||||
tr := f.StageCtx.ToolResults[0]
|
||||
if tr.Success != c.wantSucc {
|
||||
t.Errorf("Success = %v,期望 %v(返回值 %#v)", tr.Success, c.wantSucc, c.toolRet)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user