import time from json import JSONDecodeError from typing import Callable, Dict, Optional import requests from .config import AUTH0_DOMAIN, AUTH0_MANAGEMENT_CLIENT_ID, \ AUTH0_MANAGEMENT_CLIENT_SECRET, DEFAULT_REQUEST_TIMEOUT, \ PAGE_SIZE, RETRY_COUNT, RETRY_WAIT_RATE from requests.exceptions import ConnectionError, Timeout class Auth0: def __init__(self, logger): """Auth0 API and utils""" self._token = None self.logger = logger def _execute_request( self, url: str, headers: dict or None = None, json_body: dict or None = None, data: dict or None = None, files: dict or None = None, method: str = "GET", token_required: bool = True, ) -> Optional[dict]: """Execute request to auth0. Args: url (str): Request URL. headers (dict or None): Request headers. json_body (dict or None): Request json body. data (dict or None): Request data. files (dict or None): Multipart form data. method (str): Request method. token_required (bool): If auth required or not. Returns: Optional[dict]: Response body. """ # get auth token if required if not self._token and token_required: self._token = self.get_management_token() # add auth header if self._token: if not headers: headers = {} headers["Authorization"] = f"Bearer {self._token}" # make http request and retry on rate limits retry_count = 0 result = None while True: try: response = requests.request( method=method, url=url, json=json_body, data=data, headers=headers, files=files, timeout=DEFAULT_REQUEST_TIMEOUT, ) if method != "DELETE": result = response.json() except (ConnectionError, Timeout, ValueError, JSONDecodeError) as exc: if retry_count >= RETRY_COUNT: raise exc self.logger.warning( f"{retry_count}/{RETRY_COUNT} Got exception {exc} on request to {url} with headers: {headers} and" f"data: {data}." ) retry_count += 1 time.sleep(retry_count * RETRY_WAIT_RATE) continue return result def _handle_pagination(self, base_url: str, callback: Callable, params: Dict or None = None): """Handle Auth0 pagination. Args: base_url (str): Base endpoint URL. callback (Callable): Process response data. params (Dict or None): Search params. """ page_index = 0 while True: url = base_url + f"?page={page_index}&per_page={PAGE_SIZE}&include_totals=true" if params: url = "{}&{}".format(url, "&".join(["{}={}".format(k, p) for k, p in params.items()])) response = self._execute_request(url=url) callback(response) if response["limit"] + response["start"] >= response["total"]: break page_index += 1 def get_management_token(self) -> str: """Get auth0 management API bearer token. Returns: str: Bearer token. """ url = f"https://{AUTH0_DOMAIN}/oauth/token" body = { "client_id": AUTH0_MANAGEMENT_CLIENT_ID, "client_secret": AUTH0_MANAGEMENT_CLIENT_SECRET, "audience": f"https://{AUTH0_DOMAIN}/api/v2/", "grant_type": "client_credentials", } token_data = self._execute_request(url=url, data=body, method="POST", token_required=False) self._token = token_data["access_token"] return token_data["access_token"] def get_logs(self, fields: str or None = None, query: str or None = None) -> Dict[str, dict]: url = f"https://{AUTH0_DOMAIN}/api/v2/logs" params = {} if query: params["q"] = query users = {} self._handle_pagination( url, lambda response: users.update({u["log_id"]: u for u in response["logs"]}), params=params ) return users def unblock_user_by_identifier(self, identifier: str): """Unblock user by identifier. Should be any of a username, phone number, or email. Args: :param identifier: """ url = f"https://{AUTH0_DOMAIN}/api/v2/user-blocks?identifier={requests.utils.quote(identifier)}" response = self._execute_request(url=url, method="DELETE") return response