"""JWT authentication helpers.""" from datetime import datetime, timedelta, timezone from typing import Optional from jose import JWTError, jwt from passlib.context import CryptContext from fastapi import Depends, HTTPException, status, Request from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select from app.config import get_settings from app.database import get_db from app.models.admin_user import AdminUser settings = get_settings() pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") def verify_password(plain_password: str, hashed_password: str) -> bool: """Verify a plain password against a hash.""" return pwd_context.verify(plain_password, hashed_password) def get_password_hash(password: str) -> str: """Hash a password using bcrypt.""" return pwd_context.hash(password) def create_access_token(data: dict, expires_delta: timedelta | None = None) -> str: """Create a JWT access token.""" to_encode = data.copy() expire = datetime.now(timezone.utc) + ( expires_delta or timedelta(hours=settings.jwt_expire_hours) ) to_encode.update({"exp": expire}) return jwt.encode(to_encode, settings.jwt_secret, algorithm=settings.jwt_algorithm) def verify_token(token: str) -> dict | None: """Verify a JWT token and return its payload.""" try: payload = jwt.decode(token, settings.jwt_secret, algorithms=[settings.jwt_algorithm]) return payload except JWTError: return None async def get_current_user( request: Request, db: AsyncSession = Depends(get_db), ) -> AdminUser: """Dependency: extract JWT from cookie, verify, return AdminUser.""" token = request.cookies.get("hms_admin_token") if not token: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated", ) payload = verify_token(token) if not payload: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired token", ) username = payload.get("sub") if not username: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid token payload", ) result = await db.execute(select(AdminUser).where(AdminUser.username == username)) user = result.scalar_one_or_none() if not user: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="User not found", ) return user