Files
XingHuo/tests/test_auth_jwt.py
T

290 lines
11 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.
"""T-01 JWT 鉴权(auth_service 验签/吊销 + deps 工厂 + Agent 准入矩阵)。
通道矩阵:JWT(全环境,强制 X-Agent-Type 交叉校验)/ dev debug 头兜底 /
非 dev 无 Bearer 401 留痕。401/403 全部经 deps 审计(P-05),sqlite 断言。
RS256 用 cryptography 现场生成密钥对(生产形态走通);HS256 走 dev secret。
"""
from __future__ import annotations
import time
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
from jose import jwt as jose_jwt
from sqlalchemy import text
from _ddl import create_sqlite_engine
from app.api import deps as deps_mod
from app.api import risk as risk_api
from app.config.settings import settings
from app.main import app
from app.repository.risk_repository import RiskRepository
from app.service import auth_service
from app.service.risk import redis_gateway
class FakeGateway:
"""Redis fake:吊销黑名单语义(exists/set_ex 记录)。"""
def __init__(self):
self.revoked: set[str] = set()
self.messages: list = []
self.deletes: list = []
def publish(self, channel, payload):
self.messages.append((channel, payload))
def delete(self, *keys):
self.deletes.append(keys)
def exists(self, key: str) -> bool:
return key in self.revoked
def set_ex(self, key, value, ttl):
self.revoked.add(key)
@pytest.fixture()
def env(monkeypatch):
engine = create_sqlite_engine()
repo = RiskRepository(engine=engine)
monkeypatch.setattr(risk_api, "_repo", lambda: repo)
monkeypatch.setattr(deps_mod, "RiskRepository", lambda: repo)
monkeypatch.setattr(redis_gateway, "_gateway", FakeGateway())
yield {"client": TestClient(app, raise_server_exceptions=False), "repo": repo, "engine": engine}
engine.dispose()
def _bearer(tok: str, agent: str = "risk") -> dict:
return {"Authorization": f"Bearer {tok}", "X-Agent-Type": agent}
def _audit_rows(engine, decision: str) -> list:
with engine.connect() as conn:
return conn.execute(
text("SELECT event_type, actor_id, decision, input_summary FROM audit_log WHERE decision = :d"),
{"d": decision},
).mappings().all()
# ---------- JWT 通道基本链路 ----------
def test_jwt_risk_officer_access_alerts(env):
tok = auth_service.issue_dev_token(sub="STAFF-30001", roles=["risk_officer"])
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 200
assert r.json()["disclaimer"]
def test_jwt_missing_agent_type_401_with_audit(env):
tok = auth_service.issue_dev_token(sub="STAFF-30001", roles=["risk_officer"])
r = env["client"].get("/api/risk/alerts", headers={"Authorization": f"Bearer {tok}"})
assert r.status_code == 401
assert r.json()["error_code"] == "AUTH_401_MISSING_AGENT_TYPE"
assert _audit_rows(env["engine"], "unauthenticated")
def test_jwt_invalid_agent_type_400(env):
tok = auth_service.issue_dev_token(sub="STAFF-30001", roles=["risk_officer"])
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok, agent="backoffice"))
assert r.status_code == 400 and r.json()["error_code"] == "BAD_REQUEST"
def test_jwt_customer_token_cannot_enter_risk_agent(env):
"""手册 §5.4:customer token 只能进 customer agent(交叉校验 403 + 审计)。"""
tok = auth_service.issue_dev_token(
sub="CUST-9527", roles=["customer"], token_type="customer", customer_id="CUST-9527"
)
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 403
assert r.json()["error_code"] == "AUTH_403_AGENT_MISMATCH"
assert _audit_rows(env["engine"], "forbidden")
def test_jwt_advisor_cannot_enter_risk_agent(env):
tok = auth_service.issue_dev_token(sub="STAFF-10086", roles=["advisor"])
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 403 and r.json()["error_code"] == "AUTH_403_AGENT_MISMATCH"
# ---------- 验签失败矩阵(统一 401 INVALID_TOKEN + 留痕) ----------
def test_jwt_tampered_signature_rejected(env):
tok = auth_service.issue_dev_token(sub="STAFF-30001", roles=["risk_officer"])
head, payload, sig = tok.split(".")
r = env["client"].get("/api/risk/alerts", headers=_bearer(f"{head}.{payload}.AAAA{sig[4:]}"))
assert r.status_code == 401 and r.json()["error_code"] == "AUTH_401_INVALID_TOKEN"
assert _audit_rows(env["engine"], "unauthenticated")
def test_jwt_expired_rejected(env):
now = int(time.time())
payload = {
"iss": settings.jwt_issuer, "sub": "STAFF-30001", "aud": settings.jwt_audience,
"exp": now - 10, "iat": now - 100, "jti": "jti-exp", "token_type": "staff",
"roles": ["risk_officer"], "tenant_id": "TENANT-001",
}
tok = jose_jwt.encode(payload, settings.jwt_dev_secret, algorithm="HS256")
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 401 and r.json()["error_code"] == "AUTH_401_INVALID_TOKEN"
def test_jwt_wrong_issuer_rejected(env):
now = int(time.time())
payload = {
"iss": "https://evil.example", "sub": "STAFF-30001", "aud": settings.jwt_audience,
"exp": now + 600, "iat": now, "jti": "jti-iss", "token_type": "staff",
"roles": ["risk_officer"], "tenant_id": "TENANT-001",
}
tok = jose_jwt.encode(payload, settings.jwt_dev_secret, algorithm="HS256")
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 401
def test_jwt_customer_id_mismatch_rejected(env):
"""手册 §4.2:customer token 的 customer_id 必须等于 sub(防冒名)。"""
tok = auth_service.issue_dev_token(
sub="CUST-9527", roles=["customer"], token_type="customer", customer_id="CUST-1001"
)
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 401 and r.json()["error_code"] == "AUTH_401_INVALID_TOKEN"
def test_jwt_missing_required_claim_rejected(env):
now = int(time.time())
payload = { # 缺 tenant_id(手册 §4.2 必填)
"iss": settings.jwt_issuer, "sub": "STAFF-30001", "aud": settings.jwt_audience,
"exp": now + 600, "iat": now, "jti": "jti-noclaim", "token_type": "staff",
"roles": ["risk_officer"],
}
tok = jose_jwt.encode(payload, settings.jwt_dev_secret, algorithm="HS256")
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 401
# ---------- 吊销(jti 黑名单) ----------
def test_jwt_revoked_jti_rejected(env):
tok = auth_service.issue_dev_token(sub="STAFF-30001", roles=["risk_officer"])
claims = auth_service.verify_token(tok) # 先取 jti
redis_gateway._gateway.revoked.add(f"auth:revoked:{claims.jti}")
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 401 and r.json()["error_code"] == "AUTH_401_REVOKED"
def test_revocation_fail_open_when_redis_down(monkeypatch):
"""Redis 不可用 fail-open(redis_gateway 总口径;TTL 为权威兜底)。"""
monkeypatch.setattr(redis_gateway, "_gateway", None)
monkeypatch.setattr(
redis_gateway.RedisGateway, "exists", lambda self, key: (_ for _ in ()).throw(ConnectionError())
)
assert auth_service.is_revoked("jti-x") is False
# ---------- debug 头兜底(dev)与非 dev 行为 ----------
def test_dev_debug_headers_still_work(env):
r = env["client"].get(
"/api/risk/alerts", headers={"X-Debug-Role": "risk_officer", "X-Debug-Actor": "STAFF-30001"}
)
assert r.status_code == 200 # B6 演示 SOP 通道不回归
def test_non_dev_without_bearer_401(env, monkeypatch):
monkeypatch.setattr(settings, "app_env", "production")
r = env["client"].get(
"/api/risk/alerts", headers={"X-Debug-Role": "risk_officer", "X-Debug-Actor": "STAFF-30001"}
)
assert r.status_code == 401 and r.json()["error_code"] == "AUTH_401_MISSING_BEARER"
assert _audit_rows(env["engine"], "unauthenticated")
def test_non_dev_jwt_still_accepted(env, monkeypatch):
"""JWT 是唯一生产通道:非 dev 下有效 JWT + 交叉头照常放行。"""
monkeypatch.setattr(settings, "app_env", "production")
tok = auth_service.issue_dev_token(sub="STAFF-30001", roles=["risk_officer"])
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 200
def test_non_dev_hs256_secret_refused_inline(env, monkeypatch):
"""HS256 上生产防护:jwt_ready 必须给出拒绝原因(lifespan 消费)。"""
monkeypatch.setattr(settings, "app_env", "production")
monkeypatch.setattr(settings, "jwt_public_key_path", "")
assert auth_service.jwt_ready() is not None
# ---------- RS256(生产形态) ----------
_RSA_CACHE: dict = {}
def _rsa_keypair() -> tuple[bytes, bytes]:
if not _RSA_CACHE:
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
_RSA_CACHE["priv"] = key.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
)
_RSA_CACHE["pub"] = key.public_key().public_bytes(
serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo
)
return _RSA_CACHE["priv"], _RSA_CACHE["pub"]
def test_rs256_public_key_roundtrip(env, monkeypatch, tmp_path: Path):
priv, pub = _rsa_keypair()
key_file = tmp_path / "idp_pub.pem"
key_file.write_bytes(pub)
monkeypatch.setattr(settings, "jwt_public_key_path", str(key_file))
assert auth_service.jwt_ready() is None
now = int(time.time())
payload = {
"iss": settings.jwt_issuer, "sub": "STAFF-30001", "aud": settings.jwt_audience,
"exp": now + 600, "iat": now, "jti": "jti-rs256", "token_type": "staff",
"roles": ["risk_officer"], "tenant_id": "TENANT-001",
"permissions": ["risk:alert:write"],
}
tok = jose_jwt.encode(payload, priv, algorithm="RS256", headers={"kid": "2026-09-key-1"})
r = env["client"].get("/api/risk/alerts", headers=_bearer(tok))
assert r.status_code == 200
# HS256 token 在 RS256 模式下必须被拒(算法白名单,防 alg 混淆)
hs_tok = auth_service.issue_dev_token(sub="STAFF-30001", roles=["risk_officer"])
r2 = env["client"].get("/api/risk/alerts", headers=_bearer(hs_tok))
assert r2.status_code == 401
# ---------- assert_agent_access 纯函数矩阵 ----------
def test_agent_access_matrix():
from app.api.deps import AGENT_ACCESS_MATRIX, assert_agent_access
for agent, rule in AGENT_ACCESS_MATRIX.items():
for tt in rule["token_types"]:
for role in rule["roles"]:
assert_agent_access(
deps_mod.AuthContext(actor_id="X", roles=[role], token_type=tt), agent
)
with pytest.raises(deps_mod.PermissionDenied):
assert_agent_access(
deps_mod.AuthContext(actor_id="X", roles=["ops"], token_type="staff"), "analyst"
)
with pytest.raises(deps_mod.PermissionDenied):
assert_agent_access(
deps_mod.AuthContext(actor_id="X", roles=["risk_officer"], token_type="customer"), "risk"
)