198 lines
7.6 KiB
Python
198 lines
7.6 KiB
Python
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
|