codekingpro/portable-devtools
114k
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 