aps-agent/tests/auth_provider.py

111 lines
4.1 KiB
Python
Raw Normal View History

2026-07-29 23:22:40 +08:00
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