from __future__ import annotations import hashlib from dataclasses import dataclass import jwt from app.config import settings from fastapi import Depends, HTTPException, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer from jwt import ExpiredSignatureError, InvalidTokenError from packages.domain.entities import User from packages.ports.user_repository import UserRepository from .dependencies import get_user_repository bearer_scheme = HTTPBearer(auto_error=False) @dataclass(frozen=True, slots=True) class AuthenticatedUser: user: User session_id: str | None = None token_type: str | None = None def _get_redis_client(): """获取 Redis 客户端用于 JWT 黑名单""" import redis as redis_lib return redis_lib.from_url(settings.REDIS_URL, decode_responses=True) def _token_fingerprint(token: str) -> str: """计算 token 的哈希指纹""" return hashlib.sha256(token.encode()).hexdigest() def blacklist_token(token: str, exp: int) -> None: """将 token 加入黑名单,TTL 为 token 剩余有效期""" import time redis_client = _get_redis_client() key = f"jwt:blacklist:{_token_fingerprint(token)}" ttl = max(exp - int(time.time()), 1) redis_client.setex(key, ttl, "revoked") def is_token_blacklisted(token: str) -> bool: """检查 token 是否在黑名单中""" redis_client = _get_redis_client() key = f"jwt:blacklist:{_token_fingerprint(token)}" return redis_client.exists(key) > 0 async def get_current_user( credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme), user_repository: UserRepository = Depends(get_user_repository), ) -> AuthenticatedUser: if credentials is None or credentials.scheme.lower() != "bearer": raise _unauthorized("Missing bearer token") payload = _decode_user_token(credentials.credentials) user_id = payload.get("sub") if not isinstance(user_id, str) or not user_id: raise _unauthorized("Invalid token subject") user = user_repository.find_by_id(user_id) if user is None: raise _unauthorized("User no longer exists") return AuthenticatedUser( user=user, session_id=payload.get("sid"), token_type=payload.get("type"), ) def _decode_user_token(token: str) -> dict: try: payload = jwt.decode(token, settings.JWT_SECRET_KEY, algorithms=["HS256"]) except ExpiredSignatureError: raise _unauthorized("Token expired") from None except InvalidTokenError: raise _unauthorized("Invalid token") from None if payload.get("type") not in {"user_auth", "access"}: raise _unauthorized("Invalid token type") # 检查 token 是否在黑名单中 if is_token_blacklisted(token): raise _unauthorized("Token has been revoked") return payload def _unauthorized(detail: str) -> HTTPException: return HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail=detail, headers={"WWW-Authenticate": "Bearer"}, )