summaryrefslogtreecommitdiffstats
path: root/authentication/backend.py
diff options
context:
space:
mode:
Diffstat (limited to 'authentication/backend.py')
-rw-r--r--authentication/backend.py40
1 files changed, 31 insertions, 9 deletions
diff --git a/authentication/backend.py b/authentication/backend.py
index 29408e3..9434906 100644
--- a/authentication/backend.py
+++ b/authentication/backend.py
@@ -1,4 +1,4 @@
1from datetime import datetime, UTC 1import re
2from typing import Optional 2from typing import Optional
3 3
4from sqlalchemy import update 4from sqlalchemy import update
@@ -6,11 +6,11 @@ from sqlmodel import select
6from starlette.authentication import AuthenticationBackend, AuthCredentials, BaseUser, UnauthenticatedUser 6from starlette.authentication import AuthenticationBackend, AuthCredentials, BaseUser, UnauthenticatedUser
7from starlette.requests import HTTPConnection 7from starlette.requests import HTTPConnection
8 8
9from db.models import Session, User 9from db.models import Session, User, AuthToken
10from db.session import get_session_context 10from db.session import get_session_context
11 11
12 12
13class AuthBackend(AuthenticationBackend): 13class SessionAuthBackend(AuthenticationBackend):
14 async def authenticate(self, request: HTTPConnection) -> Optional[tuple[AuthCredentials, BaseUser]]: 14 async def authenticate(self, request: HTTPConnection) -> Optional[tuple[AuthCredentials, BaseUser]]:
15 if "id" not in request.session: 15 if "id" not in request.session:
16 return None 16 return None
@@ -33,16 +33,38 @@ class AuthBackend(AuthenticationBackend):
33 ) 33 )
34 db_session.commit() 34 db_session.commit()
35 35
36
36 return ( 37 return (
37 AuthCredentials( 38 AuthCredentials(
38 [ 39 ["authenticated", "session"]
39 "authenticated",
40 *[
41 connection.type for connection in session.user.app_connections
42 ]
43 ]
44 if session is not None 40 if session is not None
45 else [] 41 else []
46 ), 42 ),
47 session.user if session is not None else UnauthenticatedUser() 43 session.user if session is not None else UnauthenticatedUser()
48 ) 44 )
45
46
47class BearerTokenAuthBackend(AuthenticationBackend):
48 regex = re.compile(r"^[Bb]earer\s(?P<token>\S+)$")
49 async def authenticate(self, request: HTTPConnection) -> Optional[tuple[AuthCredentials, BaseUser]]:
50 token = None
51 if "Authorization" in request.headers:
52 match = self.regex.match(request.headers["Authorization"])
53
54 query = (
55 select(AuthToken)
56 .join(User)
57 .where(AuthToken.access_token == match.group('token'))
58 )
59
60 with get_session_context() as db_session:
61 token: AuthToken = db_session.exec(query).first()
62
63 return (
64 AuthCredentials(
65 ["authenticated", "token"]
66 if token is not None
67 else []
68 ),
69 token.user if token is not None else UnauthenticatedUser()
70 )