8935196fcd
Deploy / Staging E2E Tests (push) Has been skipped
Deploy / Build Production Runtime Images (push) Has been skipped
Deploy / Deploy Production (push) Has been skipped
Deploy / Production Browser E2E (push) Has been skipped
Deploy / Deploy Staging (push) Failing after 138h4m33s
CI/CD Pipeline / Frontend Lint (push) Failing after 138h4m39s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 138h4m39s
103 lines
3.0 KiB
Python
103 lines
3.0 KiB
Python
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"},
|
|
)
|