181 lines
6.2 KiB
Go
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)
|
|
}
|
|
}
|
|
}
|