from __future__ import annotations 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 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") return payload def _unauthorized(detail: str) -> HTTPException: return HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail=detail, headers={"WWW-Authenticate": "Bearer"}, )