Files
group_xinghuo_jinrong/tests/test_alert_service.py
T

171 lines
6.1 KiB
Python
Raw Normal View History

"""alert_service 单测(B2 · 聚合/去重/审计/推送 + 并发冒烟)。
sqlite StaticPool 单连接共享内存库,跨线程可用;publisher 注入 fake。
"""
import json
from datetime import datetime
from decimal import Decimal
from threading import Thread
import pytest
from sqlalchemy import text
from _ddl import create_sqlite_engine
from app.repository.risk_repository import RiskRepository
from app.service.risk import alert_service
from app.service.risk.alert_service import (
record_aml_alert,
record_suitability_alert,
record_trade_alerts,
set_publisher,
)
from app.service.risk.rules import RULE_ALERT_TYPES, RuleHit
class FakePublisher:
"""对齐 AlertPublisher 契约:publish(channel, payload_dict),序列化由实现方负责。"""
def __init__(self):
self.messages = []
self.deletes = []
def publish(self, channel, payload):
assert isinstance(payload, dict)
self.messages.append((channel, payload))
def delete(self, *keys):
self.deletes.append(keys)
@pytest.fixture()
def env():
engine = create_sqlite_engine() # DDL 单一事实源(B4 评审 P3-12)
repo = RiskRepository(engine=engine)
pub = FakePublisher()
set_publisher(pub)
yield repo, pub, engine
set_publisher(None)
engine.dispose()
def _trade(trade_id, amount="600000", customer="C1"):
return {
"trade_id": trade_id,
"customer_id": customer,
"product_id": "P1",
"trade_type": "subscribe",
"amount": Decimal(amount),
"traded_at": datetime(2026, 9, 6, 14, 0, 0),
}
def _hit(rule_id, score):
return RuleHit(rule_id, RULE_ALERT_TYPES[rule_id], score, "detail")
def _counts(engine, table, where="1=1", params=None):
with engine.connect() as conn:
return conn.execute(text(f"SELECT COUNT(*) FROM {table} WHERE {where}"), params or {}).scalar_one()
def test_no_hits_records_pass_audit(env):
repo, _, engine = env
result = record_trade_alerts(_trade("T1"), [], risk_repo=repo)
assert result is None
assert _counts(engine, "risk_alert") == 0
assert _counts(engine, "audit_log", "decision = 'pass'") == 1
def test_first_large_trade_creates_agg_alert(env):
repo, pub, engine = env
hits = [_hit("RISK-001", 70), _hit("RISK-002", 70)]
alert = record_trade_alerts(_trade("T1"), hits, risk_repo=repo)
assert alert["alert_type"] == "large_amount"
assert alert["triggered_rules"] == ["RISK-001", "RISK-002"]
assert alert["risk_score"] == 70
assert alert["status"] == "pending_review"
assert len(alert["payload"]["events"]) == 1
assert alert["trace_id"].startswith("trc-")
assert _counts(engine, "risk_alert") == 1
assert _counts(engine, "audit_log", "decision = 'alert_created'") == 1
(channel, payload), = pub.messages
assert channel == "risk:pub:alert"
assert payload["alert_id"] == alert["alert_id"]
assert payload["customer_id_mask"].endswith("**")
def test_same_day_second_trade_appends_not_duplicates(env):
repo, pub, engine = env
record_trade_alerts(_trade("T1"), [_hit("RISK-001", 70), _hit("RISK-002", 70)], risk_repo=repo)
alert = record_trade_alerts(_trade("T2"), [_hit("RISK-001", 70)], risk_repo=repo)
assert _counts(engine, "risk_alert") == 1 # 同日仅一张事件类单
got = repo.get_alert(alert["alert_id"])
assert len(got["payload"]["events"]) == 2
assert got["triggered_rules"] == ["RISK-001", "RISK-002"]
assert _counts(engine, "audit_log", "decision = 'alert_appended'") == 1
assert len(pub.messages) == 2 # 每次事件都广播,消费端按 alert_id 聚合
def test_alert_type_follows_highest_score_rule(env):
"""类型随最高分规则动态更新(PRD FR-4;评审 P1-3 口径)。"""
repo, _, _ = env
first = record_trade_alerts(_trade("T1"), [_hit("RISK-001", 70)], risk_repo=repo)
assert first["alert_type"] == "large_amount"
updated = record_trade_alerts(_trade("T2", "800000"), [_hit("RISK-005", 80)], risk_repo=repo)
assert updated["alert_id"] == first["alert_id"]
got = repo.get_alert(first["alert_id"])
assert got["alert_type"] == "pattern" and got["risk_score"] == 80
assert set(got["triggered_rules"]) == {"RISK-001", "RISK-005"}
def test_suitability_dedup_per_product(env):
repo, _, engine = env
req1 = _trade("T1", "100000")
req2 = _trade("T2", "100000")
a1 = record_suitability_alert(req1, "SUIT-001", "C1 不可购 R4", risk_repo=repo)
a2 = record_suitability_alert(req2, "SUIT-001", "C1 不可购 R4", risk_repo=repo)
assert a1["alert_id"] == a2["alert_id"] # 同客户+产品+日去重
got = repo.get_alert(a1["alert_id"])
assert got["alert_type"] == "suitability" and len(got["payload"]["events"]) == 2
# 不同产品独立出单
req3 = _trade("T3", "100000")
req3["product_id"] = "P2"
a3 = record_suitability_alert(req3, "SUIT-001", "x", risk_repo=repo)
assert a3["alert_id"] != a1["alert_id"]
assert _counts(engine, "risk_alert") == 2
def test_aml_standalone_alert_with_compliance_notify(env):
repo, pub, _ = env
alert = record_aml_alert("C1", {"list_type": "sanction", "match": 0.93}, risk_repo=repo)
assert alert["alert_type"] == "aml" and alert["risk_score"] == 95
(_, payload), = pub.messages
assert "compliance" in payload["notify_role"]
def test_concurrent_first_alerts_single_row(env):
"""并发冒烟(B2 验收):两线程同客户同日首单 → 预警单 1 张、events 含两笔。"""
repo, _, engine = env
errors = []
def worker(trade_id):
try:
record_trade_alerts(
_trade(trade_id), [_hit("RISK-001", 70), _hit("RISK-002", 70)], risk_repo=repo
)
except Exception as exc: # pragma: no cover
errors.append(exc)
threads = [Thread(target=worker, args=(f"T{i}",)) for i in (1, 2)]
for t in threads:
t.start()
for t in threads:
t.join()
assert not errors, errors
assert _counts(engine, "risk_alert") == 1
with engine.connect() as conn:
events = conn.execute(text("SELECT payload FROM risk_alert")).scalar_one()
assert len(json.loads(events)["events"]) == 2