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 }