Files
XingHuo/app/repository/risk_repository.py
T

473 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""jinrong_agent 风控四表读写(risk_alert / risk_suitability_log / customer_profile_l3 / risk_aml_list)。
仅服务风控模块(PRD v1.0 §5.1 权限总表);不含业务判定逻辑(聚合合并/状态机入参
由 service 层决定后传入)。JSON 列以 dict 交互,内部序列化。
"""
from __future__ import annotations
import json
from datetime import datetime
from decimal import Decimal
from typing import Any
from sqlalchemy import text
from sqlalchemy.engine import Engine
from app.config.settings import settings
from app.utils.db import get_engine
EVENT_ALERT_TYPES = ("large_amount", "freq_trade", "pattern")
_AUDIT_SQL = text(
"""
INSERT INTO audit_log
(trace_id, event_type, agent_type, actor_id, customer_id, rule_id,
input_summary, decision, risk_score, handler_id, handler_result, handler_comment)
VALUES (:trace_id, :event_type, :agent_type, :actor_id, :customer_id, :rule_id,
:input_summary, :decision, :risk_score, :handler_id, :handler_result, :handler_comment)
"""
)
class RiskRepository:
"""风控产出表读写;不修改表结构(PRD 冻结约束)。"""
def __init__(self, engine: Engine | None = None) -> None:
self._engine = engine or get_engine(settings.mysql_database)
# ---------- risk_alert ----------
def find_pending_event_alert(self, customer_id: str, day_start: datetime) -> dict | None:
"""当日该客户的事件类 pending 单(聚合锚点;PRD FR-4 同日一张)。"""
return self._find_pending(customer_id, day_start, types=EVENT_ALERT_TYPES)
def find_pending_suitability_alert(
self, customer_id: str, product_id: str, day_start: datetime
) -> dict | None:
row = self._find_pending(customer_id, day_start, types=("suitability",))
if row and row["payload"].get("product_id") == product_id:
return row
return None
def _find_pending(
self, customer_id: str, day_start: datetime, types: tuple[str, ...]
) -> dict | None:
placeholders = ", ".join(f":t{i}" for i in range(len(types)))
params: dict[str, Any] = {"cid": customer_id, "day_start": day_start}
for i, t in enumerate(types):
params[f"t{i}"] = t
sql = text(
f"""
SELECT * FROM risk_alert
WHERE customer_id = :cid
AND status = 'pending_review'
AND alert_type IN ({placeholders})
AND created_at >= :day_start
ORDER BY created_at DESC
LIMIT 1
"""
)
with self._engine.connect() as conn:
row = conn.execute(sql, params).mappings().first()
return self._parse_alert(dict(row)) if row else None
def find_recent_aml_alert(self, customer_id: str, day_start: datetime) -> dict | None:
"""该客户当日最近一张 aml 单(scan 幂等检查 · B9b 核查单②/B6 评审 P3-6)。
不限 status——已处置单也算「已出」,同扫描日重扫不重复出单(模拟每日
批量口径:跨日再扫出新单)。交易触发的当日 aml 单同样命中(当日命中
已留痕,scan 无需重复出单)。
"""
sql = text(
"""
SELECT * FROM risk_alert
WHERE customer_id = :cid AND alert_type = 'aml' AND created_at >= :day_start
ORDER BY created_at DESC
LIMIT 1
"""
)
with self._engine.connect() as conn:
row = conn.execute(sql, {"cid": customer_id, "day_start": day_start}).mappings().first()
return self._parse_alert(dict(row)) if row else None
def insert_alert(self, alert: dict[str, Any]) -> None:
sql = text(
"""
INSERT INTO risk_alert
(alert_id, trace_id, customer_id, trade_id, alert_type, triggered_rules,
risk_score, status, payload)
VALUES (:alert_id, :trace_id, :customer_id, :trade_id, :alert_type, :triggered_rules,
:risk_score, :status, :payload)
"""
)
with self._engine.begin() as conn:
conn.execute(sql, self._dump_alert(alert))
def append_alert_event(
self,
alert_id: str,
event: dict[str, Any],
triggered_rules: list[str],
risk_score: int,
alert_type: str,
) -> None:
"""读改写:追加 payload.events、合并 triggered_rules、risk_score 取 max、alert_type 更新。"""
with self._engine.begin() as conn:
row = conn.execute(
text("SELECT payload, triggered_rules, risk_score FROM risk_alert WHERE alert_id = :aid"),
{"aid": alert_id},
).mappings().first()
if row is None:
raise ValueError(f"alert not found: {alert_id}")
payload = json.loads(row["payload"]) if isinstance(row["payload"], str) else row["payload"]
events = payload.setdefault("events", [])
events.append(event)
old_rules = (
json.loads(row["triggered_rules"])
if isinstance(row["triggered_rules"], str)
else row["triggered_rules"]
)
merged_rules = sorted(set(old_rules) | set(triggered_rules))
old_score = int(row["risk_score"] or 0)
new_score = max(old_score, int(risk_score))
conn.execute(
text(
"""
UPDATE risk_alert
SET payload = :payload, triggered_rules = :rules,
risk_score = :score, alert_type = :atype
WHERE alert_id = :aid
"""
),
{
"payload": json.dumps(payload, ensure_ascii=False),
"rules": json.dumps(merged_rules),
"score": new_score,
"atype": alert_type,
"aid": alert_id,
},
)
def get_alert(self, alert_id: str) -> dict | None:
with self._engine.connect() as conn:
row = conn.execute(
text("SELECT * FROM risk_alert WHERE alert_id = :aid"), {"aid": alert_id}
).mappings().first()
return self._parse_alert(dict(row)) if row else None
def find_alerts_by_trade(self, trade_id: str) -> list[dict]:
"""trade_id 已入的预警单(rebuild_alerts 幂等检查;含已处置单——重放不得绕过处置结论)。
payload 为 JSON 列(MySQL)/TEXT JSON(sqlite 测试),LIKE 匹配 trade_id 字符串;
trade_id 形如 TRD-{date}-{uuid8} 唯一性强,无误匹配面。首笔交易的 trade_id
恒写 payload.events[0].trade_id,LIKE 覆盖事件类/suitability/aml 全部出单路径。
"""
sql = text(
"""
SELECT alert_id, alert_type, status
FROM risk_alert
WHERE payload LIKE :pat
ORDER BY created_at
"""
)
with self._engine.connect() as conn:
return [
dict(r)
for r in conn.execute(sql, {"pat": f"%{trade_id}%"}).mappings()
]
def has_engine_error_audit(self, trade_id: str) -> bool:
"""存在 decision='risk_engine_error' 的该笔审计(rebuild_alerts 警示用)。
audit_log 无 trade_id 列,trade_id 在 input_summary(JSON 列/TEXT JSON)
内,同 find_alerts_by_trade 的 LIKE 口径。
"""
sql = text(
"""
SELECT 1 FROM audit_log
WHERE decision = 'risk_engine_error' AND input_summary LIKE :pat
LIMIT 1
"""
)
with self._engine.connect() as conn:
return conn.execute(sql, {"pat": f"%{trade_id}%"}).first() is not None
def list_alerts(
self,
status: str | None = None,
alert_type: str | None = None,
customer_id: str | None = None,
start: datetime | None = None,
end: datetime | None = None,
page: int = 1,
page_size: int = 20,
) -> tuple[list[dict], int]:
"""预警台账分页查询(PRD FR-4);过滤参数全部可选。"""
where = ["1=1"]
params: dict[str, Any] = {}
if status:
where.append("status = :status")
params["status"] = status
if alert_type:
where.append("alert_type = :atype")
params["atype"] = alert_type
if customer_id:
where.append("customer_id = :cid")
params["cid"] = customer_id
if start:
where.append("created_at >= :start")
params["start"] = start
if end:
where.append("created_at < :end")
params["end"] = end
where_sql = " AND ".join(where)
with self._engine.connect() as conn:
total = conn.execute(
text(f"SELECT COUNT(*) FROM risk_alert WHERE {where_sql}"), params
).scalar_one()
rows = conn.execute(
text(
f"""
SELECT * FROM risk_alert WHERE {where_sql}
ORDER BY created_at DESC
LIMIT :lim OFFSET :off
"""
),
{**params, "lim": page_size, "off": (page - 1) * page_size},
).mappings()
return [self._parse_alert(dict(r)) for r in rows], int(total)
def update_alert_status(
self, alert_id: str, handler_result: str, handler_id: str, handler_comment: str | None
) -> bool:
"""状态机:仅 pending_review 可处置(PRD FR-4);返回是否发生变更。"""
sql = text(
"""
UPDATE risk_alert
SET status = :result, handler_id = :hid, handler_result = :result,
handler_comment = :comment, handled_at = :handled_at
WHERE alert_id = :aid AND status = 'pending_review'
"""
)
with self._engine.begin() as conn:
res = conn.execute(
sql,
{
"result": handler_result,
"hid": handler_id,
"comment": handler_comment,
"handled_at": datetime.now(),
"aid": alert_id,
},
)
return res.rowcount == 1
def handle_alert_with_audit(
self,
alert_id: str,
handler_result: str,
handler_id: str,
handler_comment: str | None,
audit_entry: dict[str, Any],
) -> bool:
"""状态机变更 + 处置审计**同事务**(B7 挂账⑦)。
UPDATE ... WHERE status='pending_review'(防并发抢先)命中后才同事务
INSERT audit_log——审计失败整体回滚,不产生无痕状态变更(P-05);返回
False 表示已被并发处置,service 层转 StateConflict。
"""
update_sql = text(
"""
UPDATE risk_alert
SET status = :result, handler_id = :hid, handler_result = :result,
handler_comment = :comment, handled_at = :handled_at
WHERE alert_id = :aid AND status = 'pending_review'
"""
)
with self._engine.begin() as conn:
res = conn.execute(
update_sql,
{
"result": handler_result,
"hid": handler_id,
"comment": handler_comment,
"handled_at": datetime.now(),
"aid": alert_id,
},
)
if res.rowcount != 1:
return False
conn.execute(_AUDIT_SQL, self._dump_audit(audit_entry))
return True
@staticmethod
def _dump_alert(alert: dict[str, Any]) -> dict[str, Any]:
out = dict(alert)
out["triggered_rules"] = json.dumps(alert["triggered_rules"])
out["payload"] = json.dumps(alert["payload"], ensure_ascii=False, default=str)
return out
@staticmethod
def _parse_alert(row: dict[str, Any]) -> dict[str, Any]:
row["triggered_rules"] = json.loads(row["triggered_rules"])
row["payload"] = json.loads(row["payload"]) if isinstance(row["payload"], str) else row["payload"]
return row
# ---------- audit_log(风控判定审计 · 只 INSERT,PRD §7.3)----------
def insert_audit_log(self, entry: dict[str, Any]) -> None:
with self._engine.begin() as conn:
conn.execute(_AUDIT_SQL, self._dump_audit(entry))
# ---------- input_guard_log(输入安全防护 · 只 INSERT,手册 P-05 双写)----------
def insert_input_guard_log(
self,
trace_id: str,
agent_type: str,
actor_id: str,
guard_type: str,
action: str,
raw_excerpt: str | None = None,
session_id: str | None = None,
) -> None:
"""安全防护留痕(T-02 挂账⑧收口);ENUM 口径见 01-mysql-共用底座.sql。"""
sql = text(
"""
INSERT INTO input_guard_log
(trace_id, session_id, agent_type, actor_id, guard_type, raw_excerpt, action)
VALUES (:trace_id, :session_id, :agent_type, :actor_id, :guard_type, :raw_excerpt, :action)
"""
)
with self._engine.begin() as conn:
conn.execute(
sql,
{
"trace_id": trace_id,
"session_id": session_id,
"agent_type": agent_type,
"actor_id": actor_id,
"guard_type": guard_type,
"raw_excerpt": (raw_excerpt or "")[:1024],
"action": action,
},
)
@staticmethod
def _dump_audit(entry: dict[str, Any]) -> dict[str, Any]:
params = dict(entry)
params["input_summary"] = json.dumps(
entry.get("input_summary") or {}, ensure_ascii=False, default=str
)
return params
# ---------- risk_suitability_log ----------
def insert_suitability_log(self, log: dict[str, Any]) -> None:
sql = text(
"""
INSERT INTO risk_suitability_log
(trace_id, customer_id, product_id, customer_risk_level, product_risk_level,
is_matched, is_blocked, block_reason, request_ref, profile_l1_version)
VALUES (:trace_id, :customer_id, :product_id, :customer_risk_level, :product_risk_level,
:is_matched, :is_blocked, :block_reason, :request_ref, :profile_l1_version)
"""
)
with self._engine.begin() as conn:
conn.execute(sql, log)
# ---------- customer_profile_l3 ----------
def get_l3(self, customer_id: str) -> dict | None:
with self._engine.connect() as conn:
row = conn.execute(
text("SELECT * FROM customer_profile_l3 WHERE customer_id = :cid"), {"cid": customer_id}
).mappings().first()
if not row:
return None
data = dict(row)
for key in ("score_dimensions", "monitor_tags"):
if isinstance(data.get(key), str):
data[key] = json.loads(data[key])
return data
def insert_l3(
self,
customer_id: str,
monitor_tier: str,
risk_score: int | None,
score_dimensions: dict | list | None,
monitor_tags: list[str],
last_alert_id: str | None,
computed_at: datetime,
) -> None:
"""新行插入;行已存在时由 service 层读后改走 update_l3(合并逻辑在 service)。"""
sql = text(
"""
INSERT INTO customer_profile_l3
(customer_id, monitor_tier, risk_score, score_dimensions, monitor_tags,
last_alert_id, computed_at)
VALUES (:cid, :tier, :score, :dims, :tags, :last_alert, :computed_at)
"""
)
with self._engine.begin() as conn:
conn.execute(
sql,
{
"cid": customer_id,
"tier": monitor_tier,
"score": risk_score,
"dims": json.dumps(score_dimensions or {}, ensure_ascii=False),
"tags": json.dumps(monitor_tags, ensure_ascii=False),
"last_alert": last_alert_id,
"computed_at": computed_at,
},
)
def update_l3(
self,
customer_id: str,
monitor_tier: str,
risk_score: int | None,
score_dimensions: dict | list | None,
monitor_tags: list[str],
last_alert_id: str | None,
computed_at: datetime,
expected_computed_at: datetime | None = None,
) -> bool:
"""整行更新(service 层完成最高档/tags 合并后调用;computed_at NOT NULL 必传)。
expected_computed_at 传入时为乐观锁(比对读时的 computed_at),行已被并发
修改则不命中,返回 False 由 service 层重读重试(B3 评审 P1-1 防丢更新)。
"""
sql = """
UPDATE customer_profile_l3
SET monitor_tier = :tier, risk_score = :score, score_dimensions = :dims,
monitor_tags = :tags, last_alert_id = :last_alert, computed_at = :computed_at
WHERE customer_id = :cid"""
params: dict[str, Any] = {
"cid": customer_id,
"tier": monitor_tier,
"score": risk_score,
"dims": json.dumps(score_dimensions or {}, ensure_ascii=False),
"tags": json.dumps(monitor_tags, ensure_ascii=False),
"last_alert": last_alert_id,
"computed_at": computed_at,
}
if expected_computed_at is not None:
sql += " AND computed_at = :expected"
params["expected"] = expected_computed_at
with self._engine.begin() as conn:
res = conn.execute(text(sql), params)
return res.rowcount == 1
# ---------- risk_aml_list ----------
def list_active_aml_entries(self) -> list[dict]:
with self._engine.connect() as conn:
rows = conn.execute(
text("SELECT * FROM risk_aml_list WHERE is_active = 1")
).mappings()
return [dict(r) for r in rows]