aps-agent/server/auth/providers.py

198 lines
7.6 KiB
Python
Raw Normal View History

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