Files
xk-ai-agent/test/agent_test.go
2026-08-14 21:50:48 +08:00

511 lines
14 KiB
Go
Raw Permalink 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 test
import (
"context"
"encoding/json"
"testing"
"tcm-agent/internal/agent"
"tcm-agent/internal/config"
"tcm-agent/internal/llm"
"tcm-agent/internal/rule"
)
// ========================================================================
// 测试文件
// ========================================================================
// 覆盖:
// - 模型工厂注册和创建
// - 模型路由(场景→模型)
// - 降级链
// - Agent 会话管理
// - 规则引擎(十八反/十九畏/孕妇/过敏)— 使用 rule 包独立类型
// - 结构化 JSON 序列化
// - 配置加载(多模型)
// ========================================================================
// ---------- 辅助函数 ----------
// newTestConfig 创建测试用配置(多模型)
func newTestConfig() *config.Config {
return &config.Config{
Server: config.ServerConfig{Port: "8080"},
MaxKB: config.MaxKBConfig{
BaseURL: "http://mock", APIKey: "test", AppID: "test",
},
LLM: config.LLMConfig{
DefaultProvider: "mock-primary",
Models: map[string]config.LLMConfigEx{
"mock-primary": {
Provider: "mock", APIKey: "test",
BaseURL: "http://mock", Model: "mock-model-1",
},
"mock-backup": {
Provider: "mock", APIKey: "test",
BaseURL: "http://mock", Model: "mock-model-2",
},
"deepseek": {
Provider: "deepseek", APIKey: "test",
BaseURL: "http://mock", Model: "deepseek-chat",
},
},
Routes: map[string]string{
"emr-generator": "mock-primary",
"prescription": "mock-primary",
"knowledge-qa": "deepseek",
},
FallbackChains: map[string][]string{
"emr-generator": {"mock-primary", "mock-backup"},
},
},
Agent: config.AgentConfig{MaxIterations: 5, Timeout: 30},
}
}
// newTestRouter 创建测试用模型路由
func newTestRouter(cfg *config.Config) (*llm.ModelRouter, *llm.FallbackChain) {
configs := make(map[string]*config.LLMConfigEx)
for name, m := range cfg.LLM.Models {
m2 := m
configs[name] = &m2
}
factory := llm.NewProviderFactory(configs)
router := llm.NewModelRouter(factory, cfg.LLM.Routes, cfg.LLM.DefaultProvider)
fallback := llm.NewFallbackChain(router, cfg.LLM.FallbackChains)
return router, fallback
}
// ---------- 模型工厂测试 ----------
// TestProviderFactory 测试工厂注册和创建
func TestProviderFactory(t *testing.T) {
cfg := newTestConfig()
configs := make(map[string]*config.LLMConfigEx)
for name, m := range cfg.LLM.Models {
m2 := m
configs[name] = &m2
}
factory := llm.NewProviderFactory(configs)
// 验证内置供应商已注册
providers := factory.ListProviders()
t.Logf("已注册供应商: %v", providers)
expectedProviders := []string{"deepseek", "openai", "azure", "ollama", "qwen", "mock"}
for _, ep := range expectedProviders {
found := false
for _, p := range providers {
if p == ep {
found = true
break
}
}
if !found {
t.Errorf("供应商 %s 未注册", ep)
}
}
// 验证创建客户端
client, err := factory.Create("deepseek")
if err != nil {
t.Fatalf("创建 DeepSeek 客户端失败: %v", err)
}
if client.Provider() != "deepseek" {
t.Errorf("期望 provider=deepseek, 实际=%s", client.Provider())
}
if !client.Supports(llm.CapFunctionCalling) {
t.Error("DeepSeek 应支持 function_calling")
}
t.Logf("✅ DeepSeek 客户端创建成功: %s", client.Name())
client.Close()
// 验证未知供应商报错
_, err = factory.Create("unknown-provider")
if err == nil {
t.Error("未知供应商应返回错误")
}
t.Logf("✅ 未知供应商正确报错: %v", err)
}
// TestCustomProvider 测试注册自定义供应商
func TestCustomProvider(t *testing.T) {
configs := make(map[string]*config.LLMConfigEx)
configs["custom"] = &config.LLMConfigEx{
Provider: "custom", APIKey: "test", Model: "my-model",
}
factory := llm.NewProviderFactory(configs)
// 注册自定义创建函数
factory.Register("custom", func(cfg *config.LLMConfigEx) (llm.LLMClient, error) {
return &mockClientForTest{name: cfg.Model, provider: "custom"}, nil
})
client, err := factory.Create("custom")
if err != nil {
t.Fatalf("创建自定义客户端失败: %v", err)
}
if client.Name() != "my-model" {
t.Errorf("期望 name=my-model, 实际=%s", client.Name())
}
t.Logf("✅ 自定义供应商注册成功")
client.Close()
}
// ---------- 模型路由测试 ----------
// TestModelRouter 测试场景路由
func TestModelRouter(t *testing.T) {
cfg := newTestConfig()
router, _ := newTestRouter(cfg)
// 验证路由表
routes := router.ListRoutes()
for scene, provider := range routes {
t.Logf("路由: %-20s → %s", scene, provider)
}
// 验证获取客户端
client, err := router.Get("emr-generator")
if err != nil {
t.Fatalf("路由获取失败: %v", err)
}
if client.Provider() != "mock" {
t.Errorf("期望 mock provider, 实际=%s", client.Provider())
}
t.Logf("✅ 路由 emr-generator → %s", client.Name())
// 验证未知场景使用默认
client2, err := router.Get("unknown-scene")
if err != nil {
t.Fatalf("默认路由失败: %v", err)
}
t.Logf("✅ 未知场景使用默认 → %s", client2.Name())
router.Close()
}
// TestDynamicRoute 测试动态注册路由
func TestDynamicRoute(t *testing.T) {
cfg := newTestConfig()
router, _ := newTestRouter(cfg)
router.RegisterRoute("new-scene", "deepseek")
routes := router.ListRoutes()
if routes["new-scene"] != "deepseek" {
t.Error("动态路由注册失败")
}
t.Logf("✅ 动态路由注册成功: new-scene → deepseek")
router.Close()
}
// ---------- 降级链测试 ----------
// TestFallbackChain 测试降级
func TestFallbackChain(t *testing.T) {
cfg := newTestConfig()
router, fallback := newTestRouter(cfg)
// 降级链存在
resp, err := fallback.ChatWithFallback(context.Background(), "emr-generator", nil, nil)
if err != nil {
t.Logf("降级链执行(预期可能有错误,因为 mock 模式): %v", err)
} else {
t.Logf("✅ 降级链成功: %s", resp.Content)
}
router.Close()
}
// ---------- Agent 会话管理测试 ----------
// TestAgentSessionLifecycle 测试完整会话生命周期
func TestAgentSessionLifecycle(t *testing.T) {
cfg := newTestConfig()
router, fallback := newTestRouter(cfg)
runner := agent.InitRunner(router, fallback, cfg)
// 创建会话
session := runner.CreateSession("test-doctor", "emr-generator")
if session.ID == "" {
t.Fatal("会话 ID 不能为空")
}
if session.UserID != "test-doctor" {
t.Errorf("期望 user_id=test-doctor, 实际=%s", session.UserID)
}
if session.Scene != "emr-generator" {
t.Errorf("期望 scene=emr-generator, 实际=%s", session.Scene)
}
t.Logf("✅ 会话创建成功: ID=%s, Scene=%s", session.ID, session.Scene)
// 获取会话
got, ok := runner.GetSession(session.ID)
if !ok {
t.Fatal("无法获取已创建的会话")
}
if got.Status != "running" {
t.Errorf("期望 status=running, 实际=%s", got.Status)
}
// 删除会话
runner.DeleteSession(session.ID)
_, ok = runner.GetSession(session.ID)
if ok {
t.Error("删除后会话应不存在")
}
t.Logf("✅ 会话删除成功")
router.Close()
}
// TestAgentToolRegistration 测试工具注册
func TestAgentToolRegistration(t *testing.T) {
cfg := newTestConfig()
router, fallback := newTestRouter(cfg)
runner := agent.InitRunner(router, fallback, cfg)
tools := runner.ListTools()
t.Logf("已注册工具: %v", tools)
expectedTools := []string{"maxkb_retrieve", "his_query", "pharmacopoeia_query", "rule_check"}
for _, et := range expectedTools {
found := false
for _, toolName := range tools {
if toolName == et {
found = true
break
}
}
if !found {
t.Errorf("工具 %s 未注册", et)
}
}
t.Logf("✅ 默认工具全部注册成功 (%d 个)", len(tools))
router.Close()
}
// ---------- 规则引擎测试(使用 rule 包独立类型)----------
// TestEighteenAnti 专项测试:十八反
func TestEighteenAnti(t *testing.T) {
tests := []struct {
name string
prescription string
shouldBlock bool
}{
{"甘草反甘遂", "甘草6g 甘遂3g", true},
{"甘草反大戟", "甘草6g 大戟3g", true},
{"乌头反贝母", "乌头5g 贝母6g", true},
{"藜芦反人参", "藜芦3g 人参10g", true},
{"正常处方", "桂枝10g 白芍10g 甘草6g", false},
}
validator := rule.NewPrescriptionValidator()
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
patient := &rule.PatientInfo{}
warnings, blocked := validator.Validate(tt.prescription, patient)
if blocked != tt.shouldBlock {
t.Errorf("[%s] 期望 blocked=%v, 实际=%v, warnings=%v",
tt.name, tt.shouldBlock, blocked, warnings)
} else {
t.Logf("✓ [%s] blocked=%v, warnings=%v", tt.name, blocked, warnings)
}
})
}
}
// TestNineteenFear 测试:十九畏
func TestNineteenFear(t *testing.T) {
tests := []struct {
name string
prescription string
}{
{"硫黄畏朴硝", "硫黄3g 朴硝6g"},
{"丁香畏郁金", "丁香3g 郁金10g"},
{"巴豆畏牵牛", "巴豆3g 牵牛10g"},
}
validator := rule.NewPrescriptionValidator()
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
patient := &rule.PatientInfo{}
warnings, blocked := validator.Validate(tt.prescription, patient)
if !blocked {
t.Errorf("[%s] 应被拦截但未拦截, warnings=%v", tt.name, warnings)
} else {
t.Logf("✓ [%s] 正确拦截: %v", tt.name, warnings)
}
})
}
}
// TestPregnancyRisk 测试:孕妇禁忌(使用 rule.PatientInfo
func TestPregnancyRisk(t *testing.T) {
tests := []struct {
name string
prescription string
shouldBlock bool
}{
{"孕妇禁用桃仁", "桃仁9g 红花6g", true},
{"孕妇慎用大黄", "大黄6g 芒硝3g", false}, // 慎用不拦截
{"正常处方", "桂枝10g 白芍10g", false},
}
validator := rule.NewPrescriptionValidator()
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
patient := &rule.PatientInfo{IsPregnant: true}
warnings, blocked := validator.Validate(tt.prescription, patient)
if blocked != tt.shouldBlock {
t.Errorf("[%s] 期望 blocked=%v, 实际=%v, warnings=%v",
tt.name, tt.shouldBlock, blocked, warnings)
} else {
t.Logf("✓ [%s] blocked=%v, warnings=%v", tt.name, blocked, warnings)
}
})
}
}
// TestAllergyCheck 测试:过敏史冲突
func TestAllergyCheck(t *testing.T) {
validator := rule.NewPrescriptionValidator()
patient := &rule.PatientInfo{Allergies: []string{"麻黄"}}
warnings, blocked := validator.Validate("麻黄9g 杏仁6g", patient)
if !blocked {
t.Error("过敏冲突应被拦截")
}
if len(warnings) == 0 {
t.Error("应产生警告")
}
t.Logf("✅ 过敏检查正确拦截: %v", warnings)
}
// TestEMRQualityChecker 测试病历质控(使用 rule 包)
func TestEMRQualityChecker(t *testing.T) {
checker := rule.NewEMRQualityChecker()
// 完整病历应通过
goodEMR := `主诉反复头晕3个月
现病史患者3个月前出现头晕
舌象:舌质暗红,苔白腻
脉象:脉弦滑
诊断:痰湿中阻证`
issues := checker.Check(goodEMR)
if len(issues) > 0 {
t.Errorf("完整病历不应有问题: %v", issues)
}
t.Logf("✅ 完整病历质控通过")
// 缺失字段应被检出
badEMR := "患者头晕,开了点药。"
issues = checker.Check(badEMR)
if len(issues) == 0 {
t.Error("缺失字段应被检出")
}
t.Logf("✅ 缺失字段正确检出: %v", issues)
}
// ---------- JSON 序列化测试 ----------
// TestStructToJSON 测试 API 响应格式
func TestStructToJSON(t *testing.T) {
resp := &agent.PrescriptionResponse{
SessionID: "sess-123",
Draft: "桂枝汤桂枝10g...",
Warnings: []string{},
Blocked: false,
Status: "success",
Prescription: &agent.Prescription{
FormulaName: "桂枝汤",
Herbs: []agent.Herb{
{Name: "桂枝", Dose: 10, Unit: "g"},
{Name: "白芍", Dose: 10, Unit: "g"},
},
Instructions: "水煎服,日一剂",
Duration: 7,
},
}
data, err := json.MarshalIndent(resp, "", " ")
if err != nil {
t.Fatalf("JSON 序列化失败: %v", err)
}
t.Logf("API 响应 JSON 示例:\n%s", string(data))
// 验证关键字段
var parsed map[string]any
json.Unmarshal(data, &parsed)
if parsed["status"] != "success" {
t.Error("status 字段不正确")
}
if parsed["blocked"] != false {
t.Error("blocked 字段不正确")
}
}
// ---------- 配置加载测试 ----------
// TestConfigLoading 测试多模型配置加载
func TestConfigLoading(t *testing.T) {
cfg := newTestConfig()
// 验证模型池
if len(cfg.LLM.Models) != 3 {
t.Errorf("期望 3 个模型,实际 %d 个", len(cfg.LLM.Models))
}
// 验证默认供应商
if cfg.LLM.DefaultProvider != "mock-primary" {
t.Errorf("期望默认=mock-primary, 实际=%s", cfg.LLM.DefaultProvider)
}
// 验证路由
if cfg.LLM.Routes["emr-generator"] != "mock-primary" {
t.Error("emr-generator 路由不正确")
}
// 验证降级链
chain, ok := cfg.LLM.FallbackChains["emr-generator"]
if !ok || len(chain) != 2 {
t.Error("降级链配置不正确")
}
t.Logf("✅ 配置加载正确 | 模型数:%d | 路由数:%d",
len(cfg.LLM.Models), len(cfg.LLM.Routes))
}
// ========================================================================
// 测试辅助
// ========================================================================
// mockClientForTest 测试用自定义客户端
type mockClientForTest struct {
name string
provider string
}
func (m *mockClientForTest) Chat(ctx context.Context, messages []agent.Message, tools []agent.Tool) (*agent.Message, error) {
return &agent.Message{Role: "assistant", Content: "mock reply", Timestamp: 123}, nil
}
func (m *mockClientForTest) StreamChat(ctx context.Context, messages []agent.Message, tools []agent.Tool) (<-chan string, error) {
ch := make(chan string, 1)
ch <- "mock"
close(ch)
return ch, nil
}
func (m *mockClientForTest) Embed(ctx context.Context, texts []string) ([][]float32, error) {
return make([][]float32, len(texts)), nil
}
func (m *mockClientForTest) Name() string { return m.name }
func (m *mockClientForTest) Provider() string { return m.provider }
func (m *mockClientForTest) Supports(cap string) bool { return true }
func (m *mockClientForTest) Close() error { return nil }