179 lines
5.8 KiB
Go
179 lines
5.8 KiB
Go
package runtime
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/krow/krow-backend/go-api/internal/gateway"
|
|
"github.com/krow/krow-backend/go-api/internal/tools"
|
|
)
|
|
|
|
// TestIsSmalltalkGreetings covers what the fix is for: the messages that were
|
|
// costing a full operational turn.
|
|
func TestIsSmalltalkGreetings(t *testing.T) {
|
|
for _, q := range []string{
|
|
"hi", "Hi", "HI", "hi!", "hi.", " hi ", "hi :)", "hi 👋",
|
|
"hello", "Hello!", "hey", "Hey there", "hi there",
|
|
"good morning", "Good Morning!", "good evening",
|
|
"thanks", "Thank you", "thank-you", "thank you",
|
|
"bye", "Goodbye", "good night",
|
|
"how are you", "How's it going?",
|
|
} {
|
|
if !isSmalltalk(q) {
|
|
t.Errorf("isSmalltalk(%q) = false, want true", q)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestIsSmalltalkRealQuestions is the half that matters for correctness. A
|
|
// false positive answers a real operational question with a greeting, so
|
|
// anything carrying a request must fall through — including the ones that
|
|
// merely START with a greeting.
|
|
func TestIsSmalltalkRealQuestions(t *testing.T) {
|
|
for _, q := range []string{
|
|
"", " ",
|
|
"how many open positions?",
|
|
"what happened today",
|
|
"hi, how many open positions?",
|
|
"hello there, which shifts are uncovered?",
|
|
"hey can you check the screening backlog",
|
|
"thanks — now show me the overtime report",
|
|
"good morning, what needs attention right now?",
|
|
"say hi to the new starters",
|
|
"how are you handling the uncovered shifts",
|
|
"bye week coverage",
|
|
} {
|
|
if isSmalltalk(q) {
|
|
t.Errorf("isSmalltalk(%q) = true, want false", q)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestNormaliseSmalltalk(t *testing.T) {
|
|
cases := map[string]string{
|
|
"Hi!": "hi",
|
|
" HELLO ": "hello",
|
|
"thank-you": "thank you",
|
|
"thank you": "thank you",
|
|
"How's it go?": "hows it go",
|
|
"👋": "",
|
|
"shift 12": "shift",
|
|
}
|
|
for in, want := range cases {
|
|
if got := normaliseSmalltalk(in); got != want {
|
|
t.Errorf("normaliseSmalltalk(%q) = %q, want %q", in, got, want)
|
|
}
|
|
}
|
|
}
|
|
|
|
// TestSmalltalkSendsNoToolsAndSkipsRetrieval is the fix as the user meets it:
|
|
// "hi" reaches the model as "hi", with nothing attached.
|
|
func TestSmalltalkSendsNoToolsAndSkipsRetrieval(t *testing.T) {
|
|
var toolCalls int
|
|
reg := tools.NewRegistry()
|
|
reg.MustRegister(countingTool("activity_breakdown", &toolCalls))
|
|
|
|
ret := &scriptedRetriever{results: onePassage("Staff must arrive fifteen minutes early.")}
|
|
gw := &scriptedGateway{}
|
|
|
|
agent := knowledgeAgent()
|
|
agent.Tools = []string{"activity_breakdown"}
|
|
|
|
exec := NewModelExecutor(gw, &MemorySink{}, reg).WithRetriever(ret)
|
|
res, err := exec.ExecuteAgent(context.Background(), agent, testInput("Hi"))
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if res.Termination != TerminationCompleted {
|
|
t.Fatalf("Termination = %q, want Completed", res.Termination)
|
|
}
|
|
|
|
if ret.calls != 0 {
|
|
t.Errorf("retriever was called %d times for a greeting; want 0", ret.calls)
|
|
}
|
|
if len(gw.seen) != 1 {
|
|
t.Fatalf("a greeting took %d model calls, want 1", len(gw.seen))
|
|
}
|
|
req := gw.seen[0]
|
|
if len(req.Tools) != 0 {
|
|
t.Errorf("greeting carried %d tool definitions, want 0", len(req.Tools))
|
|
}
|
|
if req.Messages[0].Text != "Hi" {
|
|
t.Errorf("model saw %q, want the bare greeting", req.Messages[0].Text)
|
|
}
|
|
if !strings.Contains(req.System, "greeted you") {
|
|
t.Error("the smalltalk directive did not reach the system prompt")
|
|
}
|
|
if req.MaxOutputTokens != smalltalkMaxOutputTokens {
|
|
t.Errorf("MaxOutputTokens = %d, want the smalltalk cap %d",
|
|
req.MaxOutputTokens, smalltalkMaxOutputTokens)
|
|
}
|
|
}
|
|
|
|
// TestOperationalQuestionKeepsToolsAndRetrieval is the guard on the fix above.
|
|
// The cheap path must not swallow a question that needs evidence.
|
|
func TestOperationalQuestionKeepsToolsAndRetrieval(t *testing.T) {
|
|
var toolCalls int
|
|
reg := tools.NewRegistry()
|
|
reg.MustRegister(countingTool("activity_breakdown", &toolCalls))
|
|
|
|
ret := &scriptedRetriever{results: onePassage("Staff must arrive fifteen minutes early.")}
|
|
gw := &scriptedGateway{}
|
|
|
|
agent := knowledgeAgent()
|
|
agent.Tools = []string{"activity_breakdown"}
|
|
|
|
exec := NewModelExecutor(gw, &MemorySink{}, reg).WithRetriever(ret)
|
|
if _, err := exec.ExecuteAgent(
|
|
context.Background(), agent, testInput("hi, how many open positions?"),
|
|
); err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
if ret.calls != 1 {
|
|
t.Errorf("retriever called %d times for a real question, want 1", ret.calls)
|
|
}
|
|
if len(gw.seen[0].Tools) == 0 {
|
|
t.Error("a real question was sent with no tools")
|
|
}
|
|
}
|
|
|
|
// TestCatalogueWithheldOnceToolBudgetIsSpent covers the other half of the cost
|
|
// work: the synthesis turn at the end of a tool-using run is sent without a
|
|
// catalogue the model is no longer permitted to use.
|
|
func TestCatalogueWithheldOnceToolBudgetIsSpent(t *testing.T) {
|
|
var toolCalls int
|
|
reg := tools.NewRegistry()
|
|
reg.MustRegister(countingTool("activity_breakdown", &toolCalls))
|
|
|
|
gw := &scriptedGateway{steps: []*gateway.Response{{
|
|
ToolCalls: []gateway.ToolCall{{ID: "call_1", Name: "activity_breakdown", Input: json.RawMessage(`{}`)}},
|
|
StopReason: "tool_use", Model: "fake-model",
|
|
}}}
|
|
|
|
agent := testAgent()
|
|
agent.Tools = []string{"activity_breakdown"}
|
|
|
|
exec := NewModelExecutor(gw, &MemorySink{}, reg)
|
|
res, err := exec.executeWithLimits(context.Background(), agent, testInput("what happened today"),
|
|
Limits{MaxSteps: 4, MaxToolCalls: 1, MaxTokens: 100_000, Deadline: 30 * time.Second})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if res.Termination != TerminationCompleted {
|
|
t.Fatalf("Termination = %q, want Completed", res.Termination)
|
|
}
|
|
if len(gw.seen) < 2 {
|
|
t.Fatalf("expected at least 2 model calls, got %d", len(gw.seen))
|
|
}
|
|
if len(gw.seen[0].Tools) == 0 {
|
|
t.Error("the first call must offer the catalogue")
|
|
}
|
|
if n := len(gw.seen[len(gw.seen)-1].Tools); n != 0 {
|
|
t.Errorf("the final call carried %d tool definitions; the tool budget was spent", n)
|
|
}
|
|
}
|