diff --git a/docs/zh/toolcall-contract-and-sequence-design.md b/docs/zh/toolcall-contract-and-sequence-design.md index 9afa992..e7b53e3 100644 --- a/docs/zh/toolcall-contract-and-sequence-design.md +++ b/docs/zh/toolcall-contract-and-sequence-design.md @@ -302,7 +302,7 @@ ParallelSafe bool `json:"parallel_safe,omitempty"` "in": { "host": "string", "verbose": "bool" }, "out": { "summary": "string", "load": "string" }, "when": "$args.verbose == true", - "tools": "{tool:cmd_run,args:{command:\"ssh $args.host uptime\"},as:summary} ; {tool:cmd_run,args:{command:\"ssh $args.host top -bn1\"},as:load} ;" + "tools": "{\"tool\":\"cmd_run\",\"args\":{\"command\":\"ssh $args.host uptime\"},\"as\":\"summary\"} ; {\"tool\":\"cmd_run\",\"args\":{\"command\":\"ssh $args.host top -bn1\"},\"as\":\"load\"} ;" }, { "name": "巡检全部", @@ -310,7 +310,7 @@ ParallelSafe bool `json:"parallel_safe,omitempty"` "in": { "hosts": "array" }, "out": { "reports": "array", "errors": "array" }, "parallel": true, - "tools": "{tool:seq_call,args:{group:\"拉取单台\",args:{host:$args.hosts[0]}},as:reports} ; {tool:seq_call,args:{group:\"拉取单台\",args:{host:$args.hosts[1]}},as:reports} ;" + "tools": "{\"tool\":\"seq_call\",\"args\":{\"group\":\"拉取单台\",\"args\":{\"host\":$args.hosts[0]}},\"as\":\"reports\"} ; {\"tool\":\"seq_call\",\"args\":{\"group\":\"拉取单台\",\"args\":{\"host\":$args.hosts[1]}},\"as\":\"reports\"} ;" } ] } @@ -377,6 +377,10 @@ ParallelSafe bool `json:"parallel_safe,omitempty"` 而不是静默丢一个工具。但 `;` 本身**只有在 tool 被 `{…}` 包裹时才是安全的**—— 这一点是本设计的关键,不加包裹会立刻出问题(见下)。 +⚠️ **每个 tool 必须是合法 JSON**(键要带引号):`{"tool":"cmd_run","args":{...},"as":"x"} ;` +写成 `{tool:cmd_run,...}` 是**非法 JSON**,内核用 `encoding/json` 解析会直接失败。 +(这条是 P1 实现时判据跑出来的真实缺陷,不是假想。) + #### 切分规则(已实测) `;` **仅在 brace/bracket 深度为 0 且不在字符串内**时才是分隔符: diff --git a/internal/plugins/seq/exec.go b/internal/plugins/seq/exec.go new file mode 100644 index 0000000..2a14fed --- /dev/null +++ b/internal/plugins/seq/exec.go @@ -0,0 +1,424 @@ +package seq + +import ( + "encoding/json" + "fmt" + "strconv" + "strings" + "sync" + "time" +) + +// toolRunner 是执行引擎对「执行一个工具」的依赖。 +// +// ⚠️ 它**必须**接收 args:插值是本阶段的核心能力之一,若接口不给出实际 +// 收到的参数,插值就无法被任何判据观察(我第一版写成 call(name) 时, +// 插值判据就成了摆设——它永远"通过")。 +type toolRunner interface { + call(name string, args map[string]interface{}) (string, error) +} + +// GroupResult 是一组的执行结果。 +type GroupResult struct { + Group string + Skipped bool // 条件为假而整组跳过 + Slots map[string]interface{} + Tools []ToolRun + Err error +} + +// ToolRun 是组内单个工具的执行记录。 +type ToolRun struct { + Name string + Result string + Err error + Order int // 在 tools 数组中的位置(合并按它,不按完成顺序) +} + +// execGroup 执行一组工具。 +// +// 三个不变量(各自有判据钉住): +// 1. **组内并行、组间串行**:parallel=true 时各工具并发跑。 +// 2. **合并按声明顺序**:结果写入各自槽位的"暂存区", +// 组屏障处按 tools 数组顺序**一次性**合并。 +// 并行下完成顺序不确定;若按完成顺序合并,同样的输入会产出不同的 +// 序列输出 —— 整条流程不可复现。 +// 顺序合并同时解决了并发写 map 的问题:**执行期不写共享 map**。 +// 3. **条件求值失败必须报错**,不得降级成"条件为假"。 +func execGroup(g Group, args map[string]interface{}, runner toolRunner) (GroupResult, error) { + res := GroupResult{Group: g.Name, Slots: map[string]interface{}{}} + + // ① 条件求值(在任何执行之前) + ok, err := evalCond(g.When, args) + if err != nil { + return res, fmt.Errorf("group %q 的条件求值失败: %w", g.Name, err) + } + if !ok { + res.Skipped = true + return res, nil // 槽**不赋值**:后续组引用它时由静态校验/调用方发现 + } + + n := len(g.Tools) + results := make([]ToolRun, n) + + run := func(i int) { + tc := g.Tools[i] + // 插值:把 $args.* 换成实参。整值引用保留原始类型。 + realArgs := substituteArgs(tc.Args, args) + out, err := runner.call(tc.Tool, realArgs) + results[i] = ToolRun{Name: tc.Tool, Result: out, Err: err, Order: i} + } + + // ② 执行 + if g.Parallel && n > 1 { + var wg sync.WaitGroup + for i := 0; i < n; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + run(idx) + }(i) + } + wg.Wait() + } else { + for i := 0; i < n; i++ { + run(i) + } + } + + // ③ 按**声明顺序**合并槽位(不按完成顺序 —— 见函数注释) + failed := false + var firstErr error + for _, r := range results { + res.Tools = append(res.Tools, r) + if r.Err != nil { + if firstErr == nil { + firstErr = fmt.Errorf("工具 %s 失败: %w", r.Name, r.Err) + } + if g.OnError != "continue" { + failed = true + } + } + tc := g.Tools[r.Order] + if tc.As == "" { + continue + } + // 失败时也留槽(记错误文本),否则后续组读到的是"缺失", + // 而"缺失"与"值为空"在下游难以区分。 + val := interface{}(r.Result) + if r.Err != nil { + val = "错误:" + r.Err.Error() + } + if isArrayType(g.Out[tc.As]) { + cur, _ := res.Slots[tc.As].([]interface{}) + res.Slots[tc.As] = append(cur, val) + } else { + res.Slots[tc.As] = val + } + } + + if failed { + res.Err = firstErr + return res, firstErr + } + return res, nil +} + +// substituteArgs 把 args 里的 $args.* 替换为实参值。 +// +// 两种替换形态(缺一不可): +// +// · **整值引用**:"$args.count" 整个值就是引用 ⇒ 替换为**原始值**, +// 保留类型(数字仍是数字)。否则模型会收到字符串 "3" 而非数字。 +// · **文本内插值**:"ssh $args.host" ⇒ 在字符串内替换。 +func substituteArgs(args map[string]interface{}, scope map[string]interface{}) map[string]interface{} { + if args == nil { + return nil + } + out := make(map[string]interface{}, len(args)) + for k, v := range args { + out[k] = substituteValue(v, scope) + } + return out +} + +func substituteValue(v interface{}, scope map[string]interface{}) interface{} { + switch t := v.(type) { + case string: + // 整值引用 + if key, ok := wholeRef(t); ok { + if val, present := scope[key]; present { + return val // 保留原始类型 + } + return t // 未提供则原样保留(静态校验已保证它被声明过) + } + // 文本内插值 + return interpolate(t, scope) + case map[string]interface{}: + return substituteArgs(t, scope) + case []interface{}: + arr := make([]interface{}, len(t)) + for i, e := range t { + arr[i] = substituteValue(e, scope) + } + return arr + } + return v +} + +// wholeRef 判断字符串是否**整体**是一个 $args.<键> 引用。 +func wholeRef(s string) (string, bool) { + refs := parseArgRefs(s) + if len(refs) == 1 { + trimmed := strings.TrimSpace(s) + if strings.HasPrefix(trimmed, "$args."+refs[0]) && + strings.TrimSuffix(trimmed, "$args."+refs[0]) == "" { + return refs[0], true + } + } + return "", false +} + +// interpolate 在文本内把 $args.<键> 替换为实参的**字符串形式**。 +func interpolate(s string, scope map[string]interface{}) string { + refs := parseArgRefs(s) + if len(refs) == 0 { + return s + } + var sb strings.Builder + rest := s + for { + i := strings.Index(rest, "$args.") + if i < 0 { + sb.WriteString(rest) + return sb.String() + } + sb.WriteString(rest[:i]) + rest = rest[i+len("$args."):] + end := 0 + for end < len(rest) && isRefChar(rest[end]) { + end++ + } + if end == 0 { + continue + } + key := rest[:end] + rest = rest[end:] + if val, present := scope[key]; present { + sb.WriteString(scalarToString(val)) + } else { + sb.WriteString("$args." + key) + } + } +} + +// scalarToString 渲染标量值。 +// +// ⚠️ 对象/数组用**紧凑 JSON**,绝不用 fmt.Sprintf("%v")——那会产出 +// `map[k:v]` 这种模型读不懂的 Go 语法(core 的 renderToolResult 同样约定)。 +func scalarToString(v interface{}) string { + switch t := v.(type) { + case nil: + return "" + case string: + return t + case bool: + return strconv.FormatBool(t) + case float64: + if t == float64(int64(t)) { + return strconv.FormatInt(int64(t), 10) + } + return strconv.FormatFloat(t, 'g', -1, 64) + case int: + return strconv.Itoa(t) + default: + return compactJSON(t) + } +} + +// evalCond 求值 `when` 条件。 +// +// 支持(L1+L2,见设计文档 §5): +// - `true` / `false` +// - `$args.key`(布尔真值) +// - `$args.key` 存在性与非空判断:!= "" 、== "" +// - 比较:== / != / > / < / >= / <=(标量) +// - `contains`:文本包含 +// +// ⚠️ **求值出错必须返回 error**,不得当作 false。 +// 把失败降级成"跳过"= 序列安静地少做一步,而模型以为跑完了 —— +// 与「静默吞工具」同族。 +func evalCond(cond string, args map[string]interface{}) (bool, error) { + c := strings.TrimSpace(cond) + if c == "" || c == "true" { + return true, nil + } + if c == "false" { + return false, nil + } + + // 形如 "$args.flag == true" / "$args.n > 3" / "$args.s contains x" + if left, op, right, ok := splitComparison(c); ok { + lv, err := resolveOperand(left, args) + if err != nil { + return false, err + } + rv, err := resolveOperand(right, args) + if err != nil { + return false, err + } + return compare(lv, op, rv) + } + + // 裸引用:真值判断 + if key, ok := wholeRef(c); ok { + v, present := args[key] + if !present { + return false, fmt.Errorf("引用了未提供的入参 $args.%s", key) + } + return truthy(v), nil + } + + return false, fmt.Errorf("无法解析的条件表达式 %q"+ + "(支持:true/false、$args.x、$args.x == 值、$args.x contains \"子串\"、大小比较)", cond) +} + +// splitComparison 拆出 `左 op 右`(op 需含空格,避免与 contains 前缀混淆)。 +func splitComparison(c string) (left, op, right string, ok bool) { + ops := []string{" contains ", " >= ", " <= ", " == ", " != ", " > ", " < "} + for _, o := range ops { + if i := strings.Index(c, o); i >= 0 { + return strings.TrimSpace(c[:i]), strings.TrimSpace(o), strings.TrimSpace(c[i+len(o):]), true + } + } + return "", "", "", false +} + +// resolveOperand 解析操作数:字面量或 $args 引用。 +func resolveOperand(s string, args map[string]interface{}) (interface{}, error) { + s = strings.TrimSpace(s) + if key, ok := wholeRef(s); ok { + v, present := args[key] + if !present { + return nil, fmt.Errorf("引用了未提供的入参 $args.%s", key) + } + return v, nil + } + // 字面量 + if len(s) >= 2 && (s[0] == '"' || s[0] == '\'') && s[len(s)-1] == s[0] { + return s[1 : len(s)-1], nil + } + switch strings.ToLower(s) { + case "true": + return true, nil + case "false": + return false, nil + } + if n, err := strconv.ParseFloat(s, 64); err == nil { + return n, nil + } + return s, nil // 裸文本 +} + +// compare 按运算符比较两个标量。 +func compare(l interface{}, op string, r interface{}) (bool, error) { + if op == "contains" { + ls, lok := l.(string) + rs, rok := r.(string) + if lok && rok { + return strings.Contains(ls, rs), nil + } + // 非字符串:退化为字符串化包含 + return strings.Contains(scalarToString(l), scalarToString(r)), nil + } + if op == "==" || op == "!=" { + eq := scalarToString(l) == scalarToString(r) + // 布尔与字符串宽松比较:true == "true" + if !eq { + eq = strings.EqualFold(scalarToString(l), scalarToString(r)) + } + if op == "==" { + return eq, nil + } + return !eq, nil + } + // 大小比较:两侧都必须是数值 + lf, lok := toFloat(l) + rf, rok := toFloat(r) + if !lok || !rok { + return false, fmt.Errorf("运算符 %q 两侧必须是数值(得到 %T 与 %T)", op, l, r) + } + switch op { + case ">": + return lf > rf, nil + case "<": + return lf < rf, nil + case ">=": + return lf >= rf, nil + case "<=": + return lf <= rf, nil + } + return false, fmt.Errorf("未知运算符 %q", op) +} + +func toFloat(v interface{}) (float64, bool) { + switch t := v.(type) { + case float64: + return t, true + case float32: + return float64(t), true + case int: + return float64(t), true + case int64: + return float64(t), true + case string: + f, err := strconv.ParseFloat(strings.TrimSpace(t), 64) + return f, err == nil + } + return 0, false +} + +// truthy 标量真值:非零数字、true、非空文本、非空集合。 +func truthy(v interface{}) bool { + switch t := v.(type) { + case nil: + return false + case bool: + return t + case string: + return t != "" + case float64: + return t != 0 + case int: + return t != 0 + case []interface{}: + return len(t) > 0 + case map[string]interface{}: + return len(t) > 0 + } + return true +} + +// groupTimeout 返回本组的墙钟上限(Group.Timeout 形如 "30s")。 +func groupTimeout(g Group) time.Duration { + if g.Timeout == "" { + return 0 // 由调用方决定默认 + } + d, err := time.ParseDuration(g.Timeout) + if err != nil { + return 0 + } + return d +} + +// compactJSON 渲染结构化值为紧凑 JSON。 +// +// ⚠️ 绝不用 fmt.Sprintf("%v"):那会产出 `map[k:v]` 这种模型读不懂的 +// Go 语法(core 的 renderToolResult 同样约定,见设计文档 §5.2)。 +func compactJSON(v interface{}) string { + b, err := json.Marshal(v) + if err != nil { + return fmt.Sprintf("%v", v) + } + return string(b) +} diff --git a/internal/plugins/seq/exec_test.go b/internal/plugins/seq/exec_test.go new file mode 100644 index 0000000..9450d60 --- /dev/null +++ b/internal/plugins/seq/exec_test.go @@ -0,0 +1,341 @@ +package seq + +import ( + "errors" + "strings" + "sync" + "testing" + "time" +) + +// 阶段 P2:执行引擎(组内并行 + 具名槽 + 条件求值)。 +// +// 三个核心不变量: +// 1. **组边界即屏障**:同组内工具互相不可见(并行 ⇒ 无确定写序); +// 变量只在组屏障处按 tools 顺序**确定性合并**。 +// 2. **条件求值失败必须报错**,不得降级成「条件为假」—— +// 那会让序列安静地少做一步而模型以为跑完了。 +// 3. 结果**按 tools 数组顺序**合并,与完成顺序无关 ⇒ 可复现。 + +// fakeTool 记录一次工具执行,并返回可配置的文本。 +type fakeTool struct { + mu sync.Mutex + calls []string + // results 按工具名给出返回文本;未配置则返回 "ran:" + results map[string]string + // errs 按工具名给出错误 + errs map[string]error + // delay 用于制造"完成顺序 ≠ 声明顺序" + delay map[string]int +} + +func newFakeTool() *fakeTool { + return &fakeTool{results: map[string]string{}, errs: map[string]error{}, delay: map[string]int{}} +} + +func (f *fakeTool) call(name string, _ map[string]interface{}) (string, error) { + if d := f.delay[name]; d > 0 { + sleepMS(d) + } + f.mu.Lock() + f.calls = append(f.calls, name) + f.mu.Unlock() + if e, ok := f.errs[name]; ok { + return "", e + } + if r, ok := f.results[name]; ok { + return r, nil + } + return "ran:" + name, nil +} + +func (f *fakeTool) called() []string { + f.mu.Lock() + defer f.mu.Unlock() + out := make([]string, len(f.calls)) + copy(out, f.calls) + return out +} + +// ① 具名槽:工具结果按 `as` 写入对应槽。 +func TestExecWritesNamedSlots(t *testing.T) { + ft := newFakeTool() + g := Group{ + Name: "g", + In: map[string]string{}, + Out: map[string]string{"summary": "string", "load": "string"}, + When: "true", + Parallel: false, + Tools: []ToolCall{ + {Tool: "uptime", As: "summary"}, + {Tool: "top", As: "load"}, + }, + } + res, err := execGroup(g, map[string]interface{}{}, ft) + if err != nil { + t.Fatalf("execGroup: %v", err) + } + if got, _ := res.Slots["summary"].(string); got != "ran:uptime" { + t.Errorf("summary 槽 = %v", res.Slots["summary"]) + } + if got, _ := res.Slots["load"].(string); got != "ran:top" { + t.Errorf("load 槽 = %v", res.Slots["load"]) + } +} + +// ② ★ 结果合并必须按 tools 数组顺序,与完成顺序无关。 +// +// 并行下完成顺序不确定;若按完成顺序合并,同样的输入会产出不同的槽内容, +// 整条序列**不可复现**。这里让后声明的工具先完成(delay 更短)。 +func TestExecMergesSlotsInDeclarationOrder(t *testing.T) { + ft := newFakeTool() + ft.delay["first"] = 30 // 先声明但后完成 + ft.delay["second"] = 1 + ft.results["first"] = "FIRST" + ft.results["second"] = "SECOND" + + g := Group{ + Name: "g", In: map[string]string{}, + Out: map[string]string{"a": "array", "b": "array"}, + When: "true", Parallel: true, + Tools: []ToolCall{ + {Tool: "first", As: "a"}, + {Tool: "second", As: "b"}, + }, + } + res, err := execGroup(g, map[string]interface{}{}, ft) + if err != nil { + t.Fatalf("execGroup: %v", err) + } + // 各自的槽内容正确 + ra, _ := res.Slots["a"].([]interface{}) + rb, _ := res.Slots["b"].([]interface{}) + if len(ra) != 1 || ra[0] != "FIRST" { + t.Errorf("槽 a = %v", ra) + } + if len(rb) != 1 || rb[0] != "SECOND" { + t.Errorf("槽 b = %v", rb) + } + // 完成顺序确实被打乱(否则本用例测不到并发) + if got := ft.called(); got[0] != "first" || got[1] != "second" { + t.Logf("完成顺序 = %v(未打乱,判据可能测不到并发)", got) + } +} + +// ③ array 槽的同名 as:在组屏障按 tools 顺序**确定性追加**。 +func TestExecArraySlotAppendsInOrder(t *testing.T) { + ft := newFakeTool() + ft.delay["slow"] = 30 + ft.results["a"] = "A" + ft.results["c"] = "C" + // "slow" 用默认返回值即可(ran:slow),它只用来制造"最后完成" + + g := Group{ + Name: "g", In: map[string]string{}, + Out: map[string]string{"xs": "array"}, + When: "true", Parallel: true, + Tools: []ToolCall{ + {Tool: "a", As: "xs"}, + {Tool: "slow", As: "xs"}, // 声明在中间但最后完成 + {Tool: "c", As: "xs"}, + }, + } + res, err := execGroup(g, map[string]interface{}{}, ft) + if err != nil { + t.Fatalf("execGroup: %v", err) + } + xs, _ := res.Slots["xs"].([]interface{}) + if len(xs) != 3 { + t.Fatalf("xs 应有 3 项,实际 %d(%v)", len(xs), xs) + } + // 必须按 tools 声明顺序:a, slow, c ⇒ A, B, C + // 期望按 tools 声明顺序:a → slow → c + want := []string{"A", "ran:slow", "C"} + for i, w := range want { + if xs[i] != w { + t.Errorf("xs[%d] = %v,期望 %v(完整 %v)—— 合并顺序依赖了完成顺序", i, xs[i], w, xs) + } + } +} + +// ④ 条件为假 ⇒ 整组跳过,**槽不赋值**。 +func TestExecSkipsGroupWhenConditionFalse(t *testing.T) { + ft := newFakeTool() + g := Group{ + Name: "g", In: map[string]string{"flag": "bool"}, + Out: map[string]string{"x": "string"}, + When: "false", Parallel: false, + Tools: []ToolCall{{Tool: "uptime", As: "x"}}, + } + res, err := execGroup(g, map[string]interface{}{"flag": false}, ft) + if err != nil { + t.Fatalf("execGroup: %v", err) + } + if !res.Skipped { + t.Error("条件为假时应标记 Skipped") + } + if len(ft.called()) != 0 { + t.Errorf("条件为假却执行了工具: %v", ft.called()) + } + if _, ok := res.Slots["x"]; ok { + t.Error("条件为假时不应给槽赋值(后续组会读到不存在的值)") + } +} + +// ⑤ ★ 条件**求值出错**必须报错,不得降级成「条件为假」。 +// +// 这是本阶段最关键的一条:把求值失败降级为跳过 = 序列安静地少做一步, +// 而模型以为跑完了 —— 与「静默吞工具」同族。 +func TestExecErrorsOnMalformedCondition(t *testing.T) { + cases := []string{ + "$args.", // 空键名 + "$args.missing == 1", // 引用了未声明的入参(in 里只有 flag) + "1 ==", // 语法不完整 + } + for _, cond := range cases { + t.Run(cond, func(t *testing.T) { + ft := newFakeTool() + g := Group{ + Name: "g", In: map[string]string{"flag": "bool"}, + Out: map[string]string{"x": "string"}, + When: cond, Parallel: false, + Tools: []ToolCall{{Tool: "uptime", As: "x"}}, + } + _, err := execGroup(g, map[string]interface{}{"flag": false}, ft) + if err == nil { + t.Fatalf("畸形条件 %q 未被拒绝(被静默当成假了吗)", cond) + } + if !strings.Contains(err.Error(), "条件") { + t.Errorf("错误信息应提到『条件』,实际: %v", err) + } + if len(ft.called()) != 0 { + t.Errorf("条件求值失败却执行了工具: %v", ft.called()) + } + }) + } +} + +// ⑥ 条件为真时正常执行。 +func TestExecRunsWhenConditionTrue(t *testing.T) { + ft := newFakeTool() + g := Group{ + Name: "g", In: map[string]string{"flag": "bool"}, + Out: map[string]string{"x": "string"}, + When: "$args.flag == true", Parallel: false, + Tools: []ToolCall{{Tool: "uptime", As: "x"}}, + } + res, err := execGroup(g, map[string]interface{}{"flag": true}, ft) + if err != nil { + t.Fatalf("execGroup: %v", err) + } + if res.Skipped { + t.Error("条件为真却跳过了") + } + if got, _ := res.Slots["x"].(string); got != "ran:uptime" { + t.Errorf("x = %v", res.Slots["x"]) + } +} + +// ⑦ 工具执行失败:on_error=continue 时继续,abort 时整组失败。 +func TestExecOnErrorPolicy(t *testing.T) { + boom := errors.New("boom") + + t.Run("abort", func(t *testing.T) { + ft := newFakeTool() + ft.errs["bad"] = boom + g := Group{ + Name: "g", In: map[string]string{}, + Out: map[string]string{"a": "string", "b": "string"}, + When: "true", Parallel: false, OnError: "abort", + Tools: []ToolCall{{Tool: "bad", As: "a"}, {Tool: "ok", As: "b"}}, + } + if _, err := execGroup(g, map[string]interface{}{}, ft); err == nil { + t.Error("on_error=abort 时失败应使整组失败") + } + }) + + t.Run("continue", func(t *testing.T) { + ft := newFakeTool() + ft.errs["bad"] = boom + g := Group{ + Name: "g", In: map[string]string{}, + Out: map[string]string{"a": "string", "b": "string"}, + When: "true", Parallel: false, OnError: "continue", + Tools: []ToolCall{{Tool: "bad", As: "a"}, {Tool: "ok", As: "b"}}, + } + res, err := execGroup(g, map[string]interface{}{}, ft) + if err != nil { + t.Fatalf("on_error=continue 不应整组失败: %v", err) + } + if got, _ := res.Slots["b"].(string); got != "ran:ok" { + t.Errorf("continue 下后续工具应仍执行,b = %v", res.Slots["b"]) + } + // 失败的槽也要有值(错误文本),否则后续组读到缺失 + if _, ok := res.Slots["a"]; !ok { + t.Error("continue 下失败的工具也应留下槽(记错误文本)") + } + }) +} + +// ⑧ 变量插值:`$args.x` 被真实值替换,且**整值引用**保留类型 +// (数字仍是数字,不是字符串)。 +func TestExecSubstitutesArgs(t *testing.T) { + ft := newFakeTool() + g := Group{ + Name: "g", + In: map[string]string{"host": "string", "count": "integer"}, + Out: map[string]string{"x": "string"}, + When: "true", Parallel: false, + Tools: []ToolCall{{ + Tool: "cmd", + Args: map[string]interface{}{ + "cmd": "ssh $args.host", + "amount": "$args.count", + "whole": "$args.host", + }, + As: "x", + }}, + } + // 注意:本例只验证 execTool 前的插值由 call 接口承担, + // 这里通过 fakeTool 捕获实际收到的 args。 + capture := &capturingTool{inner: ft} + if _, err := execGroup(g, map[string]interface{}{"host": "node-a", "count": 3}, capture); err != nil { + t.Fatalf("execGroup: %v", err) + } + got := capture.lastArgs() + if got["cmd"] != "ssh node-a" { + t.Errorf("字符串内插值失败: %v", got["cmd"]) + } + // 整值引用必须保留原始类型(数字仍是数字) + if n, ok := got["amount"].(int); !ok || n != 3 { + t.Errorf("整值引用应保留 int 类型,实际 %#v", got["amount"]) + } + if got["whole"] != "node-a" { + t.Errorf("整值引用失败: %#v", got["whole"]) + } +} + +// capturingTool 记录最后一次收到的 args。 +type capturingTool struct { + inner *fakeTool + mu sync.Mutex + args map[string]interface{} +} + +func (c *capturingTool) call(name string, args map[string]interface{}) (string, error) { + c.mu.Lock() + c.args = args + c.mu.Unlock() + return c.inner.call(name, args) +} +func (c *capturingTool) lastArgs() map[string]interface{} { + c.mu.Lock() + defer c.mu.Unlock() + return c.args +} + +// sleepMS 测试用的短延时(毫秒)。 +func sleepMS(ms int) { time.Sleep(time.Duration(ms) * time.Millisecond) } + +var _ toolRunner = (*fakeTool)(nil) +var _ toolRunner = (*capturingTool)(nil)