package service import ( "context" "fmt" "strings" "testing" "ai-agent-scaffold-go/internal/model" "github.com/stretchr/testify/assert" ) // ============================================================ // Stub 实现 // ============================================================ // stubChatModel 实现 ChatModelWithTools 接口 type stubChatModel struct { generateReply model.ChatReply generateErr error streamDelta string streamErr error toolResult string toolErr error } func (m *stubChatModel) Generate(ctx context.Context, msgs []model.ChatMessage) (model.ChatReply, error) { return m.generateReply, m.generateErr } func (m *stubChatModel) Stream(ctx context.Context, msgs []model.ChatMessage) (<-chan model.ChatStreamEvent, <-chan error) { events := make(chan model.ChatStreamEvent, 4) errs := make(chan error, 1) go func() { defer close(events) defer close(errs) if m.streamErr != nil { errs <- m.streamErr return } // 模拟逐字输出 for _, ch := range m.streamDelta { events <- model.ChatStreamEvent{Delta: string(ch)} } events <- model.ChatStreamEvent{Done: true} }() return events, errs } func (m *stubChatModel) CallTool(ctx context.Context, name, arguments string) (string, error) { return m.toolResult, m.toolErr } // stubToolCallModel 每次都返回工具调用的 stub type stubToolCallModel struct { callCount int maxCalls int // 达到此次数后返回文本 } func (m *stubToolCallModel) Generate(ctx context.Context, msgs []model.ChatMessage) (model.ChatReply, error) { m.callCount++ if m.callCount > m.maxCalls { return model.ChatReply{Content: "done"}, nil } return model.ChatReply{ToolCalls: []model.ChatToolCall{ {ID: fmt.Sprintf("call_%d", m.callCount), Name: "tool1", Arguments: `{"query":"test"}`}, }}, nil } func (m *stubToolCallModel) Stream(ctx context.Context, msgs []model.ChatMessage) (<-chan model.ChatStreamEvent, <-chan error) { events := make(chan model.ChatStreamEvent, 4) errs := make(chan error, 1) m.callCount++ if m.callCount > m.maxCalls { events <- model.ChatStreamEvent{Delta: "done", Done: true} } else { events <- model.ChatStreamEvent{ToolCalls: []model.ChatToolCall{ {ID: fmt.Sprintf("call_%d", m.callCount), Name: "tool1", Arguments: `{"query":"test"}`}, }} } close(events) close(errs) return events, errs } func (m *stubToolCallModel) CallTool(ctx context.Context, name, arguments string) (string, error) { return "tool-result", nil } // ============================================================ // 辅助函数 // ============================================================ func newTestLLMAgent(cm ChatModelWithTools) *LLMAgent { return NewLLMAgent("test-agent", "you are a test agent", "test desc", "", cm) } func testContent(msg string) model.ChatContent { return model.ChatContent{Texts: []model.TextPart{{Message: msg}}} } // ============================================================ // LLMAgent 测试 // ============================================================ func TestLLMAgent_Run_ReturnsContent(t *testing.T) { cm := &stubChatModel{generateReply: model.ChatReply{Content: "hello"}} agent := newTestLLMAgent(cm) result, err := agent.Run(context.Background(), testContent("hi")) assert.NoError(t, err) assert.Equal(t, "hello", result) } func TestLLMAgent_Run_ToolCallLoop_TwoRounds(t *testing.T) { cm := &stubToolCallModel{maxCalls: 1} agent := newTestLLMAgent(cm) result, err := agent.Run(context.Background(), testContent("hi")) assert.NoError(t, err) assert.Equal(t, "done", result) assert.Equal(t, 2, cm.callCount) } func TestLLMAgent_Run_ExceedsIterationLimit_ReturnsError(t *testing.T) { cm := &stubToolCallModel{maxCalls: 100} // 永远返回工具调用 agent := newTestLLMAgent(cm) _, err := agent.Run(context.Background(), testContent("hi")) assert.Error(t, err) assert.Contains(t, err.Error(), "exceeded tool-call iteration limit") } func TestLLMAgent_Run_GenerateError_ReturnsError(t *testing.T) { cm := &stubChatModel{generateErr: fmt.Errorf("llm down")} agent := newTestLLMAgent(cm) _, err := agent.Run(context.Background(), testContent("hi")) assert.Error(t, err) assert.Contains(t, err.Error(), "llm down") } func TestLLMAgent_Run_ToolCallError_ReturnsError(t *testing.T) { cm := &stubChatModel{ generateReply: model.ChatReply{ToolCalls: []model.ChatToolCall{ {ID: "c1", Name: "bad-tool", Arguments: "{}"}, }}, toolErr: fmt.Errorf("tool failed"), } agent := newTestLLMAgent(cm) _, err := agent.Run(context.Background(), testContent("hi")) assert.Error(t, err) assert.Contains(t, err.Error(), "tool failed") } func TestLLMAgent_Stream_ReturnsChunks(t *testing.T) { cm := &stubChatModel{streamDelta: "hello"} agent := newTestLLMAgent(cm) out := make(chan string, 10) err := agent.Stream(context.Background(), testContent("hi"), out) assert.NoError(t, err) close(out) var texts []string for s := range out { texts = append(texts, s) } assert.Equal(t, []string{"h", "e", "l", "l", "o"}, texts) } func TestLLMAgent_Stream_Error_ReturnsError(t *testing.T) { cm := &stubChatModel{streamErr: fmt.Errorf("stream failed")} agent := newTestLLMAgent(cm) out := make(chan string, 10) err := agent.Stream(context.Background(), testContent("hi"), out) assert.Error(t, err) } func TestLLMAgent_Name_ReturnsName(t *testing.T) { cm := &stubChatModel{} agent := NewLLMAgent("my-agent", "inst", "desc", "key", cm) assert.Equal(t, "my-agent", agent.Name()) assert.Equal(t, "key", agent.OutputKey()) } // ============================================================ // SequentialAgent 测试 // ============================================================ func TestSequentialAgent_Run_ExecutesInOrder(t *testing.T) { cm1 := &stubChatModel{generateReply: model.ChatReply{Content: "first"}} cm2 := &stubChatModel{generateReply: model.ChatReply{Content: "second"}} sub1 := NewLLMAgent("a1", "inst1", "", "out1", cm1) sub2 := NewLLMAgent("a2", "inst2", "", "", cm2) seq := NewSequentialAgent("seq", "desc", []model.Agent{sub1, sub2}) result, err := seq.Run(context.Background(), testContent("hi")) assert.NoError(t, err) assert.Equal(t, "second", result) // 返回最后一个的结果 } func TestSequentialAgent_Run_PassesOutputKey(t *testing.T) { // sub1 输出 "first",存入 vars["out1"] // sub2 的 instruction 包含 {out1},应被替换 cm1 := &stubChatModel{generateReply: model.ChatReply{Content: "first"}} cm2 := &stubChatModel{generateReply: model.ChatReply{Content: "got-first"}} sub1 := NewLLMAgent("a1", "inst1", "", "out1", cm1) sub2 := NewLLMAgent("a2", "instruction with {out1}", "", "", cm2) seq := NewSequentialAgent("seq", "desc", []model.Agent{sub1, sub2}) result, err := seq.Run(context.Background(), testContent("hi")) assert.NoError(t, err) assert.Equal(t, "got-first", result) } func TestSequentialAgent_Run_SubAgentError_StopsExecution(t *testing.T) { cm1 := &stubChatModel{generateErr: fmt.Errorf("fail")} cm2 := &stubChatModel{generateReply: model.ChatReply{Content: "second"}} sub1 := NewLLMAgent("a1", "inst1", "", "", cm1) sub2 := NewLLMAgent("a2", "inst2", "", "", cm2) seq := NewSequentialAgent("seq", "desc", []model.Agent{sub1, sub2}) _, err := seq.Run(context.Background(), testContent("hi")) assert.Error(t, err) } func TestSequentialAgent_Stream_LastAgentStreams(t *testing.T) { cm1 := &stubChatModel{generateReply: model.ChatReply{Content: "first"}} cm2 := &stubChatModel{streamDelta: "stream"} sub1 := NewLLMAgent("a1", "inst1", "", "out1", cm1) sub2 := NewLLMAgent("a2", "inst2", "", "", cm2) seq := NewSequentialAgent("seq", "desc", []model.Agent{sub1, sub2}) out := make(chan string, 10) err := seq.Stream(context.Background(), testContent("hi"), out) close(out) assert.NoError(t, err) var texts []string for s := range out { texts = append(texts, s) } assert.Equal(t, []string{"s", "t", "r", "e", "a", "m"}, texts) } // ============================================================ // ParallelAgent 测试 // ============================================================ func TestParallelAgent_Run_ConcatenatesResults(t *testing.T) { cm1 := &stubChatModel{generateReply: model.ChatReply{Content: "aaa"}} cm2 := &stubChatModel{generateReply: model.ChatReply{Content: "bbb"}} sub1 := NewLLMAgent("a1", "inst1", "", "", cm1) sub2 := NewLLMAgent("a2", "inst2", "", "", cm2) par := NewParallelAgent("par", "desc", []model.Agent{sub1, sub2}) result, err := par.Run(context.Background(), testContent("hi")) assert.NoError(t, err) assert.Contains(t, result, "[a1] aaa") assert.Contains(t, result, "[a2] bbb") } func TestParallelAgent_Run_SubAgentError_ReturnsError(t *testing.T) { cm1 := &stubChatModel{generateReply: model.ChatReply{Content: "ok"}} cm2 := &stubChatModel{generateErr: fmt.Errorf("fail")} sub1 := NewLLMAgent("a1", "inst1", "", "", cm1) sub2 := NewLLMAgent("a2", "inst2", "", "", cm2) par := NewParallelAgent("par", "desc", []model.Agent{sub1, sub2}) _, err := par.Run(context.Background(), testContent("hi")) assert.Error(t, err) } func TestParallelAgent_Run_ConcurrentExecution(t *testing.T) { // 验证并发执行:两个 agent 都被调用 cm1 := &stubChatModel{generateReply: model.ChatReply{Content: "first"}} cm2 := &stubChatModel{generateReply: model.ChatReply{Content: "second"}} sub1 := NewLLMAgent("a1", "inst1", "", "", cm1) sub2 := NewLLMAgent("a2", "inst2", "", "", cm2) par := NewParallelAgent("par", "desc", []model.Agent{sub1, sub2}) result, err := par.Run(context.Background(), testContent("hi")) assert.NoError(t, err) // 结果应包含两个 agent 的输出 assert.True(t, strings.Contains(result, "first")) assert.True(t, strings.Contains(result, "second")) } func TestParallelAgent_Stream_OutputsConcurrently(t *testing.T) { // 使用单次输出的 stub 避免逐字符并发竞争 cm1 := &singleShotChatModel{content: "result-a"} cm2 := &singleShotChatModel{content: "result-b"} sub1 := NewLLMAgent("a1", "inst1", "", "", cm1) sub2 := NewLLMAgent("a2", "inst2", "", "", cm2) par := NewParallelAgent("par", "desc", []model.Agent{sub1, sub2}) out := make(chan string, 20) err := par.Stream(context.Background(), testContent("hi"), out) close(out) assert.NoError(t, err) var texts []string for s := range out { texts = append(texts, s) } full := strings.Join(texts, "") assert.Contains(t, full, "[a1]") assert.Contains(t, full, "[a2]") assert.Contains(t, full, "result-a") assert.Contains(t, full, "result-b") } // ============================================================ // LoopAgent 测试 // ============================================================ func TestLoopAgent_Run_RepeatsSubAgents(t *testing.T) { cm := &stubChatModel{generateReply: model.ChatReply{Content: "tick"}} sub := NewLLMAgent("a1", "inst1", "", "", cm) loop := NewLoopAgent("loop", "desc", []model.Agent{sub}, 3) result, err := loop.Run(context.Background(), testContent("hi")) assert.NoError(t, err) assert.Contains(t, result, "[a1] tick") } func TestLoopAgent_Run_DefaultMaxIterations(t *testing.T) { cm := &stubChatModel{generateReply: model.ChatReply{Content: "ok"}} sub := NewLLMAgent("a1", "inst1", "", "", cm) loop := NewLoopAgent("loop", "desc", []model.Agent{sub}, 0) // 0 → 默认 3 assert.Equal(t, 3, loop.maxIterations) } func TestLoopAgent_Run_SubAgentError_StopsLoop(t *testing.T) { callCount := 0 errModel := &stubChatModel{} errModel.generateErr = fmt.Errorf("fail on call") // 用一个计数 stub countModel := &countingChatModel{reply: model.ChatReply{Content: "ok"}, failAfter: 2} sub := NewLLMAgent("a1", "inst1", "", "", countModel) loop := NewLoopAgent("loop", "desc", []model.Agent{sub}, 5) _, err := loop.Run(context.Background(), testContent("hi")) assert.Error(t, err) _ = callCount _ = errModel } func TestLoopAgent_Stream_ExecutesAndStreams(t *testing.T) { cm := &stubChatModel{streamDelta: "loop"} sub := NewLLMAgent("a1", "inst1", "", "", cm) loop := NewLoopAgent("loop", "desc", []model.Agent{sub}, 2) out := make(chan string, 20) err := loop.Stream(context.Background(), testContent("hi"), out) close(out) assert.NoError(t, err) var texts []string for s := range out { texts = append(texts, s) } full := strings.Join(texts, "") assert.Contains(t, full, "[a1]") } // ============================================================ // 工具函数测试 // ============================================================ func TestCloneVars_CreatesIndependentCopy(t *testing.T) { orig := map[string]string{"a": "1", "b": "2"} cloned := cloneVars(orig) cloned["c"] = "3" assert.NotContains(t, orig, "c") } func TestApplyVars_ReplacesPlaceholders(t *testing.T) { vars := map[string]string{"name": "world", "greeting": "hello"} result := applyVars("{greeting} {name}!", vars) assert.Equal(t, "hello world!", result) } func TestApplyVars_EmptyTemplate_ReturnsEmpty(t *testing.T) { result := applyVars("", map[string]string{"a": "1"}) assert.Equal(t, "", result) } func TestApplyVars_NoVars_ReturnsOriginal(t *testing.T) { result := applyVars("hello {name}", nil) assert.Equal(t, "hello {name}", result) } func TestInitialMessages_WithInstruction(t *testing.T) { msgs := initialMessages("system instruction", "user text") assert.Len(t, msgs, 2) assert.Equal(t, model.ChatRoleSystem, msgs[0].Role) assert.Equal(t, "system instruction", msgs[0].Content) assert.Equal(t, model.ChatRoleUser, msgs[1].Role) } func TestInitialMessages_EmptyInstruction_SkipsSystem(t *testing.T) { msgs := initialMessages("", "user text") assert.Len(t, msgs, 1) assert.Equal(t, model.ChatRoleUser, msgs[0].Role) } func TestFirstText_WithContent(t *testing.T) { content := model.ChatContent{Texts: []model.TextPart{{Message: "hi"}, {Message: "bye"}}} assert.Equal(t, "hi", firstText(content)) } func TestFirstText_EmptyContent(t *testing.T) { assert.Equal(t, "", firstText(model.ChatContent{})) } // ============================================================ // 辅助 stub // ============================================================ // singleShotChatModel 一次性输出完整内容的 stub(适合并发测试) type singleShotChatModel struct { content string } func (m *singleShotChatModel) Generate(ctx context.Context, msgs []model.ChatMessage) (model.ChatReply, error) { return model.ChatReply{Content: m.content}, nil } func (m *singleShotChatModel) Stream(ctx context.Context, msgs []model.ChatMessage) (<-chan model.ChatStreamEvent, <-chan error) { events := make(chan model.ChatStreamEvent, 2) errs := make(chan error, 1) go func() { events <- model.ChatStreamEvent{Delta: m.content, Done: true} close(events) close(errs) }() return events, errs } func (m *singleShotChatModel) CallTool(ctx context.Context, name, arguments string) (string, error) { return "ok", nil } // countingChatModel 记录调用次数,超过 failAfter 后返回错误 type countingChatModel struct { reply model.ChatReply failAfter int calls int } func (m *countingChatModel) Generate(ctx context.Context, msgs []model.ChatMessage) (model.ChatReply, error) { m.calls++ if m.calls > m.failAfter { return model.ChatReply{}, fmt.Errorf("fail at call %d", m.calls) } return m.reply, nil } func (m *countingChatModel) Stream(ctx context.Context, msgs []model.ChatMessage) (<-chan model.ChatStreamEvent, <-chan error) { events := make(chan model.ChatStreamEvent, 4) errs := make(chan error, 1) m.calls++ if m.calls > m.failAfter { errs <- fmt.Errorf("fail at call %d", m.calls) } else { events <- model.ChatStreamEvent{Delta: m.reply.Content, Done: true} } close(events) close(errs) return events, errs } func (m *countingChatModel) CallTool(ctx context.Context, name, arguments string) (string, error) { return "ok", nil }