Files
water-management-system/tests/security/test_security.py
T

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()