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