初始化
This commit is contained in:
510
test/agent_test.go
Normal file
510
test/agent_test.go
Normal file
@@ -0,0 +1,510 @@
|
||||
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 }
|
||||
Reference in New Issue
Block a user