473 lines
18 KiB
Python
473 lines
18 KiB
Python
"""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]
|