111 lines
4.1 KiB
Python
111 lines
4.1 KiB
Python
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
|