201 lines
5.6 KiB
Go
201 lines
5.6 KiB
Go
package service
|
||
|
||
import (
|
||
"fmt"
|
||
"testing"
|
||
|
||
"ai-agent-scaffold-go/internal/model"
|
||
"ai-agent-scaffold-go/pkg/types"
|
||
"github.com/stretchr/testify/assert"
|
||
)
|
||
|
||
// ============================================================
|
||
// Stub Runner
|
||
// ============================================================
|
||
|
||
type stubRunner struct {
|
||
sessionID string
|
||
runResult []string
|
||
runErr error
|
||
}
|
||
|
||
func (r *stubRunner) CreateSession(userID string) (string, error) {
|
||
if r.sessionID != "" {
|
||
return r.sessionID, nil
|
||
}
|
||
return "sess:" + userID + ":1", nil
|
||
}
|
||
|
||
func (r *stubRunner) Run(userID, sessionID string, content model.ChatContent) ([]string, error) {
|
||
return r.runResult, r.runErr
|
||
}
|
||
|
||
func (r *stubRunner) Stream(userID, sessionID string, content model.ChatContent) (<-chan string, <-chan error) {
|
||
outputs := make(chan string, 4)
|
||
errs := make(chan error, 1)
|
||
go func() {
|
||
defer close(outputs)
|
||
defer close(errs)
|
||
if r.runErr != nil {
|
||
errs <- r.runErr
|
||
return
|
||
}
|
||
for _, s := range r.runResult {
|
||
outputs <- s
|
||
}
|
||
}()
|
||
return outputs, errs
|
||
}
|
||
|
||
// ============================================================
|
||
// ChatService 测试
|
||
// ============================================================
|
||
|
||
func newTestChatService() (*ChatService, *model.InMemoryAgentRegistry, *model.InMemorySessionStore) {
|
||
registry := model.NewInMemoryAgentRegistry()
|
||
sessions := model.NewInMemorySessionStore()
|
||
svc := NewChatService(registry, sessions)
|
||
return svc, registry, sessions
|
||
}
|
||
|
||
func TestChatService_QueryAgentConfigList_ReturnsSorted(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
|
||
registry.Register(model.RegisteredAgent{AgentID: "2", AgentName: "b", AgentDesc: "desc b"})
|
||
registry.Register(model.RegisteredAgent{AgentID: "1", AgentName: "a", AgentDesc: "desc a"})
|
||
|
||
agents := svc.QueryAgentConfigList()
|
||
assert.Len(t, agents, 2)
|
||
assert.Equal(t, "1", agents[0].AgentID)
|
||
assert.Equal(t, "2", agents[1].AgentID)
|
||
}
|
||
|
||
func TestChatService_QueryAgentConfigList_Empty(t *testing.T) {
|
||
svc, _, _ := newTestChatService()
|
||
agents := svc.QueryAgentConfigList()
|
||
assert.Len(t, agents, 0)
|
||
}
|
||
|
||
func TestChatService_CreateSession_NewSession(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
registry.Register(model.RegisteredAgent{
|
||
AgentID: "1",
|
||
Runner: &stubRunner{sessionID: "sess:u1:1"},
|
||
})
|
||
|
||
sessionID, err := svc.CreateSession("1", "user1")
|
||
assert.NoError(t, err)
|
||
assert.Equal(t, "sess:u1:1", sessionID)
|
||
}
|
||
|
||
func TestChatService_CreateSession_ReusesExisting(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
registry.Register(model.RegisteredAgent{
|
||
AgentID: "1",
|
||
Runner: &stubRunner{sessionID: "sess:u1:1"},
|
||
})
|
||
|
||
id1, _ := svc.CreateSession("1", "user1")
|
||
id2, _ := svc.CreateSession("1", "user1")
|
||
assert.Equal(t, id1, id2)
|
||
}
|
||
|
||
func TestChatService_CreateSession_AgentNotFound_ReturnsError(t *testing.T) {
|
||
svc, _, _ := newTestChatService()
|
||
|
||
_, err := svc.CreateSession("nonexistent", "user1")
|
||
assert.Error(t, err)
|
||
|
||
var appErr *types.AppError
|
||
assert.ErrorAs(t, err, &appErr)
|
||
assert.Equal(t, types.CodeAgentNotFound, appErr.Code)
|
||
}
|
||
|
||
func TestChatService_HandleMessage_AgentNotFound_ReturnsError(t *testing.T) {
|
||
svc, _, _ := newTestChatService()
|
||
|
||
_, err := svc.HandleMessage("nonexistent", "user1", "", "hello")
|
||
assert.Error(t, err)
|
||
}
|
||
|
||
func TestChatService_HandleMessage_NormalCall(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
registry.Register(model.RegisteredAgent{
|
||
AgentID: "1",
|
||
Runner: &stubRunner{sessionID: "s1", runResult: []string{"reply"}},
|
||
})
|
||
|
||
outputs, err := svc.HandleMessage("1", "user1", "s1", "hello")
|
||
assert.NoError(t, err)
|
||
assert.Equal(t, []string{"reply"}, outputs)
|
||
}
|
||
|
||
func TestChatService_HandleMessage_EmptyMessage_StillPassesThrough(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
registry.Register(model.RegisteredAgent{
|
||
AgentID: "1",
|
||
Runner: &stubRunner{sessionID: "s1", runResult: []string{}},
|
||
})
|
||
|
||
// 空消息仍会创建 TextPart,当前实现不校验空消息内容
|
||
outputs, err := svc.HandleMessage("1", "user1", "s1", "")
|
||
assert.NoError(t, err)
|
||
assert.Empty(t, outputs)
|
||
}
|
||
|
||
func TestChatService_HandleMessage_RunnerError_Propagates(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
registry.Register(model.RegisteredAgent{
|
||
AgentID: "1",
|
||
Runner: &stubRunner{sessionID: "s1", runErr: fmt.Errorf("run failed")},
|
||
})
|
||
|
||
_, err := svc.HandleMessage("1", "user1", "s1", "hello")
|
||
assert.Error(t, err)
|
||
assert.Contains(t, err.Error(), "run failed")
|
||
}
|
||
|
||
func TestChatService_HandleMessageStream_AgentNotFound_ReturnsError(t *testing.T) {
|
||
svc, _, _ := newTestChatService()
|
||
|
||
outputs, errs := svc.HandleMessageStream("nonexistent", "user1", "", "hello")
|
||
// 消费通道
|
||
var streamErr error
|
||
for range outputs {
|
||
}
|
||
for err := range errs {
|
||
streamErr = err
|
||
}
|
||
assert.Error(t, streamErr)
|
||
}
|
||
|
||
func TestChatService_HandleMessageStream_NormalCall(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
registry.Register(model.RegisteredAgent{
|
||
AgentID: "1",
|
||
Runner: &stubRunner{sessionID: "s1", runResult: []string{"chunk1", "chunk2"}},
|
||
})
|
||
|
||
outputs, errs := svc.HandleMessageStream("1", "user1", "s1", "hello")
|
||
var texts []string
|
||
for s := range outputs {
|
||
texts = append(texts, s)
|
||
}
|
||
for range errs {
|
||
}
|
||
assert.Equal(t, []string{"chunk1", "chunk2"}, texts)
|
||
}
|
||
|
||
func TestChatService_CreateSession_AutoCreatesWhenSessionEmpty(t *testing.T) {
|
||
svc, registry, _ := newTestChatService()
|
||
registry.Register(model.RegisteredAgent{
|
||
AgentID: "1",
|
||
Runner: &stubRunner{sessionID: "auto-sess", runResult: []string{"ok"}},
|
||
})
|
||
|
||
// sessionID 为空时自动创建
|
||
outputs, err := svc.HandleMessage("1", "user1", "", "hello")
|
||
assert.NoError(t, err)
|
||
assert.Equal(t, []string{"ok"}, outputs)
|
||
}
|