345 lines
13 KiB
Python
345 lines
13 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 create_engine, text
|
||
from sqlalchemy.engine import Engine
|
||
|
||
from app.config.settings import settings
|
||
|
||
EVENT_ALERT_TYPES = ("large_amount", "freq_trade", "pattern")
|
||
|
||
|
||
class RiskRepository:
|
||
"""风控产出表读写;不修改表结构(PRD 冻结约束)。"""
|
||
|
||
def __init__(self, engine: Engine | None = None) -> None:
|
||
self._engine = engine or self._default_engine()
|
||
|
||
@staticmethod
|
||
def _default_engine() -> Engine:
|
||
pwd = settings.mysql_password
|
||
auth = f"{settings.mysql_user}:{pwd}" if pwd else settings.mysql_user
|
||
url = (
|
||
f"mysql+pymysql://{auth}@{settings.mysql_host}:{settings.mysql_port}"
|
||
f"/{settings.mysql_database}?charset=utf8mb4"
|
||
)
|
||
return create_engine(url, pool_pre_ping=True)
|
||
|
||
# ---------- 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 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 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
|
||
|
||
@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:
|
||
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)
|
||
"""
|
||
)
|
||
params = dict(entry)
|
||
params["input_summary"] = json.dumps(
|
||
entry.get("input_summary") or {}, ensure_ascii=False, default=str
|
||
)
|
||
with self._engine.begin() as conn:
|
||
conn.execute(sql, 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,
|
||
) -> None:
|
||
"""整行更新(service 层完成最高档/tags 合并后调用;computed_at NOT NULL 必传)。"""
|
||
sql = text(
|
||
"""
|
||
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
|
||
"""
|
||
)
|
||
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,
|
||
},
|
||
)
|
||
|
||
# ---------- 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]
|