test: 添加各模块单元测试(config/llm/service/handler)
This commit is contained in:
499
backend/internal/service/agent_test.go
Normal file
499
backend/internal/service/agent_test.go
Normal file
@@ -0,0 +1,499 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user