Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
oauth2_client.py221 linesDownload Raw Back to httpx_client
1import typing2from contextlib import asynccontextmanager3 4import httpx5from httpx import Auth, Request, Response, USE_CLIENT_DEFAULT6from anyio import Lock  # Import after httpx so import errors refer to httpx7from authlib.common.urls import url_decode8from authlib.oauth2.client import OAuth2Client as _OAuth2Client9from authlib.oauth2.auth import ClientAuth, TokenAuth10from .utils import HTTPX_CLIENT_KWARGS, build_request11from ..base_client import (12    OAuthError,13    InvalidTokenError,14    MissingTokenError,15    UnsupportedTokenTypeError,16)17 18__all__ = [19    'OAuth2Auth', 'OAuth2ClientAuth',20    'AsyncOAuth2Client', 'OAuth2Client',21]22 23 24class OAuth2Auth(Auth, TokenAuth):25    """Sign requests for OAuth 2.0, currently only bearer token is supported."""26    requires_request_body = True27 28    def auth_flow(self, request: Request) -> typing.Generator[Request, Response, None]:29        try:30            url, headers, body = self.prepare(31                str(request.url), request.headers, request.content)32            headers['Content-Length'] = str(len(body))33            yield build_request(url=url, headers=headers, body=body, initial_request=request)34        except KeyError as error:35            description = f'Unsupported token_type: {str(error)}'36            raise UnsupportedTokenTypeError(description=description)37 38 39class OAuth2ClientAuth(Auth, ClientAuth):40    requires_request_body = True41 42    def auth_flow(self, request: Request) -> typing.Generator[Request, Response, None]:43        url, headers, body = self.prepare(44            request.method, str(request.url), request.headers, request.content)45        headers['Content-Length'] = str(len(body))46        yield build_request(url=url, headers=headers, body=body, initial_request=request)47 48 49class AsyncOAuth2Client(_OAuth2Client, httpx.AsyncClient):50    SESSION_REQUEST_PARAMS = HTTPX_CLIENT_KWARGS51 52    client_auth_class = OAuth2ClientAuth53    token_auth_class = OAuth2Auth54    oauth_error_class = OAuthError55 56    def __init__(self, client_id=None, client_secret=None,57                 token_endpoint_auth_method=None,58                 revocation_endpoint_auth_method=None,59                 scope=None, redirect_uri=None,60                 token=None, token_placement='header',61                 update_token=None, **kwargs):62 63        # extract httpx.Client kwargs64        client_kwargs = self._extract_session_request_params(kwargs)65        httpx.AsyncClient.__init__(self, **client_kwargs)66 67        # We use a Lock to synchronize coroutines to prevent68        # multiple concurrent attempts to refresh the same token69        self._token_refresh_lock = Lock()70 71        _OAuth2Client.__init__(72            self, session=None,73            client_id=client_id, client_secret=client_secret,74            token_endpoint_auth_method=token_endpoint_auth_method,75            revocation_endpoint_auth_method=revocation_endpoint_auth_method,76            scope=scope, redirect_uri=redirect_uri,77            token=token, token_placement=token_placement,78            update_token=update_token, **kwargs79        )80 81    async def request(self, method, url, withhold_token=False, auth=USE_CLIENT_DEFAULT, **kwargs):82        if not withhold_token and auth is USE_CLIENT_DEFAULT:83            if not self.token:84                raise MissingTokenError()85 86            await self.ensure_active_token(self.token)87 88            auth = self.token_auth89 90        return await super().request(91            method, url, auth=auth, **kwargs)92 93    @asynccontextmanager94    async def stream(self, method, url, withhold_token=False, auth=USE_CLIENT_DEFAULT, **kwargs):95        if not withhold_token and auth is USE_CLIENT_DEFAULT:96            if not self.token:97                raise MissingTokenError()98 99            await self.ensure_active_token(self.token)100 101            auth = self.token_auth102 103        async with super().stream(104            method, url, auth=auth, **kwargs) as resp:105            yield resp106 107    async def ensure_active_token(self, token):108        async with self._token_refresh_lock:109            if self.token.is_expired():110                refresh_token = token.get('refresh_token')111                url = self.metadata.get('token_endpoint')112                if refresh_token and url:113                    await self.refresh_token(url, refresh_token=refresh_token)114                elif self.metadata.get('grant_type') == 'client_credentials':115                    access_token = token['access_token']116                    new_token = await self.fetch_token(url, grant_type='client_credentials')117                    if self.update_token:118                        await self.update_token(new_token, access_token=access_token)119                else:120                    raise InvalidTokenError()121 122    async def _fetch_token(self, url, body='', headers=None, auth=USE_CLIENT_DEFAULT,123                           method='POST', **kwargs):124        if method.upper() == 'POST':125            resp = await self.post(126                url, data=dict(url_decode(body)), headers=headers,127                auth=auth, **kwargs)128        else:129            if '?' in url:130                url = '&'.join([url, body])131            else:132                url = '?'.join([url, body])133            resp = await self.get(url, headers=headers, auth=auth, **kwargs)134 135        for hook in self.compliance_hook['access_token_response']:136            resp = hook(resp)137 138        return self.parse_response_token(resp)139 140    async def _refresh_token(self, url, refresh_token=None, body='',141                             headers=None, auth=USE_CLIENT_DEFAULT, **kwargs):142        resp = await self.post(143            url, data=dict(url_decode(body)), headers=headers,144            auth=auth, **kwargs)145 146        for hook in self.compliance_hook['refresh_token_response']:147            resp = hook(resp)148 149        token = self.parse_response_token(resp)150        if 'refresh_token' not in token:151            self.token['refresh_token'] = refresh_token152 153        if self.update_token:154            await self.update_token(self.token, refresh_token=refresh_token)155 156        return self.token157 158    def _http_post(self, url, body=None, auth=USE_CLIENT_DEFAULT, headers=None, **kwargs):159        return self.post(160            url, data=dict(url_decode(body)),161            headers=headers, auth=auth, **kwargs)162 163 164class OAuth2Client(_OAuth2Client, httpx.Client):165    SESSION_REQUEST_PARAMS = HTTPX_CLIENT_KWARGS166 167    client_auth_class = OAuth2ClientAuth168    token_auth_class = OAuth2Auth169    oauth_error_class = OAuthError170 171    def __init__(self, client_id=None, client_secret=None,172                 token_endpoint_auth_method=None,173                 revocation_endpoint_auth_method=None,174                 scope=None, redirect_uri=None,175                 token=None, token_placement='header',176                 update_token=None, **kwargs):177 178        # extract httpx.Client kwargs179        client_kwargs = self._extract_session_request_params(kwargs)180        httpx.Client.__init__(self, **client_kwargs)181 182        _OAuth2Client.__init__(183            self, session=self,184            client_id=client_id, client_secret=client_secret,185            token_endpoint_auth_method=token_endpoint_auth_method,186            revocation_endpoint_auth_method=revocation_endpoint_auth_method,187            scope=scope, redirect_uri=redirect_uri,188            token=token, token_placement=token_placement,189            update_token=update_token, **kwargs190        )191 192    @staticmethod193    def handle_error(error_type, error_description):194        raise OAuthError(error_type, error_description)195 196    def request(self, method, url, withhold_token=False, auth=USE_CLIENT_DEFAULT, **kwargs):197        if not withhold_token and auth is USE_CLIENT_DEFAULT:198            if not self.token:199                raise MissingTokenError()200 201            if not self.ensure_active_token(self.token):202                raise InvalidTokenError()203 204            auth = self.token_auth205 206        return super().request(207            method, url, auth=auth, **kwargs)208 209    def stream(self, method, url, withhold_token=False, auth=USE_CLIENT_DEFAULT, **kwargs):210        if not withhold_token and auth is USE_CLIENT_DEFAULT:211            if not self.token:212                raise MissingTokenError()213 214            if not self.ensure_active_token(self.token):215                raise InvalidTokenError()216 217            auth = self.token_auth218 219        return super().stream(220            method, url, auth=auth, **kwargs)221 
codekingpro/portable-devtools · Team Ai