from __future__ import annotations import hmac import threading import time from collections import OrderedDict, deque from collections.abc import Callable from dataclasses import dataclass from app.auth.repository import AuthRepository, SessionRecord from app.auth.schemas import LoginRequest, PasswordChangeRequest, SetupRequest from app.core.errors import AppError, AuthenticationError, ConflictError from app.core.security import hash_password, new_token, token_hash, verify_password from app.core.time import expires_in @dataclass(slots=True) class SessionTokens: session_token: str csrf_token: str def _constant_time_equal(left: str, right: str) -> bool: return hmac.compare_digest(left.encode("utf-8"), right.encode("utf-8")) class LoginRateLimiter: def __init__( self, attempts: int = 5, window_seconds: int = 600, max_keys: int = 10_000, cleanup_interval_seconds: int = 60, clock: Callable[[], float] = time.monotonic, ) -> None: self.attempts = attempts self.window_seconds = window_seconds self.max_keys = max_keys self.cleanup_interval_seconds = min(cleanup_interval_seconds, window_seconds) self._clock = clock self._events: OrderedDict[str, deque[float]] = OrderedDict() self._lock = threading.Lock() self._last_cleanup = 0.0 def check(self, key: str) -> None: now = self._clock() with self._lock: self._cleanup(now) events = self._events.get(key) if events is None: return self._discard_expired(events, now) if not events: self._events.pop(key, None) return self._events.move_to_end(key) if len(events) >= self.attempts: raise AppError( "登录尝试过于频繁,请稍后再试", code="LOGIN_RATE_LIMITED", status_code=429, title="请求过于频繁", ) def record_failure(self, key: str) -> None: now = self._clock() with self._lock: self._cleanup(now) events = self._events.get(key) if events is None: if len(self._events) >= self.max_keys: self._cleanup(now, force=True) while len(self._events) >= self.max_keys: self._events.popitem(last=False) events = deque() self._events[key] = events else: self._discard_expired(events, now) self._events.move_to_end(key) events.append(now) def reset(self, key: str) -> None: with self._lock: self._events.pop(key, None) def _cleanup(self, now: float, *, force: bool = False) -> None: if not force and now - self._last_cleanup < self.cleanup_interval_seconds: return for key, events in list(self._events.items()): self._discard_expired(events, now) if not events: self._events.pop(key, None) self._last_cleanup = now def _discard_expired(self, events: deque[float], now: float) -> None: while events and now - events[0] >= self.window_seconds: events.popleft() class AuthService: def __init__( self, repository: AuthRepository, session_days: int, bootstrap_token: str | None = None, ) -> None: self.repository = repository self.session_days = session_days self.bootstrap_token = bootstrap_token self.rate_limiter = LoginRateLimiter() def bootstrap_state(self, token: str | None) -> dict[str, object]: requires_setup = not self.repository.has_admin() session = self.authenticate(token) if token else None return { "requires_setup": requires_setup, "bootstrap_token_required": requires_setup and self.bootstrap_token is not None, "authenticated": session is not None, "username": session.username if session else None, } def setup(self, request: SetupRequest, *, loopback_client: bool) -> SessionTokens: if self.repository.has_admin(): raise ConflictError("管理员已经初始化", code="SETUP_ALREADY_COMPLETED") if self.bootstrap_token is not None: supplied_token = request.bootstrap_token or "" if not _constant_time_equal(self.bootstrap_token, supplied_token): raise AppError( "初始化令牌无效", code="BOOTSTRAP_TOKEN_INVALID", status_code=403, title="请求已拒绝", ) elif not loopback_client: raise AppError( "未配置初始化令牌时,只能从控制器本机创建管理员", code="BOOTSTRAP_LOCAL_ONLY", status_code=403, title="请求已拒绝", ) try: self.repository.create_admin(request.username, hash_password(request.password)) except ValueError as exc: raise ConflictError("管理员已经初始化", code="SETUP_ALREADY_COMPLETED") from exc return self._new_session() def login(self, request: LoginRequest, client_key: str) -> SessionTokens: self.rate_limiter.check(client_key) user = self.repository.get_admin() if user is None or not _constant_time_equal(user.username, request.username.strip()): self.rate_limiter.record_failure(client_key) raise AuthenticationError("用户名或密码不正确") if not verify_password(user.password_hash, request.password): self.rate_limiter.record_failure(client_key) raise AuthenticationError("用户名或密码不正确") self.rate_limiter.reset(client_key) self.repository.update_login_time() return self._new_session() def authenticate(self, token: str | None) -> SessionRecord | None: if not token: return None return self.repository.get_session(token_hash(token)) def require_auth( self, token: str | None, csrf_token: str | None, unsafe: bool ) -> SessionRecord: session = self.authenticate(token) if session is None: raise AuthenticationError() if unsafe and ( not csrf_token or not hmac.compare_digest(session.csrf_hash, token_hash(csrf_token)) ): raise AppError( "页面会话校验失败,请刷新后重试", code="CSRF_VALIDATION_FAILED", status_code=403, title="请求已拒绝", ) return session def logout(self, token: str | None) -> None: if token: self.repository.delete_session(token_hash(token)) def change_password(self, request: PasswordChangeRequest) -> None: user = self.repository.get_admin() if user is None or not verify_password(user.password_hash, request.current_password): raise AuthenticationError("当前密码不正确") self.repository.update_password(hash_password(request.new_password)) def _new_session(self) -> SessionTokens: session_token = new_token() csrf_token = new_token() self.repository.create_session( token_hash=token_hash(session_token), csrf_hash=token_hash(csrf_token), expires_at=expires_in(self.session_days), ) return SessionTokens(session_token, csrf_token)