diff --git a/app/repository/risk_repository.py b/app/repository/risk_repository.py index b85481b..0dafb83 100644 --- a/app/repository/risk_repository.py +++ b/app/repository/risk_repository.py @@ -310,29 +310,33 @@ class RiskRepository: monitor_tags: list[str], last_alert_id: str | None, computed_at: datetime, - ) -> None: - """整行更新(service 层完成最高档/tags 合并后调用;computed_at NOT NULL 必传)。""" - sql = text( - """ + 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 - """ - ) + 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: - 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, - }, - ) + res = conn.execute(text(sql), params) + return res.rowcount == 1 # ---------- risk_aml_list ---------- diff --git a/app/service/risk/profile_l3.py b/app/service/risk/profile_l3.py index 7fd508f..7f8505b 100644 --- a/app/service/risk/profile_l3.py +++ b/app/service/risk/profile_l3.py @@ -1,11 +1,15 @@ """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 联动最新预警。 +monitor_tags 追加合并不覆盖,risk_score 取 max(FR-7 未定义的口径外推,见 +RISK-00n 静态分;R-05 动态评分接入时复核),last_alert_id 传入时联动最新预警 +(未传保留旧值,防调用方漏传抹掉),computed_at 每次写当前时间(列 NOT NULL, +毫秒截断对齐 DATETIME(3),保证乐观锁读写比对一致)。 -并发:进程内锁按 customer_id 串行(多进程部署换 Redis SET NX,接口不变); -锁超时降级与跨进程竞态由 IntegrityError 兜底——重读合并后再 update,不丢更新。 +并发(B3 评审 P1-1 修复):进程内锁按 customer_id 串行减少冲突(多进程部署换 +Redis SET NX,接口不变);锁超时降级与跨进程竞态由乐观锁兜底——update 比对 +读时 computed_at,未命中或 insert 撞主键则重读合并重试,3 次仍冲突抛错 +(风控数据宁失败不静默覆盖)。 """ from __future__ import annotations @@ -30,6 +34,12 @@ ALERT_TYPE_TIER = { "suitability": "normal", } AML_PENDING_TAG = "aml_hit_pending_review" +_MAX_RETRIES = 3 + + +def _ms(dt: datetime) -> datetime: + """毫秒截断(对齐表列 DATETIME(3);否则乐观锁 expected 比对永不命中)。""" + return dt.replace(microsecond=(dt.microsecond // 1000) * 1000) def tier_of(alert_type: str) -> str: @@ -60,12 +70,13 @@ def merge_l3( existing 为 None 表示新客户首写;risk_score 双方均空时保持 NULL。 """ if existing is None: - tier, old_score, old_tags, old_dims = "normal", None, [], {} + tier, old_score, old_tags, old_dims, old_alert = "normal", None, [], {}, 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 {} + old_alert = existing.get("last_alert_id") scores = [int(s) for s in (old_score, risk_score) if s is not None] return { @@ -73,8 +84,8 @@ def merge_l3( "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(), + "last_alert_id": last_alert_id if last_alert_id is not None else old_alert, + "computed_at": _ms(computed_at or datetime.now()), } @@ -88,9 +99,10 @@ def upsert_profile_l3( risk_repo: RiskRepository | None = None, computed_at: datetime | None = None, ) -> dict[str, Any]: - """预警事件 → L3 upsert(B2 预警落库后调用;aml 自动追加待复核标签)。 + """预警事件 → L3 upsert(aml 自动追加待复核标签)。 - 返回合并后的 L3 行(含 customer_id)。 + 返回合并后的 L3 行(含 customer_id)。B7 接 Redis 缓存时须在此处加 + 写侧 DEL 钩子(PRD §5.1:MySQL 更新时 DEL profile:l3:{customer_id})。 """ repo = risk_repo or RiskRepository() mapped_tier = tier_of(alert_type) @@ -98,45 +110,11 @@ def upsert_profile_l3( 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) + def _write(locked: bool) -> dict[str, Any]: # 锁原语签名要求(是否获得锁);冲突安全由乐观锁重试保证 + for attempt in range(1, _MAX_RETRIES + 1): + existing = repo.get_l3(customer_id) merged = merge_l3( - raced, + existing, mapped_tier=mapped_tier, risk_score=risk_score, monitor_tags=tags, @@ -144,8 +122,7 @@ def upsert_profile_l3( score_dimensions=score_dimensions, computed_at=computed_at, ) - repo.update_l3( - customer_id, + fields = ( merged["monitor_tier"], merged["risk_score"], merged["score_dimensions"], @@ -153,12 +130,23 @@ def upsert_profile_l3( merged["last_alert_id"], merged["computed_at"], ) - merged["customer_id"] = customer_id - return merged + if existing is None: + try: + repo.insert_l3(customer_id, *fields) + merged["customer_id"] = customer_id + return merged + except IntegrityError: + logger.warning("L3 insert race, retrying (attempt %d): %s", attempt, customer_id) + continue + if repo.update_l3(customer_id, *fields, expected_computed_at=existing["computed_at"]): + merged["customer_id"] = customer_id + return merged + logger.warning("L3 update lost race, retrying (attempt %d): %s", attempt, customer_id) + raise RuntimeError(f"L3 upsert conflicted after {_MAX_RETRIES} retries: {customer_id}") 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 一并接入)。""" + """L3 只读薄封装(对话线/引擎复用;Redis 热缓存与写侧 DEL 钩子待 B7 lifespan 一并接入)。""" return (risk_repo or RiskRepository()).get_l3(customer_id) diff --git a/docs/项目框架设计/开发计划-风控模块.md b/docs/项目框架设计/开发计划-风控模块.md index 8d72587..7c602cf 100644 --- a/docs/项目框架设计/开发计划-风控模块.md +++ b/docs/项目框架设计/开发计划-风控模块.md @@ -26,7 +26,7 @@ | B4 | `aml_service.py`(归一化+相似度匹配、scan_all)+ `engine.py`(process_trade_event 组装 + 预留客户事件钩子)+ **`scoring.py` 占位签名(FR-7 预留)** | 引擎完整 | 单测:AML 阈值边界;引擎集成冒烟 | B1、B2、B3 | | B5 | `app/gateway/`(trade_gateway + gateway_repository 仅 INSERT core_trade)+ `api/simulate.py` 薄路由 | 网关 | 集成:convert 400、阻断不落 trade | A4、B4 | | B6 | **`app/api/deps.py`:`AuthContext`(actor_id/roles/customer_id,字段按 JWT 手册冻结)+ `get_auth_context()` 工厂**——dev 模式从 `X-Debug-Role`/`X-Debug-Actor` 请求头构造、`app_env != development` 启动时检测 debug 头直接拒绝;T-01 就绪后仅替换工厂内部为 JWT 解析,签名不变。另:`api/risk.py` 4 个 API(GET alerts / POST handle / POST suitability/check / POST aml/scan)+ 归属校验(含 compliance 强制 aml 过滤)。**备注:依赖层须校验 handler_result 枚举(repo 不校验);本阶段顺手统一 `NotFoundError` 异常(utils/exceptions.py 现为占位)**(评审 P2-7①②) | 鉴权依赖 + 4 个 API | Swagger 手测 + **权限矩阵(按 debug 头切换角色/身份执行 A-7/A-9 用例)** | A4、B2、**B4**(aml/scan 依赖 scan_all) | -| B7 | `main.py` 集成:路由挂载 + lifespan(双 Engine 单例注入 + Redis 单例 + trace 中间件)。**备注:顺手提取 `utils/db.py` 引擎工厂收敛 core_ro/risk_repository 双份 _default_engine**(评审 P2-6) | 可运行应用 | `uvicorn` 启动 + `/health` + 全路由可达 | B5、B6 | +| B7 | `main.py` 集成:路由挂载 + lifespan(双 Engine 单例注入 + Redis 单例 + trace 中间件)。**备注:顺手提取 `utils/db.py` 引擎工厂收敛 core_ro/risk_repository 双份 _default_engine**(评审 P2-6);**B3 挂账(B3 评审 P2-5):① `_run_locked` 锁原语公共化(alert_service/profile_l3 现复用私有实现)② L3 写侧 Redis 缓存 DEL 钩子(PRD §5.1 `profile:l3:{customer_id}` 更新时 DEL,`profile_l3.upsert_profile_l3` 已留痕)** | 可运行应用 | `uvicorn` 启动 + `/health` + 全路由可达 | B5、B6 | | B8 | **`tests/conftest.py`**:a) session fixture 启动校验演示数据就位(CUST-4001 测评 <365 天、risk_aml_list ≥8),缺失则中止并提示先跑 FLOW §0 ③④;b) fixture 幂等代跑 `prepare_risk_demo.sql`;c) teardown 按 `TRD-TEST-` 清 core_trade + 关联 risk_alert/risk_suitability_log/audit_log + 还原 L3 行。集成测试:A-1~A-5、A-7(状态机/compliance 403/GET 强制 aml)、A-9 越权、**trace 一致性断言** | 测试套件 + fixture | `pytest` 全绿 | B7 | | B9a | 演示/运维脚本开发:`scripts/demo/subscribe_alerts.py`(订阅演示)+ `scripts/demo/rebuild_alerts.py`(按 trade_id 幂等重放补偿) | 2 个脚本 | 手工执行验证 | B2、B4(可与 B5~B8 并行) | | B9b | 演示链路走查:`reset.ps1` → `prepare_risk_demo.sql` → agent 库建表 → `seed-aml-list.sql` → Swagger 逐条过 **A-1~A-5、A-7~A-9(A-6 归 M3)** | 演示 SOP | 按 PRD §8 验收表逐条打勾 | B8、B9a | diff --git a/tests/test_profile_l3.py b/tests/test_profile_l3.py index 1a9fca9..4b3e912 100644 --- a/tests/test_profile_l3.py +++ b/tests/test_profile_l3.py @@ -29,6 +29,8 @@ def env(): poolclass=StaticPool, connect_args={"check_same_thread": False}, ) + # 口径(评审 P3-3):以 VARCHAR/TEXT 近似真实 DDL 的 ENUM/SMALLINT/JSON, + # ENUM 档位防线不在测试 DB 层,由 tier_of 白名单 + highest_tier 保证。 with engine.begin() as conn: conn.execute( text( @@ -134,7 +136,7 @@ def test_normal_upgrades_to_watch(env): def test_watch_does_not_degrade_to_normal(env): - """ Suitability 映射 normal:已 watch 的客户不被 suitability 事件拉低。""" + """suitability 映射 normal:已 watch 的客户不被 suitability 事件拉低。""" repo, engine = env _upsert(repo, "C1", "pattern", 80, alert_id="ALT-1") merged = _upsert(repo, "C1", "suitability", 90, alert_id="ALT-2") @@ -150,6 +152,15 @@ def test_aml_marks_high_with_pending_review_tag(env): assert AML_PENDING_TAG in _row(engine, "C1")["monitor_tags"] +def test_high_does_not_degrade_to_normal(env): + """B3 验收(评审 P2-4):aml 后 suitability(映射 normal)不回落。""" + repo, engine = env + _upsert(repo, "C1", "aml", 95, alert_id="ALT-1") + merged = _upsert(repo, "C1", "suitability", 90, alert_id="ALT-2") + assert merged["monitor_tier"] == "high" + assert _row(engine, "C1")["monitor_tier"] == "high" + + def test_large_amount_after_aml_does_not_fall_back(env): """B3 验收:AML 后大额 → tier 仍 high、tags 并集、score 取 max。""" repo, engine = env @@ -166,8 +177,27 @@ def test_large_amount_after_aml_does_not_fall_back(env): def test_tags_accumulate_not_overwrite(env): repo, _ = env _upsert(repo, "C1", "pattern", 80, alert_id="ALT-1", tags=["freq"]) - merged = _upsert(repo, "C1", "pattern", 80, alert_id="ALT-2", tags=["freq"]) - assert merged["monitor_tags"] == ["freq"] # set 合并去重 + merged = _upsert(repo, "C1", "pattern", 80, alert_id="ALT-2", tags=["manual_review"]) + assert merged["monitor_tags"] == ["freq", "manual_review"] # 不同 tag 跨事件追加 + + +def test_last_alert_id_kept_when_not_passed(env): + """评审 P2-2:漏传 last_alert_id 不抹掉旧值(FR-7 联动语义防御)。""" + repo, engine = env + _upsert(repo, "C1", "aml", 95, alert_id="ALT-1") + merged = upsert_profile_l3("C1", "large_amount", 70, risk_repo=repo) + assert merged["last_alert_id"] == "ALT-1" + assert _row(engine, "C1")["last_alert_id"] == "ALT-1" + + +def test_computed_at_refreshed_on_each_write(env): + """评审 P3-2:每次写 computed_at 均刷新(FR-7)。""" + repo, engine = env + _upsert(repo, "C1", "pattern", 80, alert_id="ALT-1") + upsert_profile_l3("C1", "large_amount", 70, last_alert_id="ALT-2", + computed_at=datetime(2027, 1, 1, 8, 0, 0), risk_repo=repo) + # sqlite 读回为字符串,格式无关断言(核心是值已从首写的 now 刷新为传入时间) + assert str(_row(engine, "C1")["computed_at"]).startswith("2027-01-01 08:00") def test_score_dimensions_replaced_only_when_passed(env): @@ -245,3 +275,41 @@ def test_integrity_error_falls_back_to_remerge(env, monkeypatch): assert merged["last_alert_id"] == "ALT-2" row = _row(engine, "C1") assert row["monitor_tier"] == "high" and row["last_alert_id"] == "ALT-2" + + +def test_update_lost_race_retries_and_converges(env, monkeypatch): + """评审 P1-1:update 路径丢更新——本进程读旧值后对方先写 high,乐观锁未命中 + 触发重读重试,最终收敛 high/95(修复前会被覆盖回退 watch/70)。""" + repo, engine = env + repo.insert_l3("C1", "watch", 70, {}, [], "ALT-1", datetime.now()) + + calls = {"n": 0} + orig_update = type(repo).update_l3 + + def racing_update(self, *args, **kwargs): + calls["n"] += 1 + if calls["n"] == 1: + orig_update(self, "C1", "high", 95, {}, [AML_PENDING_TAG], "ALT-AML", datetime.now()) + return False # 对方抢先提交,本进程乐观锁未命中 + return orig_update(self, *args, **kwargs) + + monkeypatch.setattr(type(repo), "update_l3", racing_update) + merged = _upsert(repo, "C1", "large_amount", 70, alert_id="ALT-2") + + assert calls["n"] >= 2 # 确实走了重试 + assert merged["monitor_tier"] == "high" and merged["risk_score"] == 95 + assert AML_PENDING_TAG in merged["monitor_tags"] + row = _row(engine, "C1") + assert row["monitor_tier"] == "high" and row["risk_score"] == 95 + assert row["last_alert_id"] == "ALT-2" + + +def test_persistent_conflict_raises_not_silent(env, monkeypatch): + """重试耗尽抛错(宁失败不静默覆盖),不留半写状态。""" + repo, engine = env + repo.insert_l3("C1", "watch", 70, {}, [], "ALT-1", datetime.now()) + monkeypatch.setattr(type(repo), "update_l3", lambda self, *a, **k: False) + with pytest.raises(RuntimeError, match="conflicted"): + _upsert(repo, "C1", "large_amount", 70, alert_id="ALT-2") + row = _row(engine, "C1") + assert row["monitor_tier"] == "watch" # 未被静默改写