"""Auth0 API and utils.""" import time import logging from typing import Callable, Dict, Optional import requests from auth0_block_users import config from auth0_block_users import constants class Auth0: def __init__(self, logger: logging.Logger, token: str or None = None): self._logger = logger self._token = token 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) -> dict or None: """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: dict or None: 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'] = 'Bearer {}'.format(self._token) # make http request and retry on rate limits retry_count = 0 while True: response = requests.request( method=method, url=url, json=json_body, data=data, headers=headers, files=files ) if response.status_code != 429 or retry_count >= constants.RETRY_COUNT: break retry_count = retry_count + 1 self._logger.debug('Retry {}'.format(retry_count)) time.sleep(retry_count * constants.RETRY_WAIT_RATE) # handle non 2xx statuses if not response: self._logger.error('{} {}'.format(response.status_code, response.text)) raise ValueError('Incorrect request result') if response.text: response_data = response.json() self._logger.debug(response_data) return response_data return None 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 + '?page={}&per_page=100&include_totals=true'.format(page_index) 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 = page_index + 1 def get_management_token(self) -> str: """Get auth0 management API bearer token. Returns: str: Bearer token. """ url = 'https://{}/oauth/token'.format(config.AUTH0_DOMAIN) body = { 'client_id': config.AUTH0_MANAGEMENT_CLIENT_ID, 'client_secret': config.AUTH0_MANAGEMENT_CLIENT_SECRET, 'audience': 'https://{}/api/v2/'.format(config.AUTH0_DOMAIN), '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_users(self, fields: str or None = None, query: str or None = None) -> Dict[str, dict]: """Get users by query. Args: fields (str or None): User's fields to get. query (str or None): Search query. Returns: Dict[str, dict]: Roles mapping name to ID. """ url = 'https://{}/api/v2/users'.format(config.AUTH0_DOMAIN) params = {} if fields: params['fields'] = fields if query: params['q'] = query users = {} self._handle_pagination( url, lambda response: users.update({u['user_id']: u for u in response['users']}), params=params) return users def update_user(self, user_id: str, user_attrs: Dict): """Update users by query. Args: user_id (str): User ID. user_attrs (Dict): Attributes to update. """ url = 'https://{}/api/v2/users/{}'.format(config.AUTH0_DOMAIN, user_id) self._execute_request( url=url, json_body=user_attrs, method='PATCH', token_required=True)