from __future__ import annotations import base64 import hashlib import hmac import json import time from typing import Any from fastapi import Request from server.auth.context import IdentityContext from server.auth.providers import AuthError, AuthProvider TEST_USERS: tuple[dict[str, Any], ...] = ( {"id": 1001, "username": "planner", "fullname": "计划员", "mobile": "13800000001"}, {"id": 1002, "username": "collaborator", "fullname": "协作成员", "mobile": "13800000002"}, {"id": 1003, "username": "viewer", "fullname": "查看成员", "mobile": "13800000003"}, ) def _b64encode(raw: bytes) -> str: return base64.urlsafe_b64encode(raw).decode("ascii").rstrip("=") def _b64decode(raw: str) -> bytes: return base64.urlsafe_b64decode(raw + "=" * (-len(raw) % 4)) class TestAuthProvider(AuthProvider): __test__ = False mode = "test" def __init__(self, tenant_uuid: str) -> None: self.tenant_uuid = tenant_uuid self.secret = b"aps-test-auth-provider" def _identity(self, username: str) -> IdentityContext: user = next(row for row in TEST_USERS if row["username"] == username) return IdentityContext( user_id=int(user["id"]), username=str(user["username"]), fullname=str(user["fullname"]), tenant_uuid=self.tenant_uuid, roles=("planner",), expires_at=int(time.time()) + 3600, ) def _issue(self, identity: IdentityContext) -> str: body = _b64encode(json.dumps(identity.to_dict(), separators=(",", ":")).encode()) signature = _b64encode(hmac.new(self.secret, body.encode(), hashlib.sha256).digest()) return f"{body}.{signature}" def _verify(self, token: str) -> IdentityContext: try: body, signature = token.split(".", 1) expected = _b64encode(hmac.new(self.secret, body.encode(), hashlib.sha256).digest()) if not hmac.compare_digest(signature, expected): raise ValueError("signature") payload = json.loads(_b64decode(body)) return IdentityContext( user_id=int(payload["user_id"]), username=str(payload["username"]), fullname=str(payload["fullname"]), tenant_uuid=str(payload["tenant_uuid"]), roles=tuple(payload.get("roles") or ()), expires_at=int(payload["expires_at"]), ) except Exception as exc: raise AuthError("AUTH_INVALID", "测试登录凭证无效", 401) from exc async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]: identity = self._identity(str(payload.get("username") or "")) return identity, self._issue(identity) async def authenticate(self, request: Request) -> IdentityContext: token = request.cookies.get("aps_session") or "" if not token: raise AuthError("AUTH_REQUIRED", "请先登录", 401) return self._verify(token) async def refresh(self, request: Request) -> tuple[IdentityContext, str]: identity = await self.authenticate(request) return identity, self._issue(identity) async def search_users(self, query: str, identity: IdentityContext) -> list[dict[str, Any]]: needle = query.strip().lower() return [ { "id": int(user["id"]), "username": str(user["username"]), "fullname": str(user["fullname"]), "mobile": str(user["mobile"]), "tenantUuid": identity.tenant_uuid, } for user in TEST_USERS if int(user["id"]) != identity.user_id and (not needle or needle in " ".join(str(value) for value in user.values()).lower()) ] def install_test_auth(monkeypatch, tenant_uuid: str) -> TestAuthProvider: import server.auth.middleware as auth_middleware import server.gateway.app as gateway_app provider = TestAuthProvider(tenant_uuid) monkeypatch.setattr(auth_middleware, "get_auth_provider", lambda: provider) monkeypatch.setattr(gateway_app, "get_auth_provider", lambda: provider) return provider