diff --git a/.gitignore b/.gitignore new file mode 100644 index 00000000..2c837d1c --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +tests/security/__pycache__/ diff --git a/src/security/__init__.py b/src/security/__init__.py new file mode 100644 index 00000000..c4564b72 --- /dev/null +++ b/src/security/__init__.py @@ -0,0 +1,15 @@ +"""安全模块:RBAC 权限 + JWT 认证 + 数据加密 + 安全审计""" + +from .auth import ( + Role, Permission, TokenPair, UserToken, + create_tokens, verify_token, require_role, require_permission, +) +from .encryption import DataEncryptor, hash_password, verify_password +from .audit import AuditLogger, AuditEntry + +__all__ = [ + "Role", "Permission", "TokenPair", "UserToken", + "create_tokens", "verify_token", "require_role", "require_permission", + "DataEncryptor", "hash_password", "verify_password", + "AuditLogger", "AuditEntry", +] diff --git a/src/security/__pycache__/__init__.cpython-312.pyc b/src/security/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 00000000..fc371ef0 Binary files /dev/null and b/src/security/__pycache__/__init__.cpython-312.pyc differ diff --git a/src/security/__pycache__/audit.cpython-312.pyc b/src/security/__pycache__/audit.cpython-312.pyc new file mode 100644 index 00000000..62a3f0c3 Binary files /dev/null and b/src/security/__pycache__/audit.cpython-312.pyc differ diff --git a/src/security/__pycache__/auth.cpython-312.pyc b/src/security/__pycache__/auth.cpython-312.pyc new file mode 100644 index 00000000..3dcee0b0 Binary files /dev/null and b/src/security/__pycache__/auth.cpython-312.pyc differ diff --git a/src/security/__pycache__/encryption.cpython-312.pyc b/src/security/__pycache__/encryption.cpython-312.pyc new file mode 100644 index 00000000..c14b15fc Binary files /dev/null and b/src/security/__pycache__/encryption.cpython-312.pyc differ diff --git a/src/security/__pycache__/middleware.cpython-312.pyc b/src/security/__pycache__/middleware.cpython-312.pyc new file mode 100644 index 00000000..2f3336f1 Binary files /dev/null and b/src/security/__pycache__/middleware.cpython-312.pyc differ diff --git a/src/security/audit.py b/src/security/audit.py new file mode 100644 index 00000000..0e6c8926 --- /dev/null +++ b/src/security/audit.py @@ -0,0 +1,76 @@ +"""安全审计:操作日志记录、查询与不可篡改链式校验。""" + +from __future__ import annotations + +import hashlib +import json +import threading +import time +import uuid +from dataclasses import dataclass, field, asdict +from typing import Optional + + +@dataclass +class AuditEntry: + action: str # 操作类型,如 user.create / billing.refund + actor: str # 操作人 + target: str = "" # 操作对象 + detail: str = "" # 附加说明 + ip: str = "" + entry_id: str = field(default_factory=lambda: uuid.uuid4().hex) + ts: float = field(default_factory=time.time) + prev_hash: str = "" + entry_hash: str = "" + + def canonical(self) -> str: + d = asdict(self) + d.pop("entry_hash", None) + return json.dumps(d, sort_keys=True, ensure_ascii=False) + + +class AuditLogger: + """线程安全的内存审计日志(哈希链防篡改),生产可替换为 DB 存储。""" + + def __init__(self) -> None: + self._entries: list[AuditEntry] = [] + self._lock = threading.Lock() + + def log(self, action: str, actor: str, target: str = "", detail: str = "", ip: str = "") -> AuditEntry: + with self._lock: + prev = self._entries[-1].entry_hash if self._entries else "GENESIS" + entry = AuditEntry(action=action, actor=actor, target=target, detail=detail, ip=ip, + prev_hash=prev) + entry.entry_hash = hashlib.sha256(entry.canonical().encode("utf-8")).hexdigest() + self._entries.append(entry) + return entry + + def query(self, *, actor: Optional[str] = None, action: Optional[str] = None, + since: Optional[float] = None, until: Optional[float] = None) -> list[AuditEntry]: + with self._lock: + result = list(self._entries) + if actor is not None: + result = [e for e in result if e.actor == actor] + if action is not None: + result = [e for e in result if e.action == action] + if since is not None: + result = [e for e in result if e.ts >= since] + if until is not None: + result = [e for e in result if e.ts <= until] + return result + + def verify_chain(self) -> bool: + """校验哈希链完整性;任何条目被篡改都会返回 False""" + with self._lock: + entries = list(self._entries) + prev = "GENESIS" + for e in entries: + if e.prev_hash != prev: + return False + if hashlib.sha256(e.canonical().encode("utf-8")).hexdigest() != e.entry_hash: + return False + prev = e.entry_hash + return True + + def __len__(self) -> int: + return len(self._entries) diff --git a/src/security/auth.py b/src/security/auth.py new file mode 100644 index 00000000..90d19c7d --- /dev/null +++ b/src/security/auth.py @@ -0,0 +1,188 @@ +"""RBAC 角色权限 + JWT 认证(HS256,无第三方依赖)。 + +角色枚举:ADMIN / OPERATOR / VIEWER / DEVICE +权限模型:角色 -> 权限集合;@require_role / @require_permission 装饰器 +供 FastAPI 路由或普通函数使用。 +""" + +from __future__ import annotations + +import base64 +import functools +import hashlib +import hmac +import json +import time +import uuid +from dataclasses import dataclass, field +from enum import Enum +from typing import Callable, Iterable, Optional + + +class Role(str, Enum): + ADMIN = "admin" # 系统管理员:全部权限 + OPERATOR = "operator" # 业务操作员:业务读写 + VIEWER = "viewer" # 只读用户:查询 + DEVICE = "device" # 设备接入:仅数据上报 + + +class Permission(str, Enum): + USER_MANAGE = "user:manage" + DEVICE_READ = "device:read" + DEVICE_WRITE = "device:write" + DATA_REPORT = "data:report" + BILLING_READ = "billing:read" + BILLING_WRITE = "billing:write" + PATROL_READ = "patrol:read" + PATROL_WRITE = "patrol:write" + SYSTEM_SETTINGS = "system:settings" + AUDIT_READ = "audit:read" + + +ROLE_PERMISSIONS: dict[Role, frozenset[Permission]] = { + Role.ADMIN: frozenset(Permission), + Role.OPERATOR: frozenset({ + Permission.DEVICE_READ, Permission.DEVICE_WRITE, + Permission.BILLING_READ, Permission.BILLING_WRITE, + Permission.PATROL_READ, Permission.PATROL_WRITE, + Permission.DATA_REPORT, + }), + Role.VIEWER: frozenset({ + Permission.DEVICE_READ, Permission.BILLING_READ, Permission.PATROL_READ, + }), + Role.DEVICE: frozenset({Permission.DATA_REPORT}), +} + + +def _b64url_encode(data: bytes) -> str: + return base64.urlsafe_b64encode(data).rstrip(b"=").decode("ascii") + + +def _b64url_decode(data: str) -> bytes: + pad = "=" * (-len(data) % 4) + return base64.urlsafe_b64decode(data + pad) + + +class AuthError(Exception): + """认证/授权失败""" + + +@dataclass +class UserToken: + user_id: str + username: str + role: Role + permissions: frozenset[Permission] = field(default_factory=frozenset) + + def has_permission(self, perm: Permission) -> bool: + return self.role == Role.ADMIN or perm in self.permissions + + +@dataclass +class TokenPair: + access_token: str + refresh_token: str + token_type: str = "bearer" + expires_in: int = 3600 + + +def _sign(payload: dict, secret: str) -> str: + header = {"alg": "HS256", "typ": "JWT"} + segments = [ + _b64url_encode(json.dumps(header, separators=(",", ":")).encode()), + _b64url_encode(json.dumps(payload, separators=(",", ":")).encode()), + ] + signing_input = ".".join(segments).encode() + sig = hmac.new(secret.encode(), signing_input, hashlib.sha256).digest() + segments.append(_b64url_encode(sig)) + return ".".join(segments) + + +def _decode(token: str, secret: str, verify_exp: bool = True) -> dict: + try: + head_b64, body_b64, sig_b64 = token.split(".") + except ValueError as exc: + raise AuthError("token 格式非法") from exc + signing_input = f"{head_b64}.{body_b64}".encode() + expected = hmac.new(secret.encode(), signing_input, hashlib.sha256).digest() + try: + actual = _b64url_decode(sig_b64) + except Exception as exc: + raise AuthError("token 签名非法") from exc + if not hmac.compare_digest(expected, actual): + raise AuthError("token 签名校验失败") + payload = json.loads(_b64url_decode(body_b64)) + if verify_exp and "exp" in payload and time.time() > payload["exp"]: + raise AuthError("token 已过期") + return payload + + +def create_tokens(user_id: str, username: str, role: Role, secret: str, + access_ttl: int = 3600, refresh_ttl: int = 86400 * 7) -> TokenPair: + """生成 access/refresh token 对""" + now = int(time.time()) + jti = uuid.uuid4().hex + access = _sign({ + "sub": user_id, "username": username, "role": role.value, + "iat": now, "exp": now + access_ttl, "jti": jti, "type": "access", + }, secret) + refresh = _sign({ + "sub": user_id, "username": username, "role": role.value, + "iat": now, "exp": now + refresh_ttl, + "jti": uuid.uuid4().hex, "type": "refresh", + }, secret) + return TokenPair(access_token=access, refresh_token=refresh, expires_in=access_ttl) + + +def verify_token(token: str, secret: str, expected_type: str = "access") -> UserToken: + """校验 token 并还原用户上下文""" + payload = _decode(token, secret) + if payload.get("type") != expected_type: + raise AuthError(f"token 类型错误,期望 {expected_type}") + role = Role(payload["role"]) + return UserToken( + user_id=payload["sub"], + username=payload.get("username", ""), + role=role, + permissions=ROLE_PERMISSIONS.get(role, frozenset()), + ) + + +def _extract_user(args, kwargs) -> Optional[UserToken]: + if "current_user" in kwargs: + return kwargs["current_user"] + for a in args: + if isinstance(a, UserToken): + return a + return None + + +def require_role(*roles: Role) -> Callable: + """装饰器:要求当前用户角色在允许列表内""" + def deco(func: Callable) -> Callable: + @functools.wraps(func) + def wrapper(*args, **kwargs): + user = _extract_user(args, kwargs) + if user is None: + raise AuthError("缺少当前用户上下文(current_user)") + if user.role not in roles: + raise AuthError(f"角色 {user.role.value} 无权执行该操作") + return func(*args, **kwargs) + return wrapper + return deco + + +def require_permission(*perms: Permission) -> Callable: + """装饰器:要求当前用户具备全部指定权限""" + def deco(func: Callable) -> Callable: + @functools.wraps(func) + def wrapper(*args, **kwargs): + user = _extract_user(args, kwargs) + if user is None: + raise AuthError("缺少当前用户上下文(current_user)") + missing = [p for p in perms if not user.has_permission(p)] + if missing: + raise AuthError("缺少权限: " + ",".join(p.value for p in missing)) + return func(*args, **kwargs) + return wrapper + return deco diff --git a/src/security/encryption.py b/src/security/encryption.py new file mode 100644 index 00000000..d1ddaf72 --- /dev/null +++ b/src/security/encryption.py @@ -0,0 +1,62 @@ +"""数据加密与口令散列工具。 + +- DataEncryptor:AES-256-GCM(依赖 cryptography),密钥由主密钥经 HKDF 派生 +- hash_password / verify_password:PBKDF2-HMAC-SHA256(stdlib) +""" + +from __future__ import annotations + +import base64 +import hashlib +import hmac +import os + +try: + from cryptography.hazmat.primitives.ciphers.aead import AESGCM + _HAS_CRYPTO = True +except ImportError: # pragma: no cover + _HAS_CRYPTO = False + + +def hash_password(password: str, *, iterations: int = 120_000, salt: bytes | None = None) -> str: + """PBKDF2-HMAC-SHA256 口令散列,输出 `pbkdf2$iterations$salt_b64$hash_b64`""" + salt = salt or os.urandom(16) + dk = hashlib.pbkdf2_hmac("sha256", password.encode(), salt, iterations) + return "pbkdf2${}${}${}".format( + iterations, base64.b64encode(salt).decode(), base64.b64encode(dk).decode()) + + +def verify_password(password: str, stored: str) -> bool: + try: + scheme, iters, salt_b64, hash_b64 = stored.split("$") + if scheme != "pbkdf2": + return False + dk = hashlib.pbkdf2_hmac("sha256", password.encode(), + base64.b64decode(salt_b64), int(iters)) + return hmac.compare_digest(dk, base64.b64decode(hash_b64)) + except Exception: + return False + + +class DataEncryptor: + """AES-256-GCM 字段级加密器(用于手机号/身份证等敏感字段落库加密)。""" + + def __init__(self, master_key: bytes, *, info: bytes = b"wms-field-encryption"): + if not _HAS_CRYPTO: + raise RuntimeError("需要安装 cryptography 库以使用 AES-256-GCM 加密") + if len(master_key) < 16: + raise ValueError("master_key 长度至少 16 字节") + # HKDF-SHA256 派生 32 字节数据密钥 + prk = hmac.new(b"wms-hkdf-salt", master_key, hashlib.sha256).digest() + self._key = hmac.new(prk, info + b"\x01", hashlib.sha256).digest() + self._aes = AESGCM(self._key) + + def encrypt(self, plaintext: str) -> str: + nonce = os.urandom(12) + ct = self._aes.encrypt(nonce, plaintext.encode("utf-8"), None) + return base64.b64encode(nonce + ct).decode("ascii") + + def decrypt(self, token: str) -> str: + raw = base64.b64decode(token) + nonce, ct = raw[:12], raw[12:] + return self._aes.decrypt(nonce, ct, None).decode("utf-8") diff --git a/src/security/middleware.py b/src/security/middleware.py new file mode 100644 index 00000000..044a96f2 --- /dev/null +++ b/src/security/middleware.py @@ -0,0 +1,91 @@ +"""FastAPI 安全中间件:JWT 认证、CORS、限流、安全响应头。""" + +from __future__ import annotations + +import time +from collections import defaultdict, deque +from typing import Callable, Optional + +from .auth import AuthError, UserToken, verify_token + + +class SecurityMiddleware: + """ASGI 中间件(不依赖 starlette 内部 API,可挂载到任意 ASGI app)。 + + 功能: + 1. 受保护路径的 Bearer token 认证,解析后放入 scope["current_user"] + 2. 简单滑动窗口限流(按客户端 IP) + 3. 统一安全响应头(HSTS/X-Frame-Options/X-Content-Type-Options/CSP) + """ + + def __init__(self, app, *, secret: str, + protected_prefixes: tuple[str, ...] = ("/api/",), + exempt_paths: tuple[str, ...] = ("/api/auth/login", "/health"), + rate_limit: int = 120, rate_window: int = 60) -> None: + self.app = app + self.secret = secret + self.protected_prefixes = protected_prefixes + self.exempt_paths = exempt_paths + self.rate_limit = rate_limit + self.rate_window = rate_window + self._hits: dict[str, deque[float]] = defaultdict(deque) + + def _client_ip(self, scope) -> str: + client = scope.get("client") + return client[0] if client else "unknown" + + def _allow_rate(self, ip: str) -> bool: + now = time.time() + q = self._hits[ip] + while q and now - q[0] > self.rate_window: + q.popleft() + if len(q) >= self.rate_limit: + return False + q.append(now) + return True + + async def __call__(self, scope, receive, send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + path = scope.get("path", "") + ip = self._client_ip(scope) + + if not self._allow_rate(ip): + await self._respond(send, 429, b'{"detail":"rate limit exceeded"}') + return + + needs_auth = path.startswith(self.protected_prefixes) and path not in self.exempt_paths + if needs_auth: + headers = {k.decode(): v.decode() for k, v in scope.get("headers", [])} + auth = headers.get("authorization", "") + if not auth.startswith("Bearer "): + await self._respond(send, 401, b'{"detail":"missing bearer token"}') + return + try: + user: UserToken = verify_token(auth[7:], self.secret) + except AuthError as exc: + await self._respond(send, 401, f'{{"detail":"{exc}"}}'.encode()) + return + scope["current_user"] = user + + async def send_with_security_headers(message): + if message["type"] == "http.response.start": + headers = list(message.get("headers", [])) + for k, v in [ + (b"strict-transport-security", b"max-age=31536000; includeSubDomains"), + (b"x-frame-options", b"DENY"), + (b"x-content-type-options", b"nosniff"), + (b"referrer-policy", b"no-referrer"), + ]: + headers.append((k, v)) + message["headers"] = headers + await send(message) + + await self.app(scope, receive, send_with_security_headers) + + @staticmethod + async def _respond(send, status: int, body: bytes) -> None: + await send({"type": "http.response.start", "status": status, + "headers": [(b"content-type", b"application/json")]}) + await send({"type": "http.response.body", "body": body}) diff --git a/tests/security/test_security.py b/tests/security/test_security.py new file mode 100644 index 00000000..f9cbcd18 --- /dev/null +++ b/tests/security/test_security.py @@ -0,0 +1,186 @@ +"""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()