128 lines
4.0 KiB
Python
128 lines
4.0 KiB
Python
|
|
import hashlib
|
||
|
|
import hmac
|
||
|
|
import secrets
|
||
|
|
import threading
|
||
|
|
import time
|
||
|
|
from collections import defaultdict, deque
|
||
|
|
from datetime import datetime, timezone
|
||
|
|
|
||
|
|
from argon2 import PasswordHasher
|
||
|
|
from argon2.exceptions import InvalidHashError, VerificationError
|
||
|
|
from fastapi import Depends, HTTPException, Request, status
|
||
|
|
from psycopg import Connection
|
||
|
|
|
||
|
|
from .config import settings
|
||
|
|
from .db import get_connection
|
||
|
|
|
||
|
|
password_hasher = PasswordHasher()
|
||
|
|
|
||
|
|
|
||
|
|
def hash_password(password: str) -> str:
|
||
|
|
return password_hasher.hash(password)
|
||
|
|
|
||
|
|
|
||
|
|
def verify_password(password: str, password_hash: str) -> bool:
|
||
|
|
try:
|
||
|
|
return password_hasher.verify(password_hash, password)
|
||
|
|
except (InvalidHashError, VerificationError):
|
||
|
|
return False
|
||
|
|
|
||
|
|
|
||
|
|
def token_digest(token: str) -> str:
|
||
|
|
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||
|
|
|
||
|
|
|
||
|
|
def new_token() -> str:
|
||
|
|
return secrets.token_urlsafe(32)
|
||
|
|
|
||
|
|
|
||
|
|
def new_csrf_token() -> str:
|
||
|
|
return secrets.token_urlsafe(32)
|
||
|
|
|
||
|
|
|
||
|
|
def code_digest(code: str) -> str:
|
||
|
|
return hmac.new(settings.csrf_secret.encode(), code.encode(), hashlib.sha256).hexdigest()
|
||
|
|
|
||
|
|
|
||
|
|
def utc_now() -> datetime:
|
||
|
|
return datetime.now(timezone.utc)
|
||
|
|
|
||
|
|
|
||
|
|
def check_origin(request: Request) -> None:
|
||
|
|
origin = request.headers.get("origin")
|
||
|
|
if origin and origin.rstrip("/") not in settings.allowed_origins:
|
||
|
|
raise HTTPException(status_code=403, detail={"code": "ORIGIN_NOT_ALLOWED", "message": "请求来源不被允许"})
|
||
|
|
|
||
|
|
|
||
|
|
def require_csrf(request: Request) -> None:
|
||
|
|
check_origin(request)
|
||
|
|
header = request.headers.get("x-csrf-token")
|
||
|
|
cookie = request.cookies.get("kaotings_csrf")
|
||
|
|
if not header or not cookie or not hmac.compare_digest(header, cookie):
|
||
|
|
raise HTTPException(status_code=403, detail={"code": "CSRF_FAILED", "message": "请求校验失败"})
|
||
|
|
|
||
|
|
|
||
|
|
def set_session_cookie(response, token: str) -> None:
|
||
|
|
response.set_cookie(
|
||
|
|
settings.session_cookie_name,
|
||
|
|
token,
|
||
|
|
max_age=settings.session_ttl_seconds,
|
||
|
|
httponly=True,
|
||
|
|
secure=settings.session_cookie_secure,
|
||
|
|
samesite="lax",
|
||
|
|
path="/",
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def clear_session_cookie(response) -> None:
|
||
|
|
response.delete_cookie(settings.session_cookie_name, path="/")
|
||
|
|
|
||
|
|
|
||
|
|
def set_csrf_cookie(response, token: str) -> None:
|
||
|
|
response.set_cookie("kaotings_csrf", token, max_age=settings.session_ttl_seconds, httponly=False, secure=settings.session_cookie_secure, samesite="lax", path="/")
|
||
|
|
|
||
|
|
|
||
|
|
def current_user(request: Request, connection: Connection = Depends(get_connection)) -> dict:
|
||
|
|
token = request.cookies.get(settings.session_cookie_name)
|
||
|
|
if not token:
|
||
|
|
raise HTTPException(status_code=401, detail={"code": "AUTH_REQUIRED", "message": "请先登录"})
|
||
|
|
user = connection.execute(
|
||
|
|
"""
|
||
|
|
SELECT u.*, s.id AS session_id
|
||
|
|
FROM sessions s JOIN users u ON u.id = s.user_id
|
||
|
|
WHERE s.token_hash = %s AND s.revoked_at IS NULL AND s.expires_at > now()
|
||
|
|
""",
|
||
|
|
(token_digest(token),),
|
||
|
|
).fetchone()
|
||
|
|
if not user:
|
||
|
|
raise HTTPException(status_code=401, detail={"code": "AUTH_REQUIRED", "message": "会话已失效"})
|
||
|
|
if user["status"] != "active":
|
||
|
|
raise HTTPException(status_code=403, detail={"code": "ACCOUNT_DISABLED", "message": "账户不可用"})
|
||
|
|
return user
|
||
|
|
|
||
|
|
|
||
|
|
def require_admin(user: dict = Depends(current_user)) -> dict:
|
||
|
|
if user["role"] != "admin":
|
||
|
|
raise HTTPException(status_code=403, detail={"code": "FORBIDDEN", "message": "无权执行此操作"})
|
||
|
|
return user
|
||
|
|
|
||
|
|
|
||
|
|
class RateLimiter:
|
||
|
|
def __init__(self) -> None:
|
||
|
|
self._events: dict[str, deque[float]] = defaultdict(deque)
|
||
|
|
self._lock = threading.Lock()
|
||
|
|
|
||
|
|
def allow(self, key: str, limit: int, window_seconds: int) -> bool:
|
||
|
|
now = time.monotonic()
|
||
|
|
with self._lock:
|
||
|
|
events = self._events[key]
|
||
|
|
while events and events[0] <= now - window_seconds:
|
||
|
|
events.popleft()
|
||
|
|
if len(events) >= limit:
|
||
|
|
return False
|
||
|
|
events.append(now)
|
||
|
|
return True
|
||
|
|
|
||
|
|
|
||
|
|
rate_limiter = RateLimiter()
|