Files
krow_backend/go-api/internal/runtime/smalltalk_test.go
Aravind a372340281
Some checks failed
CI / test (push) Failing after 4m41s
CI / fixture (push) Failing after 9s
add the greeting msg
2026-10-05 16:20:17 +05:30

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)
}
}