187 lines
7.2 KiB
Python
187 lines
7.2 KiB
Python
"""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()
|