summaryrefslogtreecommitdiffstats
path: root/authentication/backend.py
blob: 943490689e760135839f529055d6cdbbc9bb92ac (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
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<token>\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()
        )