初始化

This commit is contained in:
2026-08-14 21:50:48 +08:00
commit d7e382f2e7
114 changed files with 24123 additions and 0 deletions

510
test/agent_test.go Normal file
View 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 }