136 lines
4.0 KiB
Python
136 lines
4.0 KiB
Python
from datetime import datetime
|
|
from types import MappingProxyType, SimpleNamespace
|
|
|
|
import pytest
|
|
|
|
from app.core.contracts import RequestContext
|
|
from app.model.fund import FundRiskAlert
|
|
from app.repository.fund_query_repository import FundRecord
|
|
from app.service.risk_analysis_service import RiskAnalysisService
|
|
|
|
|
|
class FakeRepository:
|
|
async def get_alert_detail(self, _alert_no: str) -> FundRecord:
|
|
return FundRecord(
|
|
entity="risk_alert_detail",
|
|
values=MappingProxyType({
|
|
"alert": {
|
|
"alert_no": "ALERT-001",
|
|
"alert_type": "适当性错配",
|
|
"risk_level": "高",
|
|
"rule_codes": ["RW-007"],
|
|
"evidence_summary": "C2客户购买R5产品且留痕不完整。",
|
|
"due_time": "2026-09-10T12:00:00",
|
|
},
|
|
"customer": {"customer_no": "CUST-001", "name": "张*", "risk_level": "C2"},
|
|
"transaction": {"amount": "700000.00"},
|
|
}),
|
|
)
|
|
|
|
|
|
class FakeSession:
|
|
def __init__(self, alert: FundRiskAlert):
|
|
self.alert = alert
|
|
self.added = []
|
|
self.committed = False
|
|
|
|
async def scalar(self, _statement):
|
|
return self.alert
|
|
|
|
def add(self, value):
|
|
self.added.append(value)
|
|
|
|
async def commit(self):
|
|
self.committed = True
|
|
|
|
|
|
class FailingResolver:
|
|
async def resolve(self, **_kwargs):
|
|
raise RuntimeError("模型端点不可用")
|
|
|
|
|
|
class UnusedModelService:
|
|
async def generate(self, *_args, **_kwargs):
|
|
raise AssertionError("端点解析失败后不应调用模型")
|
|
|
|
|
|
def context() -> RequestContext:
|
|
return RequestContext(
|
|
user_id="990000002",
|
|
trace_id="analysis-trace",
|
|
permissions=("risk:alert:read",),
|
|
data_scope="all",
|
|
)
|
|
|
|
|
|
def alert() -> FundRiskAlert:
|
|
return FundRiskAlert(
|
|
id=1,
|
|
alert_no="ALERT-001",
|
|
customer_id=9,
|
|
alert_type="适当性错配",
|
|
alert_level="高",
|
|
trigger_rule_codes=["RW-007"],
|
|
evidence_summary="证据",
|
|
evidence_snapshot={},
|
|
priority_score=90,
|
|
event_status="正在发生",
|
|
status="待处理",
|
|
ack_status="未确认",
|
|
is_escalated=0,
|
|
created_at=datetime(2026, 9, 10),
|
|
updated_at=datetime(2026, 9, 10),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
("output_type", "expected"),
|
|
[
|
|
("预警研判", "风险结论"),
|
|
("回访话术", "开场说明"),
|
|
("工单摘要", "工单标题"),
|
|
],
|
|
)
|
|
async def test_analysis_fallback_contains_original_structure(
|
|
output_type: str,
|
|
expected: str,
|
|
) -> None:
|
|
session = FakeSession(alert())
|
|
service = RiskAnalysisService(
|
|
session,
|
|
repository=FakeRepository(), # type: ignore[arg-type]
|
|
model_service=UnusedModelService(),
|
|
endpoint_resolver=FailingResolver(),
|
|
)
|
|
|
|
result = await service.generate(context(), "ALERT-001", output_type)
|
|
|
|
assert expected in result["content"]
|
|
assert result["source"] == "模板降级输出"
|
|
assert session.alert.ai_analysis[output_type]["content"] == result["content"]
|
|
assert session.committed is True
|
|
assert session.added[0].action_type == "risk_ai_analysis_generated"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_forbidden_model_claim_falls_back_to_template() -> None:
|
|
class Resolver:
|
|
async def resolve(self, **_kwargs):
|
|
return [object()]
|
|
|
|
class ModelService:
|
|
async def generate(self, *_args, **_kwargs):
|
|
return SimpleNamespace(text="我已关闭预警")
|
|
|
|
session = FakeSession(alert())
|
|
result = await RiskAnalysisService(
|
|
session,
|
|
repository=FakeRepository(), # type: ignore[arg-type]
|
|
model_service=ModelService(),
|
|
endpoint_resolver=Resolver(),
|
|
).generate(context(), "ALERT-001", "预警研判")
|
|
|
|
assert result["source"] == "模板降级输出"
|
|
assert "风险结论" in result["content"]
|