import re from typing import Optional from sqlalchemy import update from sqlmodel import select from starlette.authentication import AuthenticationBackend, AuthCredentials, BaseUser, UnauthenticatedUser from starlette.requests import HTTPConnection from db.models import Session, User, AuthToken from db.session import get_session_context class SessionAuthBackend(AuthenticationBackend): async def authenticate(self, request: HTTPConnection) -> Optional[tuple[AuthCredentials, BaseUser]]: if "id" not in request.session: return None query = ( select(Session) .join(User) .where( Session.id == request.session['id'] ) ) with get_session_context() as db_session: session: Session = db_session.exec(query).first() if session is not None: db_session.execute( update(Session).where(Session.id == session.id) ) db_session.commit() return ( AuthCredentials( ["authenticated", "session"] if session is not None else [] ), session.user if session is not None else UnauthenticatedUser() ) class BearerTokenAuthBackend(AuthenticationBackend): regex = re.compile(r"^[Bb]earer\s(?P\S+)$") async def authenticate(self, request: HTTPConnection) -> Optional[tuple[AuthCredentials, BaseUser]]: token = None if "Authorization" in request.headers: match = self.regex.match(request.headers["Authorization"]) query = ( select(AuthToken) .join(User) .where(AuthToken.access_token == match.group('token')) ) with get_session_context() as db_session: token: AuthToken = db_session.exec(query).first() return ( AuthCredentials( ["authenticated", "token"] if token is not None else [] ), token.user if token is not None else UnauthenticatedUser() )