from typing import Dict, Any, Optional, TypeVar, Type import aiohttp import backoff import marshmallow_dataclass from aiohttp import ClientResponse, ClientResponseError, ContentTypeError from marshmallow import ValidationError from ..exceptions import BadGateway, NotFound, APIError from ..constants import HTTP_400_BAD_REQUEST, HTTP_404_NOT_FOUND, HTTP_403_FORBIDDEN,\ HTTP_401_UNAUTHORIZED class BaseClientError(Exception): code: str detail: str def __init__(self, code: str, detail: str): self.code = code self.detail = detail class ApiUnavailable(BaseClientError): pass class ApiUnauthorized(BaseClientError): pass class ApiInvalidRequestError(BaseClientError): server_error: Optional[str] = None def __init__(self, error: Optional[str] = None): self.server_error = error class ApiResponseParsingError(BaseClientError): pass T = TypeVar("T") class ApiClient: _session: Optional[aiohttp.ClientSession] = None base_headers: Optional[Dict[str, str]] = None api_host: str def __init__( self, api_host: str, headers: Optional[Dict[str, str]] = None, ): self.base_headers = headers self.api_host = api_host @property def session(self): if self._session is None: self._session = aiohttp.ClientSession(headers=self.base_headers) return self._session async def get( self, endpoint: str, params: Optional[Dict[str, Any]] = None, headers: Optional[Dict[str, str]] = None, response_type: Type[T] = None, ) -> T: return await self.execute("get", endpoint=endpoint, params=params, headers=headers, response_type=response_type) async def post( self, endpoint: str, params: Optional[Dict[str, Any]] = None, payload: Optional[Dict[str, Any]] = None, headers: Optional[Dict[str, str]] = None, response_type: Type[T] = None, ) -> T: return await self.execute("post", endpoint=endpoint, params=params, json=payload, headers=headers, response_type=response_type) async def delete( self, endpoint: str, params: Optional[Dict[str, Any]] = None, payload: Optional[Dict[str, Any]] = None, headers: Optional[Dict[str, str]] = None, response_type: Type[T] = None, ) -> T: return await self.execute("delete", endpoint=endpoint, params=params, json=payload, headers=headers, response_type=response_type) async def execute( self, method: str, endpoint: str, params: Optional[Dict[str, Any]] = None, json: Optional[Dict[str, Any]] = None, headers: Optional[Dict[str, str]] = None, response_type: Type[T] = None, ) -> T: try: response = await self._execute(method, endpoint=endpoint, params=params, json=json, headers=headers) except ClientResponseError as error: raise BadGateway(detail="Code: {}. Message: {}".format(error.status, error.message)) except BaseClientError: raise BadGateway() except APIError: raise BadGateway() if response_type is None: return None try: json_body = await response.json() except ContentTypeError as error: raise ApiResponseParsingError(code="validation_error", detail=error.message) is_response_array = isinstance(json_body, list) schema = marshmallow_dataclass.class_schema(response_type)(many=is_response_array) try: data = schema.load(json_body) except ValidationError as map_error: raise ApiResponseParsingError(code="validation_error", detail=map_error.normalized_messages()) return data @backoff.on_exception(backoff.expo, ApiUnavailable, max_time=3) async def _execute( self, method: str, endpoint: str, params: Optional[Dict[str, Any]] = None, json: Optional[Dict[str, Any]] = None, headers: Optional[Dict[str, str]] = None, ) -> Any: try: if headers is None: headers = {} response: ClientResponse = await self.session.request( method, self.api_host + endpoint, params=params, headers=headers, json=json ) if response.status in (HTTP_401_UNAUTHORIZED, HTTP_403_FORBIDDEN): raise ApiUnauthorized(code="Unauthorized", detail="Token is expired or not valid") if response.status == HTTP_404_NOT_FOUND: raise NotFound() if response.status == HTTP_400_BAD_REQUEST: server_error = await response.text() raise ApiInvalidRequestError(server_error) return response except (aiohttp.ClientConnectorError, aiohttp.ServerTimeoutError): raise ApiUnavailable