from typing import Tuple import requests from src.api_client.base import BaseApiClient from src.api_client.errors import AuthError, BaseApiError from src.api_utils.request import Request __all__ = ["AuthApiClient"] class AuthApiClient(BaseApiClient): default_error_cls = AuthError src_auth_headers: Tuple[str, ...] = ("Authorization",) dst_auth_headers: Tuple[str, ...] = ("Treatments", "X-Client-Id", "X-User-Id", "X-Userinfo") def __init__(self, *, auth_url: str, **kwargs): super().__init__(**kwargs) self._auth_url = auth_url @staticmethod def _is_auth_error(error: BaseApiError) -> bool: return False def authorize(self, request: Request): headers = {key: request.headers.get(key) or request.cookies.get(key) for key in self.src_auth_headers} auth_request = requests.Request(method="get", url=self._auth_url, headers=headers) response = self.send_request_with_retry(auth_request) for auth_header in self.dst_auth_headers: if auth_header not in response.headers: continue request.headers[auth_header] = response.headers[auth_header]