Files
XingHuo/tests/test_risk_api.py
T
GaoYiYuan_0626 a78e9d4d31 fix: B6 复审收口——测试 Redis 隔离/网关审计口径 platform/死常量清理/回归断言补强
- P2: test_risk_api env 注入 FakePublisher,aml 推送不触真 Redis
- P3: deny/_authz_audit 增 agent_type 参数,simulate 越权审计传 platform
  (与网关放行审计同口径);删 HANDLE_RESULTS 死常量(Literal 单处定义)
- P3: 补 start_date/end_date 过滤+分页 422 边界+suitability api 层审计断言;
  补 app/service/risk/__init__.py
- 挂账: B7 行新增④项(启动期拒绝/引擎工厂覆盖三实例化点/handle 原子性/
  input_guard_log);B9b 行核查单(disclaimer/scan 幂等/Swagger 手测/analyst);
  TODO.md T-30/T-31/T-32 进度同步
2026-09-06 19:21:15 +08:00

402 lines
16 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.
"""risk API 权限矩阵测试(B6 · A-7 处置/仅 aml、A-9 越权 403+audit、debug 头边界)。
TestClient 独立挂 router(main 挂载归 B7);仓储经 monkeypatch 注入 sqlite。
"""
from datetime import datetime, timedelta
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, text
from sqlalchemy.pool import StaticPool
from app.api import risk as risk_api
from app.api.deps import assert_customer_access, permission_denied_handler
from app.api.risk import router as risk_router
from app.repository.core_ro import CoreReadOnlyRepository
from app.repository.risk_repository import RiskRepository
from app.service.risk import alert_service
from app.utils.exceptions import PermissionDenied
class FakePublisher:
"""aml/scan 推送隔离(复审 P2:测试不得依赖本机 Redis 状态)。"""
def __init__(self):
self.messages = []
def publish(self, channel, payload):
self.messages.append((channel, payload))
def _alert(alert_id, customer, atype, status="pending_review", score=70):
return {
"alert_id": alert_id,
"trace_id": f"trc-{alert_id.lower()}",
"customer_id": customer,
"trade_id": None,
"alert_type": atype,
"triggered_rules": ["RISK-001"],
"risk_score": score,
"status": status,
"payload": {"events": [], "customer_context": {}},
}
@pytest.fixture()
def env():
engine = create_engine(
"sqlite://", poolclass=StaticPool, connect_args={"check_same_thread": False}
)
with engine.begin() as conn:
for ddl in [
"""CREATE TABLE core_customer (
customer_id VARCHAR(64) PRIMARY KEY, display_name VARCHAR(128), age INTEGER,
occupation VARCHAR(64), open_date DATE, is_active TINYINT DEFAULT 1)""",
"""CREATE TABLE core_customer_risk (
customer_id VARCHAR(64), risk_code VARCHAR(8), evaluated_at TIMESTAMP)""",
"""CREATE TABLE core_customer_advisor (
advisor_id VARCHAR(64), customer_id VARCHAR(64), rel_status VARCHAR(16))""",
"""CREATE TABLE core_product (
product_id VARCHAR(64) PRIMARY KEY, product_name VARCHAR(128),
min_risk_code VARCHAR(8), product_type VARCHAR(32))""",
"""CREATE TABLE risk_aml_list (
id INTEGER PRIMARY KEY AUTOINCREMENT, list_id VARCHAR(64), list_type VARCHAR(16),
full_name VARCHAR(128), match_threshold REAL, source VARCHAR(64),
list_version VARCHAR(16), effective_date DATE, is_active TINYINT DEFAULT 1,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)""",
"""CREATE TABLE risk_alert (
alert_id VARCHAR(64) PRIMARY KEY, trace_id VARCHAR(64), customer_id VARCHAR(64),
trade_id VARCHAR(64), alert_type VARCHAR(16), triggered_rules TEXT,
risk_score INTEGER, status VARCHAR(24) DEFAULT 'pending_review', payload TEXT,
handler_id VARCHAR(64), handler_result VARCHAR(64), handler_comment VARCHAR(512),
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, handled_at TIMESTAMP)""",
"""CREATE TABLE audit_log (
id INTEGER PRIMARY KEY, trace_id VARCHAR(64), event_type VARCHAR(64),
agent_type VARCHAR(16), actor_id VARCHAR(64), customer_id VARCHAR(64),
rule_id VARCHAR(64), input_summary TEXT, decision VARCHAR(64), risk_score INTEGER,
handler_id VARCHAR(64), handler_result VARCHAR(64), handler_comment VARCHAR(512),
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)""",
"""CREATE TABLE risk_suitability_log (
id INTEGER PRIMARY KEY AUTOINCREMENT, trace_id VARCHAR(64), customer_id VARCHAR(64),
product_id VARCHAR(64), customer_risk_level VARCHAR(8), product_risk_level VARCHAR(8),
is_matched TINYINT, is_blocked TINYINT, block_reason VARCHAR(512),
request_ref VARCHAR(64), profile_l1_version VARCHAR(32),
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)""",
"""CREATE TABLE customer_profile_l3 (
customer_id VARCHAR(64) PRIMARY KEY, monitor_tier VARCHAR(16) NOT NULL,
risk_score INTEGER, score_dimensions TEXT, monitor_tags TEXT,
last_alert_id VARCHAR(64), computed_at TIMESTAMP NOT NULL,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)""",
]:
conn.execute(text(ddl))
conn.execute(
text(
"INSERT INTO core_customer (customer_id, display_name, age, is_active) VALUES"
" ('CUST-1001', '客户·王**', 28, 1), ('CUST-3001', '客户·孙**', 45, 1)"
)
)
conn.execute(
text(
"INSERT INTO core_customer_risk (customer_id, risk_code, evaluated_at) VALUES"
" ('CUST-1001', 'C1', :t), ('CUST-3001', 'C3', :t)"
),
{"t": datetime.now() - timedelta(days=30)},
)
conn.execute(
text(
"INSERT INTO core_customer_advisor VALUES ('ADV-01', 'CUST-3001', 'active')"
)
)
conn.execute(
text(
"INSERT INTO core_product VALUES"
" ('PROD-161725', '科技成长主题', 'R4', 'mixed'),"
" ('PROD-510300', '沪深300指数', 'R3', 'index')"
)
)
conn.execute(
text(
"INSERT INTO risk_aml_list (list_id, list_type, full_name, match_threshold,"
" source, list_version, is_active) VALUES"
" ('SAN-1', 'sanction', '客户·孙**', 0.85, 'mock', 'v1', 1)"
)
)
repo = RiskRepository(engine=engine)
repo.insert_alert(_alert("ALT-E1", "CUST-3001", "large_amount"))
repo.insert_alert(_alert("ALT-A1", "CUST-1001", "aml", score=95))
repo.insert_alert(_alert("ALT-D1", "CUST-3001", "freq_trade", status="confirmed_normal", score=50))
yield repo, engine
engine.dispose()
@pytest.fixture()
def client(env, monkeypatch):
repo, engine = env
monkeypatch.setattr(risk_api, "_repo", lambda: repo)
monkeypatch.setattr(risk_api, "CoreReadOnlyRepository", lambda: CoreReadOnlyRepository(engine=engine))
# 401/越权审计经 deps 内 RiskRepository 兜底构造,统一注入 sqlite(B6 评审 P3-4)
from app.api import deps as deps_mod
monkeypatch.setattr(deps_mod, "RiskRepository", lambda: repo)
# aml 命中会广播 Pub/Sub:注入 fake,不触真 Redis(复审 P2)
monkeypatch.setattr(alert_service, "_publisher", FakePublisher())
app = FastAPI()
app.include_router(risk_router)
app.add_exception_handler(PermissionDenied, permission_denied_handler)
with TestClient(app) as c:
yield c
def _h(role="", actor=""):
return {"X-Debug-Role": role, "X-Debug-Actor": actor} if role else {}
OFFICER = _h("risk_officer", "STAFF-90001")
COMPLIANCE = _h("compliance", "STAFF-40001")
CUST_1001 = _h("customer", "CUST-1001")
ADV_01 = _h("advisor", "ADV-01")
def _counts(engine, table, where="1=1"):
with engine.connect() as conn:
return conn.execute(text(f"SELECT COUNT(*) FROM {table} WHERE {where}")).scalar_one()
# ---------- GET /alerts ----------
def test_officer_sees_all_compliance_forced_aml(client):
r = client.get("/api/risk/alerts", headers=OFFICER)
assert r.status_code == 200
assert r.json()["total"] == 3 # 全量(含已处置)
r = client.get("/api/risk/alerts", headers=COMPLIANCE)
assert r.status_code == 200
body = r.json()
assert body["total"] == 1 and body["items"][0]["alert_type"] == "aml" # A-7 强制过滤
r = client.get("/api/risk/alerts", params={"alert_type": "large_amount"}, headers=COMPLIANCE)
assert r.json()["items"][0]["alert_type"] == "aml" # 传入类型被覆盖
def test_other_roles_cannot_list_alerts(client):
for headers in (CUST_1001, ADV_01):
r = client.get("/api/risk/alerts", headers=headers)
assert r.status_code == 403
def test_alerts_date_filter_and_pagination(client):
"""复审 P3:start_date/end_date 过滤与分页边界回归保护。
sqlite CURRENT_TIMESTAMP 为 UTC,窗口用 utcnow 构造(localtime 收敛挂账 B8)。
"""
now = datetime.utcnow()
r = client.get("/api/risk/alerts", params={"start_date": (now - timedelta(hours=1)).isoformat()}, headers=OFFICER)
assert r.json()["total"] == 3 # 窗口内全命中
r = client.get("/api/risk/alerts", params={"start_date": (now + timedelta(hours=1)).isoformat()}, headers=OFFICER)
assert r.json()["total"] == 0 # 未来起点 → 空
r = client.get("/api/risk/alerts", params={"end_date": (now - timedelta(hours=1)).isoformat()}, headers=OFFICER)
assert r.json()["total"] == 0 # 过去终点 → 空
r = client.get("/api/risk/alerts", params={"page": 2, "page_size": 2}, headers=OFFICER)
body = r.json()
assert body["total"] == 3 and len(body["items"]) == 1 and body["page"] == 2
assert client.get("/api/risk/alerts", params={"page": 0}, headers=OFFICER).status_code == 422
assert client.get("/api/risk/alerts", params={"page_size": 101}, headers=OFFICER).status_code == 422
def test_missing_debug_headers_401(client):
assert client.get("/api/risk/alerts").status_code == 401
# ---------- POST /alerts/{id}/handle ----------
def test_officer_handle_success_state_machine_and_audit(client, env):
repo, engine = env
r = client.post(
"/api/risk/alerts/ALT-E1/handle",
json={"handler_result": "confirmed_suspicious", "handler_comment": "确认可疑"},
headers=OFFICER,
)
assert r.status_code == 200
body = r.json()
assert body["status"] == "confirmed_suspicious" and body["handler_id"] == "STAFF-90001"
assert _counts(engine, "audit_log", "event_type='alert_handle' AND decision='alert_handled'") == 1
# 状态机:已处置单禁止跳改
r2 = client.post(
"/api/risk/alerts/ALT-E1/handle",
json={"handler_result": "confirmed_normal"},
headers=OFFICER,
)
assert r2.status_code == 409
def test_compliance_handle_403(client, env):
repo, _ = env
r = client.post(
"/api/risk/alerts/ALT-E1/handle",
json={"handler_result": "confirmed_suspicious"},
headers=COMPLIANCE,
)
assert r.status_code == 403 # A-7
assert repo.get_alert("ALT-E1")["status"] == "pending_review"
def test_handle_invalid_result_422(client):
r = client.post(
"/api/risk/alerts/ALT-E1/handle",
json={"handler_result": "auto_frozen"}, # 非法枚举(依赖层校验)
headers=OFFICER,
)
assert r.status_code == 422
def test_handle_missing_alert_404(client):
r = client.post(
"/api/risk/alerts/ALT-XXXX/handle",
json={"handler_result": "confirmed_normal"},
headers=OFFICER,
)
assert r.status_code == 404
# ---------- POST /suitability/check(G-01 + A-9) ----------
def test_suitability_check_by_owner_officer_and_assigned_advisor(client, env):
repo, engine = env
r = client.post(
"/api/risk/suitability/check",
json={"customer_id": "CUST-1001", "product_id": "PROD-161725"},
headers=CUST_1001,
)
assert r.status_code == 200
body = r.json()
assert body["blocked"] is True and body["rule_id"] == "SUIT-001" # C1+R4
r = client.post(
"/api/risk/suitability/check",
json={"customer_id": "CUST-3001", "product_id": "PROD-510300"},
headers=OFFICER,
)
assert r.status_code == 200 and r.json()["is_matched"] is True
r = client.post(
"/api/risk/suitability/check",
json={"customer_id": "CUST-3001", "product_id": "PROD-510300"},
headers=ADV_01, # ADV-01 名下 CUST-3001
)
assert r.status_code == 200
# 直调路径每次校验补 api 层审计(评审 P2-3;复审 P3 回归保护)
assert _counts(engine, "audit_log", "event_type='suitability_check'") == 3
assert (
_counts(
engine,
"audit_log",
"event_type='suitability_check' AND decision='suitability_blocked'",
)
== 1
)
def test_a9_customer_other_403_with_audit(client, env):
repo, engine = env
r = client.post(
"/api/risk/suitability/check",
json={"customer_id": "CUST-3001", "product_id": "PROD-510300"},
headers=CUST_1001, # 查他人
)
assert r.status_code == 403
assert "AUTH_403_NOT_OWNER" in r.json()["detail"]
assert _counts(engine, "audit_log", "event_type='authz' AND decision='forbidden'") == 1
def test_a9_advisor_not_assigned_403_with_audit(client, env):
repo, engine = env
r = client.post(
"/api/risk/suitability/check",
json={"customer_id": "CUST-1001", "product_id": "PROD-161725"},
headers=ADV_01, # 非名下
)
assert r.status_code == 403
assert "AUTH_403_NOT_ASSIGNED" in r.json()["detail"]
assert _counts(engine, "audit_log", "event_type='authz' AND decision='forbidden'") == 1
def test_compliance_suitability_check_403(client):
r = client.post(
"/api/risk/suitability/check",
json={"customer_id": "CUST-3001", "product_id": "PROD-510300"},
headers=COMPLIANCE, # 审计角色不在客户业务数据白名单
)
assert r.status_code == 403 and "AUTH_403_SCOPE" in r.json()["detail"]
# ---------- POST /aml/scan ----------
def test_aml_scan_officer_only_with_audit(client, env):
repo, engine = env
r = client.post("/api/risk/aml/scan", headers=OFFICER)
assert r.status_code == 200
body = r.json()
assert body == {"scanned": 2, "hit_customers": 1, "alerts": body["alerts"]} # CUST-3001 命中
assert len(body["alerts"]) == 1
l3 = repo.get_l3("CUST-3001")
assert l3["monitor_tier"] == "high" # scan 命中标记 L3
assert _counts(engine, "audit_log", "event_type='aml_scan' AND decision='scan_completed'") == 1
for headers in (COMPLIANCE, CUST_1001):
assert client.post("/api/risk/aml/scan", headers=headers).status_code == 403
# ---------- 非 dev 环境拒绝 debug 鉴权 ----------
def test_non_dev_rejects_debug_auth(client, monkeypatch):
from app.config.settings import settings
monkeypatch.setattr(settings, "app_env", "production")
with pytest.raises(RuntimeError, match="debug auth disabled"):
client.get("/api/risk/alerts", headers=OFFICER)
# ---------- deps 单元:compliance 不放行客户业务数据 ----------
def test_assert_customer_access_scope_denial(env):
repo, _ = env
from app.api.deps import AuthContext
auth = AuthContext(actor_id="STAFF-40001", roles=["compliance"])
with pytest.raises(PermissionDenied) as ei:
assert_customer_access(
auth, "CUST-1001", core_ro=CoreReadOnlyRepository(engine=env[1]), risk_repo=repo
)
assert ei.value.code == "AUTH_403_SCOPE"
def test_authz_denials_are_audited(client, env):
"""B6 评审 P1-1:403/401 全部留痕(AUTH_403_ROLE / AUTH_401_MISSING_DEBUG_HEADERS)。"""
repo, engine = env
assert client.get("/api/risk/alerts").status_code == 401 # 无 debug 头
assert client.post("/api/risk/aml/scan", headers=ADV_01).status_code == 403
assert client.post(
"/api/risk/alerts/ALT-E1/handle",
json={"handler_result": "confirmed_normal"},
headers=COMPLIANCE,
).status_code == 403
assert _counts(engine, "audit_log", "decision='unauthenticated' AND input_summary LIKE '%AUTH_401%'") == 1
assert _counts(engine, "audit_log", "decision='forbidden' AND input_summary LIKE '%AUTH_403_ROLE%'") == 2
def test_multi_role_combinations_fixed_behavior(client):
"""B6 评审 P3-1:多角色口径固化——并集权限、归属 fail-closed。"""
# customer+compliance:借 compliance 角色看 aml 台账(审计类跨客户,固化允许)
r = client.get("/api/risk/alerts", headers=_h("customer,compliance", "CUST-1001"))
assert r.status_code == 200 and r.json()["items"][0]["alert_type"] == "aml"
# customer+advisor 查非本人客户:customer 分支 deny 即终止(最窄范围 fail-closed)
r = client.post(
"/api/risk/suitability/check",
json={"customer_id": "CUST-1001", "product_id": "PROD-161725"},
headers=_h("customer,advisor", "CUST-3001"),
)
assert r.status_code == 403 and "AUTH_403_NOT_OWNER" in r.json()["detail"]