From 498c79f434856aa68e9a883247ea69a22256fce6 Mon Sep 17 00:00:00 2001 From: Tom van der Lee Date: Wed, 22 Jul 2026 14:51:44 +0200 Subject: Added auth layer --- authentication/authlib.py | 167 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 167 insertions(+) create mode 100644 authentication/authlib.py (limited to 'authentication/authlib.py') diff --git a/authentication/authlib.py b/authentication/authlib.py new file mode 100644 index 0000000..8472062 --- /dev/null +++ b/authentication/authlib.py @@ -0,0 +1,167 @@ +import json +from collections import defaultdict +from unittest.mock import patch + +from authlib.common.security import generate_token as generate_random_token +from authlib.oauth2 import AuthorizationServer as BaseAuthorizationServer, OAuth2Request +from authlib.oauth2.rfc6749 import AuthorizationCodeGrant as BaseAuthorizationCodeGrant, \ + RefreshTokenGrant as BaseRefreshTokenGrant, OAuth2Payload +from authlib.oauth2.rfc6750 import BearerTokenGenerator +from authlib.oauth2.rfc7636 import CodeChallenge +from sqlalchemy.exc import NoResultFound + +from sqlmodel import select +from fastapi import Request, Response + +from db.models import ClientApplication, AuthToken, AuthCode, User +from db.session import get_session_context + + +class AuthorizationCodeGrant(BaseAuthorizationCodeGrant): + TOKEN_ENDPOINT_AUTH_METHODS = ['client_secret_basic', 'client_secret_post', 'none'] + + def save_authorization_code(self, code: str, request: OAuth2Request): + with get_session_context() as db_session: + db_session.add(request.user) + + code_challenge = request.payload.data.get('code_challenge') + code_challenge_method = request.payload.data.get('code_challenge_method') + + auth_code = AuthCode( + code=code, + client_id=request.client.client_id, + redirect_uri=request.payload.redirect_uri, + response_type=request.payload.response_type, + scope=request.payload.scope, + user_id=request.user.id, + code_challenge=code_challenge, + code_challenge_method=code_challenge_method, + ) + db_session.add(auth_code) + + return auth_code + + def query_authorization_code(self, code: str, client: ClientApplication) -> AuthCode | None: + try: + with get_session_context() as db_session: + auth_code = db_session.exec(select(AuthCode).where(AuthCode.code == code, AuthCode.client_id == client.client_id)).one() + except NoResultFound: + return None + + if auth_code.is_expired(): + return None + + return auth_code + + def delete_authorization_code(self, code: AuthCode): + with get_session_context() as db_session: + db_session.delete(code) + + def authenticate_user(self, authentication_code: AuthCode) -> User: + return authentication_code.user + + +class RefreshTokenGrant(BaseRefreshTokenGrant): + def authenticate_refresh_token(self, refresh_token) -> AuthToken | None: + try: + with get_session_context() as db_session: + token = db_session.exec(select(AuthToken).where(AuthToken.refresh_token == refresh_token)).one() + except NoResultFound: + return None + + if not token.is_refresh_token_active(): + return None + + return token + + def authenticate_user(self, credential: AuthToken): + return credential.user + + def revoke_old_credential(self, credential: AuthToken): + with get_session_context() as db_session: + credential.revoked = True + db_session.add(credential) + +async def prepare_oauth_request(request: Request) -> Request: + async with request.form() as form: + request.state.form_data = dict(form) + return request + + +class FastApiOAuth2Payload(OAuth2Payload): + def __init__(self, request: Request): + self._request = request + + @property + def data(self): + return { + **self._request.query_params, + **getattr(self._request.state, 'form_data', {}), + } + + @property + def datalist(self): + values = defaultdict(list) + for k in self.data: + values[k].extend([self.data[k]]) + return values + +class FastApiOAuth2Request(OAuth2Request): + def __init__(self, request: Request): + with patch('authlib.oauth2.rfc6749.errors.InsecureTransportError'): + super().__init__( + method=request.method, + uri=str(request.url), + headers=request.headers, + ) + self.method = request.method + self.uri = str(request.url) + self.headers = request.headers + self.payload = FastApiOAuth2Payload(request) + self.user = request.user + + self._request = request + + @property + def args(self): + return self._request.query_params + + @property + def form(self): + return getattr(self._request.state, 'form_data', {}) + + +class AuthorizationServer(BaseAuthorizationServer): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.register_grant(AuthorizationCodeGrant, extensions=[CodeChallenge()]) + self.register_grant(RefreshTokenGrant) + self.register_token_generator("default", BearerTokenGenerator( + access_token_generator=lambda *args, **kwargs: generate_random_token(42), + refresh_token_generator=lambda *args, **kwargs: generate_random_token(48), + )) + + def query_client(self, client_id: str) -> ClientApplication: + with get_session_context() as db_session: + return db_session.exec(select(ClientApplication).where(ClientApplication.client_id == client_id)).one() + + def save_token(self, token: dict, request: OAuth2Request): + with get_session_context() as db_session: + db_session.add(AuthToken( + user_id=request.user.id, + client_id=request.client.client_id, + **token + )) + + def create_oauth2_request(self, request: Request) -> OAuth2Request: + return FastApiOAuth2Request(request) + + def handle_response(self, status, body, headers) -> Response: + return Response( + status_code=status, + content=json.dumps(body) if isinstance(body, dict) else body, + headers=dict(headers), + ) + + def send_signal(self, name, *args, **kwargs): + pass -- cgit v1.2.3