Files
krow_backend/go-api/internal/tools/registry_test.go
2026-08-28 12:21:44 +05:30

181 lines
6.2 KiB
Go

package tools_test
import (
"context"
"encoding/json"
"strings"
"testing"
"github.com/krow/krow-backend/go-api/internal/tools"
)
func stub(name string, effect tools.Effect, h tools.Handler) tools.Tool {
if h == nil {
h = func(context.Context, tools.Context, json.RawMessage) tools.Result {
return tools.OK(map[string]any{"ok": true})
}
}
t := tools.Tool{
Name: name, Description: "A stub.", Effect: effect,
InputSchema: map[string]any{"type": "object"}, Handler: h,
}
if effect == tools.EffectWrite {
t.Confirm = describeStub
}
return t
}
// describeStub is the minimum a write must be able to say about itself.
func describeStub(context.Context, tools.Context, json.RawMessage) (*tools.Confirmation, *tools.Result) {
return &tools.Confirmation{Title: "Do the thing", Summary: "The thing will be done."}, nil
}
func TestWriteToolsAlwaysRequireConfirmation(t *testing.T) {
// I4, and the reason it is forced rather than trusted: a write that forgot
// to set the flag would run unconfirmed forever, indistinguishable from a
// deliberate choice.
reg := tools.NewRegistry()
w := stub("send_shift_offer", tools.EffectWrite, nil)
w.RequiresConfirmation = false // an author trying to opt out
reg.MustRegister(w)
got, _ := reg.Get("send_shift_offer")
if !got.RequiresConfirmation {
t.Fatal("a write tool must require confirmation regardless of what its author declared")
}
res := reg.Dispatch(context.Background(), tools.Context{}, "send_shift_offer", nil)
if res.Confirmation == nil {
t.Fatal("an unconfirmed write must come back as a question, not run")
}
if res.Data != nil {
t.Fatal("an unconfirmed write must not produce a result")
}
}
func TestAWriteWithNoRendererIsRefusedAtRegistration(t *testing.T) {
// §8 step 3. A confirmation nobody can read is a click, not a confirmation,
// and the only honest dialog for a tool with no renderer would show raw
// arguments. Caught at boot rather than in production.
reg := tools.NewRegistry()
w := stub("silent_write", tools.EffectWrite, nil)
w.Confirm = nil
if err := reg.Register(w); err == nil {
t.Fatal("a write with no Confirm renderer must be refused")
}
}
func TestRegisterRejectsMalformedTools(t *testing.T) {
reg := tools.NewRegistry()
cases := map[string]tools.Tool{
"no name": {Description: "x", Effect: tools.EffectRead, Handler: stub("x", tools.EffectRead, nil).Handler},
"no handler": {Name: "a", Description: "x", Effect: tools.EffectRead},
"no description": {Name: "b", Effect: tools.EffectRead, Handler: stub("b", tools.EffectRead, nil).Handler},
"bad effect": {Name: "c", Description: "x", Effect: "maybe", Handler: stub("c", tools.EffectRead, nil).Handler},
}
for name, tool := range cases {
if err := reg.Register(tool); err == nil {
t.Errorf("%s: should have been refused", name)
}
}
if err := reg.Register(stub("ok", tools.EffectRead, nil)); err != nil {
t.Errorf("a well-formed tool was refused: %v", err)
}
if err := reg.Register(stub("ok", tools.EffectRead, nil)); err == nil {
t.Error("a duplicate name should be refused")
}
}
func TestResolveEnforcesThePerAgentCap(t *testing.T) {
reg := tools.NewRegistry()
names := make([]string, 0, tools.MaxToolsPerAgent+1)
for i := 0; i <= tools.MaxToolsPerAgent; i++ {
n := string(rune('a'+i%26)) + strings.Repeat("x", i)
reg.MustRegister(stub(n, tools.EffectRead, nil))
names = append(names, n)
}
if _, _, err := reg.Resolve(names); err == nil {
t.Errorf("%d tools should exceed the cap of %d", len(names), tools.MaxToolsPerAgent)
}
}
func TestResolveReportsUnknownNames(t *testing.T) {
reg := tools.NewRegistry()
reg.MustRegister(stub("real", tools.EffectRead, nil))
resolved, unknown, err := reg.Resolve([]string{"real", "imaginary", "real"})
if err != nil {
t.Fatalf("resolve: %v", err)
}
if len(resolved) != 1 {
t.Errorf("resolved %d tools, want 1 — a repeated name is not two tools", len(resolved))
}
if len(unknown) != 1 || unknown[0] != "imaginary" {
t.Errorf("unknown = %v, want [imaginary] — a missing name must be reported, not dropped", unknown)
}
}
func TestDispatchContainsAPanickingHandler(t *testing.T) {
// A tool is the least trusted code in the runtime. One bad handler costs
// its own call, not the run.
reg := tools.NewRegistry()
reg.MustRegister(stub("explodes", tools.EffectRead,
func(context.Context, tools.Context, json.RawMessage) tools.Result {
panic("boom")
}))
res := reg.Dispatch(context.Background(), tools.Context{}, "explodes", nil)
if res.Error == nil || res.Error.Code != tools.CodeFailed {
t.Fatalf("a panicking handler should return a failure, got %+v", res)
}
}
func TestDispatchWithholdsOversizedResults(t *testing.T) {
// Cut JSON is not JSON, and a model handed a fragment reasons about it as
// if it were whole. The result is withheld and the fact declared.
reg := tools.NewRegistry()
big := strings.Repeat("x", 2000)
tool := stub("huge", tools.EffectRead,
func(context.Context, tools.Context, json.RawMessage) tools.Result {
return tools.OK(map[string]any{"blob": big})
})
tool.MaxResultBytes = 500
reg.MustRegister(tool)
res := reg.Dispatch(context.Background(), tools.Context{}, "huge", nil)
if !res.Truncated {
t.Fatal("an oversized result must be marked truncated")
}
encoded, _ := json.Marshal(res.Data)
if strings.Contains(string(encoded), big) {
t.Error("the oversized payload was returned anyway")
}
if !strings.Contains(string(encoded), "narrower") {
t.Error("the model was not told how to ask again")
}
}
func TestDispatchUnknownToolIsAResultNotAPanic(t *testing.T) {
reg := tools.NewRegistry()
res := reg.Dispatch(context.Background(), tools.Context{}, "nope", nil)
if res.Error == nil || res.Error.Code != tools.CodeUnavailable {
t.Fatalf("want unavailable, got %+v", res)
}
}
func TestNamesAreSortedForCacheStability(t *testing.T) {
// The tool list is part of the cached prefix. An unstable order would
// invalidate the cache on every request for no reason at all.
reg := tools.NewRegistry()
for _, n := range []string{"zulu", "alpha", "mike"} {
reg.MustRegister(stub(n, tools.EffectRead, nil))
}
got := reg.Names()
want := []string{"alpha", "mike", "zulu"}
for i := range want {
if got[i] != want[i] {
t.Fatalf("Names() = %v, want %v", got, want)
}
}
}