Files

78 lines
2.7 KiB
Python

import secrets
from datetime import datetime, timedelta
import bcrypt
from fastapi import Depends, HTTPException, Request, status
from sqlalchemy.orm import Session as DbSession
from app.config import settings
from app.database import get_db
from app.models import Session, User
SESSION_COOKIE = "session_token"
def hash_password(password: str) -> str:
return bcrypt.hashpw(password.encode(), bcrypt.gensalt()).decode()
def verify_password(plain: str, hashed: str) -> bool:
return bcrypt.checkpw(plain.encode(), hashed.encode())
def create_session_token() -> str:
return secrets.token_urlsafe(32)
def create_user_session(db: DbSession, user: User) -> Session:
token = create_session_token()
expires_at = datetime.utcnow() + timedelta(hours=settings.session_lifetime_hours)
session = Session(user_id=user.id, token=token, expires_at=expires_at)
db.add(session)
db.commit()
db.refresh(session)
return session
def get_session_by_token(db: DbSession, token: str) -> Session | None:
session = db.query(Session).filter(Session.token == token, Session.is_active.is_(True)).first()
if not session:
return None
if session.expires_at < datetime.utcnow():
session.is_active = False
db.commit()
return None
return session
def get_current_user(request: Request, db: DbSession = Depends(get_db)) -> User:
token = request.cookies.get(SESSION_COOKIE)
if not token:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
session = get_session_by_token(db, token)
if not session:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Session expired")
user = db.query(User).filter(User.id == session.user_id).first()
if not user:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="User not found")
return user
def get_current_user_optional(request: Request, db: DbSession = Depends(get_db)) -> User | None:
token = request.cookies.get(SESSION_COOKIE)
if not token:
return None
session = get_session_by_token(db, token)
if not session:
return None
return db.query(User).filter(User.id == session.user_id).first()
def verify_api_key(request: Request) -> None:
auth = request.headers.get("Authorization", "")
if not auth.startswith("Bearer "):
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid API key")
key = auth.removeprefix("Bearer ").strip()
if not secrets.compare_digest(key, settings.api_key):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Invalid API key")