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()
)
|