165 lines
5.8 KiB
Python
165 lines
5.8 KiB
Python
"""L3 监测画像写入(B3 · PRD FR-7 / R-05 最小写入)。
|
||
|
||
合并规则(防降级,PRD FR-7):monitor_tier 取最高档(normal < watch < high),
|
||
monitor_tags 追加合并不覆盖,risk_score 取 max,computed_at 每次写当前时间
|
||
(列 NOT NULL 必须显式赋值),last_alert_id 联动最新预警。
|
||
|
||
并发:进程内锁按 customer_id 串行(多进程部署换 Redis SET NX,接口不变);
|
||
锁超时降级与跨进程竞态由 IntegrityError 兜底——重读合并后再 update,不丢更新。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
from datetime import datetime
|
||
from typing import Any
|
||
|
||
from sqlalchemy.exc import IntegrityError
|
||
|
||
from app.repository.risk_repository import RiskRepository
|
||
from app.service.risk.alert_service import _run_locked # 同包复用聚合锁;B7 收敛至公共原语
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
TIER_ORDER = ("normal", "watch", "high")
|
||
ALERT_TYPE_TIER = {
|
||
"aml": "high",
|
||
"pattern": "watch",
|
||
"large_amount": "watch",
|
||
"freq_trade": "watch",
|
||
"suitability": "normal",
|
||
}
|
||
AML_PENDING_TAG = "aml_hit_pending_review"
|
||
|
||
|
||
def tier_of(alert_type: str) -> str:
|
||
"""alert_type → L3 档位映射(PRD FR-7 固定映射;未知类型拒绝写入)。"""
|
||
tier = ALERT_TYPE_TIER.get(alert_type)
|
||
if tier is None:
|
||
raise ValueError(f"unknown alert_type for L3 mapping: {alert_type}")
|
||
return tier
|
||
|
||
|
||
def highest_tier(a: str, b: str) -> str:
|
||
"""取最高档(防降级核心原语)。"""
|
||
return a if TIER_ORDER.index(a) >= TIER_ORDER.index(b) else b
|
||
|
||
|
||
def merge_l3(
|
||
existing: dict[str, Any] | None,
|
||
*,
|
||
mapped_tier: str,
|
||
risk_score: int | None,
|
||
monitor_tags: list[str],
|
||
last_alert_id: str | None,
|
||
score_dimensions: dict[str, Any] | None = None,
|
||
computed_at: datetime | None = None,
|
||
) -> dict[str, Any]:
|
||
"""纯函数:existing 行(repo.get_l3 输出)与新事件合并后的 L3 行。
|
||
|
||
existing 为 None 表示新客户首写;risk_score 双方均空时保持 NULL。
|
||
"""
|
||
if existing is None:
|
||
tier, old_score, old_tags, old_dims = "normal", None, [], {}
|
||
else:
|
||
tier = existing.get("monitor_tier") or "normal"
|
||
old_score = existing.get("risk_score")
|
||
old_tags = existing.get("monitor_tags") or []
|
||
old_dims = existing.get("score_dimensions") or {}
|
||
|
||
scores = [int(s) for s in (old_score, risk_score) if s is not None]
|
||
return {
|
||
"monitor_tier": highest_tier(tier, mapped_tier),
|
||
"risk_score": max(scores) if scores else None,
|
||
"monitor_tags": sorted(set(old_tags) | set(monitor_tags)),
|
||
"score_dimensions": score_dimensions if score_dimensions is not None else old_dims,
|
||
"last_alert_id": last_alert_id,
|
||
"computed_at": computed_at or datetime.now(),
|
||
}
|
||
|
||
|
||
def upsert_profile_l3(
|
||
customer_id: str,
|
||
alert_type: str,
|
||
risk_score: int | None = None,
|
||
monitor_tags: list[str] | None = None,
|
||
last_alert_id: str | None = None,
|
||
score_dimensions: dict[str, Any] | None = None,
|
||
risk_repo: RiskRepository | None = None,
|
||
computed_at: datetime | None = None,
|
||
) -> dict[str, Any]:
|
||
"""预警事件 → L3 upsert(B2 预警落库后调用;aml 自动追加待复核标签)。
|
||
|
||
返回合并后的 L3 行(含 customer_id)。
|
||
"""
|
||
repo = risk_repo or RiskRepository()
|
||
mapped_tier = tier_of(alert_type)
|
||
tags = list(monitor_tags or [])
|
||
if alert_type == "aml" and AML_PENDING_TAG not in tags:
|
||
tags.append(AML_PENDING_TAG)
|
||
|
||
def _write(locked: bool) -> dict[str, Any]:
|
||
existing = repo.get_l3(customer_id)
|
||
merged = merge_l3(
|
||
existing,
|
||
mapped_tier=mapped_tier,
|
||
risk_score=risk_score,
|
||
monitor_tags=tags,
|
||
last_alert_id=last_alert_id,
|
||
score_dimensions=score_dimensions,
|
||
computed_at=computed_at,
|
||
)
|
||
try:
|
||
if existing is None:
|
||
repo.insert_l3(
|
||
customer_id,
|
||
merged["monitor_tier"],
|
||
merged["risk_score"],
|
||
merged["score_dimensions"],
|
||
merged["monitor_tags"],
|
||
merged["last_alert_id"],
|
||
merged["computed_at"],
|
||
)
|
||
else:
|
||
repo.update_l3(
|
||
customer_id,
|
||
merged["monitor_tier"],
|
||
merged["risk_score"],
|
||
merged["score_dimensions"],
|
||
merged["monitor_tags"],
|
||
merged["last_alert_id"],
|
||
merged["computed_at"],
|
||
)
|
||
except IntegrityError:
|
||
# 首写竞态(锁超时降级/跨进程):另一线程已 insert,重读合并转更新
|
||
if existing is not None:
|
||
raise
|
||
raced = repo.get_l3(customer_id)
|
||
merged = merge_l3(
|
||
raced,
|
||
mapped_tier=mapped_tier,
|
||
risk_score=risk_score,
|
||
monitor_tags=tags,
|
||
last_alert_id=last_alert_id,
|
||
score_dimensions=score_dimensions,
|
||
computed_at=computed_at,
|
||
)
|
||
repo.update_l3(
|
||
customer_id,
|
||
merged["monitor_tier"],
|
||
merged["risk_score"],
|
||
merged["score_dimensions"],
|
||
merged["monitor_tags"],
|
||
merged["last_alert_id"],
|
||
merged["computed_at"],
|
||
)
|
||
merged["customer_id"] = customer_id
|
||
return merged
|
||
|
||
return _run_locked(f"l3:{customer_id}", _write)
|
||
|
||
|
||
def get_profile_l3(customer_id: str, risk_repo: RiskRepository | None = None) -> dict[str, Any] | None:
|
||
"""L3 只读薄封装(对话线/引擎复用;Redis 缓存待 B7 lifespan 一并接入)。"""
|
||
return (risk_repo or RiskRepository()).get_l3(customer_id)
|