from __future__ import annotations import base64 import hashlib import hmac import json import os import time from abc import ABC, abstractmethod from typing import Any from fastapi import Request from server.auth.context import IdentityContext class AuthError(Exception): def __init__(self, code: str, message: str, status_code: int = 401) -> None: super().__init__(message) self.code = code self.message = message self.status_code = status_code class AuthProvider(ABC): mode = "unconfigured" @abstractmethod async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]: raise NotImplementedError @abstractmethod async def authenticate(self, request: Request) -> IdentityContext: raise NotImplementedError async def refresh(self, request: Request) -> tuple[IdentityContext, str]: identity = await self.authenticate(request) raise AuthError("AUTH_NOT_CONFIGURED", "刷新接口尚未配置", 503) async def logout(self, request: Request) -> None: return None async def search_users(self, query: str, identity: IdentityContext) -> list[dict[str, Any]]: raise AuthError("AUTH_NOT_CONFIGURED", "租户用户查询接口尚未配置", 503) class UnconfiguredAuthProvider(AuthProvider): async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]: raise AuthError("AUTH_NOT_CONFIGURED", "用户管理接口尚未配置", 503) async def authenticate(self, request: Request) -> IdentityContext: raise AuthError("AUTH_NOT_CONFIGURED", "用户管理接口尚未配置", 503) DEFAULT_MOCK_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 MockAuthProvider(AuthProvider): """Development-only provider. It is enabled only by APS_AUTH_PROVIDER=mock.""" mode = "mock" def __init__(self) -> None: self.secret = (os.environ.get("APS_MOCK_AUTH_SECRET") or "aps-local-development-only").encode() self.tenant_uuid = os.environ.get("APS_MOCK_TENANT_UUID") or "demo0000000000000000000000000001" self.ttl_seconds = int(os.environ.get("APS_AUTH_TTL_SECONDS") or "28800") def _directory(self) -> list[dict[str, Any]]: configured = os.environ.get("APS_MOCK_USERS") if configured: try: rows = json.loads(configured) if isinstance(rows, list): return [row for row in rows if isinstance(row, dict)] except json.JSONDecodeError: pass return [dict(user) for user in DEFAULT_MOCK_USERS] def _identity_for(self, payload: dict[str, Any]) -> IdentityContext: identifier = str( payload.get("username") or payload.get("mobile") or payload.get("identifier") or "planner" ).strip() users = self._directory() user = next( (row for row in users if identifier in {str(row.get("username")), str(row.get("mobile")), str(row.get("id"))}), users[0], ) return IdentityContext( user_id=int(user.get("id") or 1001), username=str(user.get("username") or identifier), fullname=str(user.get("fullname") or user.get("username") or identifier), tenant_uuid=self.tenant_uuid, roles=tuple(user.get("roles") or ("planner",)), expires_at=int(time.time()) + self.ttl_seconds, ) def _issue(self, identity: IdentityContext) -> str: payload = identity.to_dict() raw = json.dumps(payload, separators=(",", ":"), ensure_ascii=False).encode("utf-8") body = _b64encode(raw) sig = _b64encode(hmac.new(self.secret, body.encode("ascii"), hashlib.sha256).digest()) return f"{body}.{sig}" def _verify(self, token: str) -> IdentityContext: try: body, signature = token.split(".", 1) expected = _b64encode(hmac.new(self.secret, body.encode("ascii"), hashlib.sha256).digest()) if not hmac.compare_digest(signature, expected): raise ValueError("signature") payload = json.loads(_b64decode(body)) expires_at = int(payload.get("expires_at") or 0) if expires_at <= int(time.time()): raise AuthError("AUTH_EXPIRED", "登录状态已过期", 401) return IdentityContext( user_id=int(payload["user_id"]), username=str(payload["username"]), fullname=str(payload.get("fullname") or payload["username"]), tenant_uuid=str(payload["tenant_uuid"]), roles=tuple(payload.get("roles") or ()), expires_at=expires_at, ) except AuthError: raise except Exception as exc: raise AuthError("AUTH_INVALID", "登录凭证无效", 401) from exc @staticmethod def _request_token(request: Request) -> str: auth = request.headers.get("authorization") or "" if auth.lower().startswith("bearer "): return auth[7:].strip() return request.cookies.get("aps_session") or "" async def login(self, payload: dict[str, Any]) -> tuple[IdentityContext, str]: identity = self._identity_for(payload) return identity, self._issue(identity) async def authenticate(self, request: Request) -> IdentityContext: token = self._request_token(request) if not token: raise AuthError("AUTH_REQUIRED", "请先登录", 401) return self._verify(token) async def refresh(self, request: Request) -> tuple[IdentityContext, str]: current = await self.authenticate(request) identity = IdentityContext( user_id=current.user_id, username=current.username, fullname=current.fullname, tenant_uuid=current.tenant_uuid, roles=current.roles, expires_at=int(time.time()) + self.ttl_seconds, ) return identity, self._issue(identity) async def search_users(self, query: str, identity: IdentityContext) -> list[dict[str, Any]]: needle = query.strip().lower() rows = [] for user in self._directory(): if int(user.get("id") or 0) == identity.user_id: continue haystack = " ".join(str(user.get(k) or "") for k in ("username", "fullname", "mobile")).lower() if needle and needle not in haystack: continue rows.append({ "id": int(user["id"]), "username": str(user.get("username") or ""), "fullname": str(user.get("fullname") or user.get("username") or ""), "mobile": str(user.get("mobile") or ""), "tenantUuid": identity.tenant_uuid, }) return rows[:20] _provider: AuthProvider | None = None _provider_mode: str | None = None def get_auth_provider() -> AuthProvider: global _provider, _provider_mode mode = (os.environ.get("APS_AUTH_PROVIDER") or "unconfigured").strip().lower() if _provider is None or _provider_mode != mode: _provider = MockAuthProvider() if mode == "mock" else UnconfiguredAuthProvider() _provider_mode = mode return _provider