"""src/security 单元测试:JWT/RBAC/加密/审计/中间件(18 个用例)""" import time import unittest from src.security.auth import ( AuthError, Permission, Role, UserToken, create_tokens, verify_token, require_role, require_permission, ) from src.security.encryption import DataEncryptor, hash_password, verify_password from src.security.audit import AuditLogger from src.security.middleware import SecurityMiddleware SECRET = "unit-test-secret-key-32bytes-long!" class TestJwtAuth(unittest.TestCase): def test_create_and_verify_access_token(self): pair = create_tokens("u1", "zhangsan", Role.OPERATOR, SECRET) user = verify_token(pair.access_token, SECRET) self.assertEqual(user.user_id, "u1") self.assertEqual(user.role, Role.OPERATOR) self.assertEqual(pair.token_type, "bearer") def test_refresh_token_type(self): pair = create_tokens("u1", "zhangsan", Role.ADMIN, SECRET) user = verify_token(pair.refresh_token, SECRET, expected_type="refresh") self.assertEqual(user.user_id, "u1") with self.assertRaises(AuthError): verify_token(pair.refresh_token, SECRET) # 类型不匹配 def test_expired_token_rejected(self): pair = create_tokens("u1", "zhangsan", Role.VIEWER, SECRET, access_ttl=-1) with self.assertRaises(AuthError): verify_token(pair.access_token, SECRET) def test_tampered_token_rejected(self): pair = create_tokens("u1", "zhangsan", Role.VIEWER, SECRET) bad = pair.access_token[:-2] + "xx" with self.assertRaises(AuthError): verify_token(bad, SECRET) def test_wrong_secret_rejected(self): pair = create_tokens("u1", "zhangsan", Role.VIEWER, SECRET) with self.assertRaises(AuthError): verify_token(pair.access_token, "another-secret") class TestRBAC(unittest.TestCase): def _user(self, role): return UserToken("u", "n", role, { Role.ADMIN: frozenset(Permission), Role.OPERATOR: frozenset({Permission.BILLING_READ, Permission.BILLING_WRITE}), Role.VIEWER: frozenset({Permission.BILLING_READ}), Role.DEVICE: frozenset({Permission.DATA_REPORT}), }[role]) def test_admin_has_all_permissions(self): self.assertTrue(self._user(Role.ADMIN).has_permission(Permission.USER_MANAGE)) def test_viewer_cannot_write(self): self.assertFalse(self._user(Role.VIEWER).has_permission(Permission.BILLING_WRITE)) def test_require_role_allows(self): @require_role(Role.ADMIN, Role.OPERATOR) def op(current_user=None): return "ok" self.assertEqual(op(current_user=self._user(Role.OPERATOR)), "ok") def test_require_role_denies(self): @require_role(Role.ADMIN) def op(current_user=None): return "ok" with self.assertRaises(AuthError): op(current_user=self._user(Role.VIEWER)) def test_require_permission_denies_missing(self): @require_permission(Permission.BILLING_WRITE) def op(current_user=None): return "ok" with self.assertRaises(AuthError): op(current_user=self._user(Role.VIEWER)) self.assertEqual(op(current_user=self._user(Role.OPERATOR)), "ok") class TestEncryption(unittest.TestCase): def test_password_hash_roundtrip(self): stored = hash_password("S3cret!") self.assertTrue(verify_password("S3cret!", stored)) self.assertFalse(verify_password("wrong", stored)) def test_password_hash_unique_salt(self): self.assertNotEqual(hash_password("same"), hash_password("same")) def test_aesgcm_roundtrip(self): enc = DataEncryptor(b"master-key-for-testing-32b") token = enc.encrypt("13800138000") self.assertEqual(enc.decrypt(token), "13800138000") def test_aesgcm_tamper_detected(self): enc = DataEncryptor(b"master-key-for-testing-32b") token = enc.encrypt("sensitive") import base64 raw = bytearray(base64.b64decode(token)); raw[-1] ^= 1 with self.assertRaises(Exception): enc.decrypt(base64.b64encode(bytes(raw)).decode()) def test_short_master_key_rejected(self): with self.assertRaises(ValueError): DataEncryptor(b"short") class TestAudit(unittest.TestCase): def test_log_and_query(self): log = AuditLogger() log.log("user.create", "admin", target="u100") log.log("billing.refund", "operator", target="r9", ip="10.0.0.1") self.assertEqual(len(log), 2) self.assertEqual(len(log.query(actor="admin")), 1) self.assertEqual(log.query(action="billing.refund")[0].ip, "10.0.0.1") def test_hash_chain_integrity(self): log = AuditLogger() for i in range(5): log.log(f"op.{i}", "tester") self.assertTrue(log.verify_chain()) def test_hash_chain_tamper_detected(self): log = AuditLogger() log.log("a", "x"); log.log("b", "x") log._entries[0].detail = "tampered" self.assertFalse(log.verify_chain()) class TestMiddleware(unittest.IsolatedAsyncioTestCase): def _make_scope(self, path="/api/devices", token=None): headers = [] if token: headers.append((b"authorization", f"Bearer {token}".encode())) return {"type": "http", "path": path, "headers": headers, "client": ("127.0.0.1", 12345)} @staticmethod def _send_collector(sent): async def _send(message): sent.append(message) return _send async def test_missing_token_401(self): async def app(scope, receive, send): raise AssertionError("不应进入业务") mw = SecurityMiddleware(app, secret=SECRET) sent = [] await mw(self._make_scope(), None, self._send_collector(sent)) self.assertEqual(sent[0]["status"], 401) async def test_valid_token_passes_and_headers(self): async def app(scope, receive, send): assert "current_user" in scope await send({"type": "http.response.start", "status": 200, "headers": []}) await send({"type": "http.response.body", "body": b"ok"}) mw = SecurityMiddleware(app, secret=SECRET) pair = create_tokens("u1", "op", Role.OPERATOR, SECRET) sent = [] await mw(self._make_scope(token=pair.access_token), None, self._send_collector(sent)) self.assertEqual(sent[0]["status"], 200) hdr_keys = {k for k, _ in sent[0]["headers"]} self.assertIn(b"strict-transport-security", hdr_keys) self.assertIn(b"x-frame-options", hdr_keys) async def test_rate_limit(self): async def app(scope, receive, send): await send({"type": "http.response.start", "status": 200, "headers": []}) await send({"type": "http.response.body", "body": b"ok"}) mw = SecurityMiddleware(app, secret=SECRET, rate_limit=3, rate_window=60) pair = create_tokens("u1", "op", Role.OPERATOR, SECRET) statuses = [] for _ in range(5): sent = [] await mw(self._make_scope(token=pair.access_token), None, self._send_collector(sent)) statuses.append(sent[0]["status"]) self.assertEqual(statuses.count(429), 2) if __name__ == "__main__": unittest.main()