201 lines
7.5 KiB
Python
201 lines
7.5 KiB
Python
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)
|