Files
XingHuo/app/repository/risk_repository.py
T

345 lines
13 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 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]