Files
group_xinghuo_jinrong/tests/test_step10_agent_tools.py
T

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