511 lines
14 KiB
Go
511 lines
14 KiB
Go
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 }
|