290 lines
9.8 KiB
Python
290 lines
9.8 KiB
Python
"""顾问 Agent 工具层测试"""
|
|
|
|
import pytest
|
|
from unittest.mock import Mock
|
|
from datetime import date
|
|
|
|
from app.service.agent_tools import (
|
|
create_query_holdings_tool,
|
|
create_query_fund_nav_tool,
|
|
create_manage_kyc_tool,
|
|
create_search_templates_tool,
|
|
create_compliance_check_tool,
|
|
)
|
|
|
|
|
|
class TestQueryCustomerHoldings:
|
|
"""测试查询客户持仓工具"""
|
|
|
|
def test_returns_holdings_with_customer_info(self):
|
|
"""应返回客户基本信息和持仓列表"""
|
|
mock_core_repo = Mock()
|
|
mock_core_repo.get_customer_l0.return_value = {
|
|
"customer_id": "CUST-001",
|
|
"display_name": "张三",
|
|
"risk_code": "C3",
|
|
}
|
|
mock_core_repo.list_holdings.return_value = [
|
|
{
|
|
"product_id": "PROD-001",
|
|
"product_name": "稳健增长基金A",
|
|
"shares": 1000.0,
|
|
"market_value": 15000.0,
|
|
},
|
|
{
|
|
"product_id": "PROD-002",
|
|
"product_name": "平衡配置基金B",
|
|
"shares": 500.0,
|
|
"market_value": 8000.0,
|
|
},
|
|
]
|
|
|
|
tool = create_query_holdings_tool(mock_core_repo)
|
|
result = tool.invoke({"customer_id": "CUST-001"})
|
|
|
|
assert result["customer_id"] == "CUST-001"
|
|
assert result["display_name"] == "张三"
|
|
assert result["risk_code"] == "C3"
|
|
assert len(result["holdings"]) == 2
|
|
assert result["holdings"][0]["product_name"] == "稳健增长基金A"
|
|
|
|
def test_returns_error_when_customer_not_found(self):
|
|
"""客户不存在时应返回错误信息"""
|
|
mock_core_repo = Mock()
|
|
mock_core_repo.get_customer_l0.return_value = None
|
|
|
|
tool = create_query_holdings_tool(mock_core_repo)
|
|
result = tool.invoke({"customer_id": "NOT-EXIST"})
|
|
|
|
assert "error" in result
|
|
assert "NOT-EXIST" in result["error"]
|
|
|
|
def test_handles_empty_holdings(self):
|
|
"""客户无持仓时应返回空列表"""
|
|
mock_core_repo = Mock()
|
|
mock_core_repo.get_customer_l0.return_value = {
|
|
"customer_id": "CUST-002",
|
|
"display_name": "李四",
|
|
"risk_code": "C2",
|
|
}
|
|
mock_core_repo.list_holdings.return_value = []
|
|
|
|
tool = create_query_holdings_tool(mock_core_repo)
|
|
result = tool.invoke({"customer_id": "CUST-002"})
|
|
|
|
assert result["customer_id"] == "CUST-002"
|
|
assert result["holdings"] == []
|
|
|
|
|
|
class TestQueryFundNav:
|
|
"""测试查询基金净值工具"""
|
|
|
|
def test_returns_latest_nav_with_history(self):
|
|
"""应返回最新净值和历史数据"""
|
|
mock_core_repo = Mock()
|
|
mock_core_repo.get_product.return_value = {
|
|
"product_id": "PROD-000001",
|
|
"product_name": "稳健增长基金A",
|
|
"product_type": "equity",
|
|
"min_risk_code": "C3",
|
|
}
|
|
mock_core_repo.get_latest_nav.return_value = {
|
|
"product_id": "PROD-000001",
|
|
"nav": 1.25,
|
|
"nav_date": date(2026, 9, 11),
|
|
"daily_chg_pct": 0.5,
|
|
}
|
|
mock_core_repo.list_product_nav_history.return_value = [
|
|
{"nav_date": date(2026, 9, 11), "nav": 1.25, "daily_chg_pct": 0.5},
|
|
{"nav_date": date(2026, 9, 10), "nav": 1.24, "daily_chg_pct": -0.2},
|
|
]
|
|
|
|
tool = create_query_fund_nav_tool(mock_core_repo)
|
|
result = tool.invoke({"fund_code": "000001", "history_days": 2})
|
|
|
|
assert result["fund_code"] == "PROD-000001"
|
|
assert result["fund_name"] == "稳健增长基金A"
|
|
assert result["latest_nav"] == 1.25
|
|
assert len(result["history"]) == 2
|
|
|
|
def test_normalizes_fund_code(self):
|
|
"""应自动补全 PROD- 前缀"""
|
|
mock_core_repo = Mock()
|
|
mock_core_repo.get_product.return_value = {
|
|
"product_id": "PROD-000001",
|
|
"product_name": "Test Fund",
|
|
"product_type": "equity",
|
|
"min_risk_code": "C3",
|
|
}
|
|
mock_core_repo.get_latest_nav.return_value = None
|
|
mock_core_repo.list_product_nav_history.return_value = []
|
|
|
|
tool = create_query_fund_nav_tool(mock_core_repo)
|
|
tool.invoke({"fund_code": "000001"})
|
|
|
|
# 验证调用时使用了 PROD-000001
|
|
mock_core_repo.get_product.assert_called_once_with("PROD-000001")
|
|
|
|
def test_returns_error_when_fund_not_found(self):
|
|
"""基金不存在时应返回错误"""
|
|
mock_core_repo = Mock()
|
|
mock_core_repo.get_product.return_value = None
|
|
|
|
tool = create_query_fund_nav_tool(mock_core_repo)
|
|
result = tool.invoke({"fund_code": "999999"})
|
|
|
|
assert "error" in result
|
|
|
|
|
|
class TestManageKycSession:
|
|
"""测试 KYC 会话管理工具"""
|
|
|
|
def test_create_session(self):
|
|
"""action=create 应创建新会话"""
|
|
mock_kyc_service = Mock()
|
|
mock_kyc_service.create_session.return_value = {
|
|
"session_id": "sess-123",
|
|
"status": "active",
|
|
"suggested_question": "请问您的投资经验有多少年?",
|
|
}
|
|
|
|
tool = create_manage_kyc_tool(mock_kyc_service)
|
|
result = tool.invoke({"action": "create", "customer_id": "CUST-001"})
|
|
|
|
assert result["session_id"] == "sess-123"
|
|
assert result["status"] == "active"
|
|
mock_kyc_service.create_session.assert_called_once()
|
|
|
|
def test_chat_session(self):
|
|
"""action=chat 应发送消息并获取回复"""
|
|
mock_kyc_service = Mock()
|
|
mock_kyc_service.chat.return_value = {
|
|
"session_id": "sess-123",
|
|
"assistant_message": "了解了,您的风险承受能力评估为C3",
|
|
"collected_fields": {"risk_tolerance": "C3"},
|
|
"progress": 60,
|
|
}
|
|
|
|
tool = create_manage_kyc_tool(mock_kyc_service)
|
|
result = tool.invoke({
|
|
"action": "chat",
|
|
"session_id": "sess-123",
|
|
"user_input": "我能接受10%的亏损",
|
|
})
|
|
|
|
assert result["session_id"] == "sess-123"
|
|
assert "assistant_message" in result
|
|
assert result["progress"] == 60
|
|
|
|
def test_get_session(self):
|
|
"""action=get 应返回会话状态"""
|
|
mock_kyc_service = Mock()
|
|
mock_kyc_service.get_session.return_value = {
|
|
"session_id": "sess-123",
|
|
"status": "active",
|
|
"progress": 40,
|
|
"collected_fields": {"age": 35},
|
|
}
|
|
|
|
tool = create_manage_kyc_tool(mock_kyc_service)
|
|
result = tool.invoke({"action": "get", "session_id": "sess-123"})
|
|
|
|
assert result["status"] == "active"
|
|
assert result["progress"] == 40
|
|
|
|
def test_complete_session(self):
|
|
"""action=complete 应完成会话"""
|
|
mock_kyc_service = Mock()
|
|
mock_kyc_service.complete_session.return_value = {
|
|
"session_id": "sess-123",
|
|
"status": "completed",
|
|
}
|
|
|
|
tool = create_manage_kyc_tool(mock_kyc_service)
|
|
result = tool.invoke({"action": "complete", "session_id": "sess-123"})
|
|
|
|
assert result["status"] == "completed"
|
|
|
|
def test_invalid_action(self):
|
|
"""无效的 action 应返回错误"""
|
|
tool = create_manage_kyc_tool(Mock())
|
|
result = tool.invoke({"action": "invalid"})
|
|
|
|
assert "error" in result
|
|
assert "invalid" in result["error"].lower()
|
|
|
|
|
|
class TestSearchTemplates:
|
|
"""测试模板搜索工具"""
|
|
|
|
def test_returns_matched_templates(self):
|
|
"""应返回匹配的模板列表"""
|
|
mock_template_service = Mock()
|
|
mock_template_service.search.return_value = [
|
|
{
|
|
"template_id": 1,
|
|
"title": "稳健投资话术",
|
|
"content": "建议采用分散投资策略...",
|
|
"score": 0.95,
|
|
},
|
|
{
|
|
"template_id": 2,
|
|
"title": "风险控制话术",
|
|
"content": "在当前市场环境下...",
|
|
"score": 0.85,
|
|
},
|
|
]
|
|
|
|
tool = create_search_templates_tool(mock_template_service)
|
|
result = tool.invoke({"query": "稳健投资", "limit": 5})
|
|
|
|
assert len(result["templates"]) == 2
|
|
assert result["templates"][0]["title"] == "稳健投资话术"
|
|
|
|
def test_handles_empty_results(self):
|
|
"""无匹配结果时应返回空列表"""
|
|
mock_template_service = Mock()
|
|
mock_template_service.search.return_value = []
|
|
|
|
tool = create_search_templates_tool(mock_template_service)
|
|
result = tool.invoke({"query": "不存在的话术"})
|
|
|
|
assert result["templates"] == []
|
|
|
|
|
|
class TestComplianceCheck:
|
|
"""测试合规检查工具"""
|
|
|
|
def test_returns_compliance_result(self):
|
|
"""应返回合规检查结果"""
|
|
mock_compliance_service = Mock()
|
|
mock_compliance_service.check.return_value = {
|
|
"risk_level": "INFO",
|
|
"hits": [],
|
|
"can_copy": True,
|
|
"message": "内容合规",
|
|
}
|
|
|
|
tool = create_compliance_check_tool(mock_compliance_service)
|
|
result = tool.invoke({"text": "建议关注稳健增长基金A", "scene": "advisor_chat"})
|
|
|
|
assert result["risk_level"] == "INFO"
|
|
assert result["can_copy"] is True
|
|
|
|
def test_detects_block_risk(self):
|
|
"""应检测到 BLOCK 级别风险"""
|
|
mock_compliance_service = Mock()
|
|
mock_compliance_service.check.return_value = {
|
|
"risk_level": "BLOCK",
|
|
"hits": [{"rule_id": 1, "pattern": "保证收益"}],
|
|
"can_copy": False,
|
|
"message": "包含违规内容",
|
|
}
|
|
|
|
tool = create_compliance_check_tool(mock_compliance_service)
|
|
result = tool.invoke({"text": "这只基金保证收益", "scene": "advisor_chat"})
|
|
|
|
assert result["risk_level"] == "BLOCK"
|
|
assert result["can_copy"] is False
|
|
assert len(result["hits"]) > 0
|