287 lines
11 KiB
Go
287 lines
11 KiB
Go
package runtime
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"testing"
|
|
|
|
"github.com/smallnest/pigo/internal/agentcore"
|
|
"github.com/smallnest/pigo/internal/agenttool"
|
|
"github.com/smallnest/pigo/internal/provider"
|
|
)
|
|
|
|
// collectStream drains a LoopEventStream, returning the event types in order
|
|
// and the run result messages.
|
|
func collectStream(t *testing.T, s *LoopEventStream) ([]string, []agentcore.AgentMessage) {
|
|
t.Helper()
|
|
var kinds []string
|
|
for ev := range s.Events() {
|
|
kinds = append(kinds, ev.EventType())
|
|
}
|
|
msgs, err := s.Result(context.Background())
|
|
if err != nil {
|
|
t.Fatalf("stream result: %v", err)
|
|
}
|
|
return kinds, msgs
|
|
}
|
|
|
|
// oneToolAssistant builds an assistant message with a single tool call.
|
|
func oneToolAssistant(id, name string) agentcore.AssistantMessage {
|
|
return agentcore.AssistantMessage{
|
|
RoleField: agentcore.RoleAssistant,
|
|
StopReason: agentcore.StopReasonToolUse,
|
|
Content: agentcore.ContentList{agentcore.NewToolCallContent(id, name, json.RawMessage(`{}`))},
|
|
}
|
|
}
|
|
|
|
// scriptedStream returns a StreamFn that emits one StreamDoneEvent per call,
|
|
// consuming msgs in order. Extra calls beyond msgs emit a plain end_turn.
|
|
func scriptedStream(msgs []agentcore.AssistantMessage) provider.StreamFn {
|
|
i := 0
|
|
return func(ctx context.Context, model string, llm provider.LlmContext, cfg provider.StreamConfig) (*provider.AssistantMessageEventStream, error) {
|
|
var msg agentcore.AssistantMessage
|
|
if i < len(msgs) {
|
|
msg = msgs[i]
|
|
} else {
|
|
msg = agentcore.AssistantMessage{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn}
|
|
}
|
|
i++
|
|
s := provider.NewAssistantMessageEventStream(0)
|
|
go func() {
|
|
_ = s.Emit(ctx, provider.StreamDoneEvent{Message: msg})
|
|
s.Close()
|
|
}()
|
|
return s, nil
|
|
}
|
|
}
|
|
|
|
func newRunCfg(stream provider.StreamFn, tools ...agentcore.AgentTool) RunConfig {
|
|
reg := agenttool.NewToolRegistry()
|
|
for _, tl := range tools {
|
|
_ = reg.Register(tl)
|
|
}
|
|
return RunConfig{
|
|
LoopConfig: LoopConfig{Model: "fake", Stream: stream},
|
|
Batch: agenttool.BatchConfig{ToolExecutorConfig: agenttool.ToolExecutorConfig{Registry: reg}},
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopNoToolCallsSingleTurn(t *testing.T) {
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn, Content: agentcore.ContentList{agentcore.NewTextContent("hi")}},
|
|
}))
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
|
|
kinds, msgs := collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
|
|
want := []string{agentcore.EventAgentStart, agentcore.EventTurnStart, agentcore.EventMessageEnd, agentcore.EventTurnEnd, agentcore.EventTelemetry, agentcore.EventAgentEnd}
|
|
assertEventKinds(t, kinds, want)
|
|
if len(msgs) != 1 {
|
|
t.Fatalf("run produced %d messages, want 1: %+v", len(msgs), msgs)
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopInnerLoopFeedsToolResults(t *testing.T) {
|
|
// Turn 1: tool call. Turn 2: no tool call → inner loop ends.
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
oneToolAssistant("c1", "echo"),
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn, Content: agentcore.ContentList{agentcore.NewTextContent("done")}},
|
|
}), echoTool("echo", agentcore.ToolExecutionParallel, false))
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
|
|
kinds, msgs := collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
|
|
// Two turns; a tool executed in the first.
|
|
if countKind(kinds, agentcore.EventTurnStart) != 2 {
|
|
t.Errorf("expected 2 turns, got kinds %v", kinds)
|
|
}
|
|
if countKind(kinds, agentcore.EventToolExecutionEnd) != 1 {
|
|
t.Errorf("expected 1 tool execution, got kinds %v", kinds)
|
|
}
|
|
// Messages produced: assistant(tool) + toolResult + assistant(done) = 3.
|
|
if len(msgs) != 3 {
|
|
t.Fatalf("expected 3 new messages, got %d: %+v", len(msgs), msgs)
|
|
}
|
|
if _, ok := msgs[1].(agentcore.ToolResultMessage); !ok {
|
|
t.Errorf("expected message[1] to be a tool result, got %T", msgs[1])
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopFollowUpMessagesContinue(t *testing.T) {
|
|
served := false
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn, Content: agentcore.ContentList{agentcore.NewTextContent("first")}},
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn, Content: agentcore.ContentList{agentcore.NewTextContent("second")}},
|
|
}))
|
|
cfg.GetFollowUpMessages = func(ctx context.Context, agentCtx *agentcore.AgentContext) []agentcore.AgentMessage {
|
|
if served {
|
|
return nil
|
|
}
|
|
served = true
|
|
return []agentcore.AgentMessage{agentcore.UserMessage{RoleField: agentcore.RoleUser, Content: agentcore.ContentList{agentcore.NewTextContent("more")}}}
|
|
}
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
|
|
kinds, _ := collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
if countKind(kinds, agentcore.EventTurnStart) != 2 {
|
|
t.Errorf("follow-up should drive a second turn, got kinds %v", kinds)
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopShouldStopAfterTurn(t *testing.T) {
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
oneToolAssistant("c1", "echo"),
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn},
|
|
}), echoTool("echo", agentcore.ToolExecutionParallel, false))
|
|
cfg.ShouldStopAfterTurn = func(ctx context.Context, agentCtx *agentcore.AgentContext) bool { return true }
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
|
|
kinds, _ := collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
// Stops after the first turn_end, so only one turn.
|
|
if countKind(kinds, agentcore.EventTurnStart) != 1 {
|
|
t.Errorf("shouldStopAfterTurn=true must stop after one turn, got %v", kinds)
|
|
}
|
|
if kinds[len(kinds)-1] != agentcore.EventAgentEnd {
|
|
t.Errorf("run must end with agent_end, got %v", kinds)
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopSteeringInjected(t *testing.T) {
|
|
var injectedSeen bool
|
|
steer := agentcore.UserMessage{RoleField: agentcore.RoleUser, Content: agentcore.ContentList{agentcore.NewTextContent("steer")}}
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
oneToolAssistant("c1", "echo"),
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn},
|
|
}), echoTool("echo", agentcore.ToolExecutionParallel, false))
|
|
pulled := false
|
|
cfg.GetSteeringMessages = func(ctx context.Context) []agentcore.AgentMessage {
|
|
if pulled {
|
|
return nil
|
|
}
|
|
pulled = true
|
|
return []agentcore.AgentMessage{steer}
|
|
}
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
for _, m := range agentCtx.Messages {
|
|
if um, ok := m.(agentcore.UserMessage); ok && len(um.Content) == 1 {
|
|
if tc, ok := um.Content[0].(agentcore.TextContent); ok && tc.Text == "steer" {
|
|
injectedSeen = true
|
|
}
|
|
}
|
|
}
|
|
if !injectedSeen {
|
|
t.Errorf("steering message was not injected into the context")
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopPrepareNextTurnSwapsModel(t *testing.T) {
|
|
var seenModels []string
|
|
streamFn := func(ctx context.Context, model string, llm provider.LlmContext, cfg provider.StreamConfig) (*provider.AssistantMessageEventStream, error) {
|
|
seenModels = append(seenModels, model)
|
|
var msg agentcore.AssistantMessage
|
|
if len(seenModels) == 1 {
|
|
msg = oneToolAssistant("c1", "echo")
|
|
} else {
|
|
msg = agentcore.AssistantMessage{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn}
|
|
}
|
|
s := provider.NewAssistantMessageEventStream(0)
|
|
go func() { _ = s.Emit(ctx, provider.StreamDoneEvent{Message: msg}); s.Close() }()
|
|
return s, nil
|
|
}
|
|
cfg := newRunCfg(streamFn, echoTool("echo", agentcore.ToolExecutionParallel, false))
|
|
newModel := "swapped-model"
|
|
cfg.PrepareNextTurn = func(ctx context.Context, agentCtx *agentcore.AgentContext) *TurnUpdate {
|
|
return &TurnUpdate{Model: &newModel}
|
|
}
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
if len(seenModels) != 2 || seenModels[1] != newModel {
|
|
t.Errorf("prepareNextTurn should swap model to %q, saw %v", newModel, seenModels)
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopLengthFailsToolCalls(t *testing.T) {
|
|
// Turn 1: tool call but truncated (length). Turn 2: end.
|
|
truncated := oneToolAssistant("c1", "echo")
|
|
truncated.StopReason = agentcore.StopReasonLength
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
truncated,
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn},
|
|
}), echoTool("echo", agentcore.ToolExecutionParallel, false))
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
|
|
kinds, msgs := collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
// The tool must NOT have executed (truncated → failed instead).
|
|
if countKind(kinds, agentcore.EventToolExecutionEnd) != 0 {
|
|
t.Errorf("truncated message must not execute tools, got %v", kinds)
|
|
}
|
|
// A failed tool result must have been synthesized.
|
|
var foundFail bool
|
|
for _, m := range msgs {
|
|
if tr, ok := m.(agentcore.ToolResultMessage); ok && tr.IsError && tr.ToolCallID == "c1" {
|
|
foundFail = true
|
|
}
|
|
}
|
|
if !foundFail {
|
|
t.Errorf("expected a synthesized failed tool result for the truncated call")
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopErrorStopEndsRun(t *testing.T) {
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonError, ErrorMessage: "boom"},
|
|
}))
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
|
|
kinds, _ := collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
if countKind(kinds, agentcore.EventTurnStart) != 1 {
|
|
t.Errorf("error stop must end after one turn, got %v", kinds)
|
|
}
|
|
if kinds[len(kinds)-1] != agentcore.EventAgentEnd {
|
|
t.Errorf("run must end with agent_end, got %v", kinds)
|
|
}
|
|
}
|
|
|
|
func TestAgentLoopAllTerminateStopsRun(t *testing.T) {
|
|
term := true
|
|
termTool := execTool{
|
|
name: "quit",
|
|
run: func(ctx context.Context, id string, args json.RawMessage, onUpdate agentcore.ToolUpdateFunc) (agentcore.AgentToolResult, error) {
|
|
return agentcore.AgentToolResult{Content: agentcore.ContentList{agentcore.NewTextContent("bye")}, Terminate: &term}, nil
|
|
},
|
|
}
|
|
cfg := newRunCfg(scriptedStream([]agentcore.AssistantMessage{
|
|
oneToolAssistant("c1", "quit"),
|
|
{RoleField: agentcore.RoleAssistant, StopReason: agentcore.StopReasonEndTurn}, // should never be reached
|
|
}), termTool)
|
|
agentCtx := &agentcore.AgentContext{Messages: agentcore.MessageList{agentcore.UserMessage{RoleField: agentcore.RoleUser}}}
|
|
|
|
kinds, _ := collectStream(t, agentLoop(context.Background(), agentCtx, cfg))
|
|
if countKind(kinds, agentcore.EventTurnStart) != 1 {
|
|
t.Errorf("terminate must end the run after one turn, got %v", kinds)
|
|
}
|
|
}
|
|
|
|
func assertEventKinds(t *testing.T, got, want []string) {
|
|
t.Helper()
|
|
if len(got) != len(want) {
|
|
t.Fatalf("event kinds = %v, want %v", got, want)
|
|
}
|
|
for i := range want {
|
|
if got[i] != want[i] {
|
|
t.Fatalf("event[%d] = %q, want %q (full %v)", i, got[i], want[i], got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func countKind(kinds []string, want string) int {
|
|
n := 0
|
|
for _, k := range kinds {
|
|
if k == want {
|
|
n++
|
|
}
|
|
}
|
|
return n
|
|
}
|