Files
GoLoom/backend/internal/service/agent_test.go

500 lines
16 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}