Files
XingHuo/tests/test_risk_api.py
T

327 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.
"""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
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))
monkeypatch.setattr(risk_api, "CoreReadOnlyRepository", lambda: CoreReadOnlyRepository(engine=engine))
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_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):
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
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", risk_repo=repo)
assert ei.value.code == "AUTH_403_SCOPE"