"""Auth0 API and utils.""" import json import time from typing import Callable, Dict, List import zlib import requests import config from logger import logger from utils import execute_chunks def get_auth0_id(user_id: str) -> str: """Add auth0 prefix to user ID. Args: user_id (str): User ID. Returns: str: User ID with auth0 prefix. """ if 'auth0|' not in user_id: return 'auth0|{}'.format(user_id) else: return user_id def get_internal_id(user_id: str) -> str: """Remove auth0 prefix for user ID. Args: user_id (str): User ID. Returns: str: User ID without auth0 prefix. """ return user_id.replace('auth0|', '') def _handle_pagination(base_url: str, proc_data: Callable): """Handle Auth0 pagination. Args: base_url (str): Base endpoint URL. proc_data (Callable): Process response data. """ page_index = 0 while True: url = base_url + '?page={}&per_page=100&include_totals=true'.format(page_index) response = _execute_request(url=url) proc_data(response) if response['limit'] + response['start'] >= response['total']: break page_index = page_index + 1 def _execute_request( 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: str or None = None, 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 (str or None): Auth bearer token. token_required (bool): If auth required or not. Returns: dict or None: Response body. """ # get auth token if required if not token and token_required: token = get_management_token() # add auth header if token: if not headers: headers = {} headers['Authorization'] = 'Bearer {}'.format(token) logger.debug(url) # 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 >= config.Retry.COUNT: break retry_count = retry_count + 1 logger.debug('Retry {}'.format(retry_count)) time.sleep(retry_count * config.Retry.WAIT_RATE) # handle non 2xx statuses if not response: logger.error('{} {}'.format(response.status_code, response.text)) raise ValueError('Incorrect request result') if response.text: response_data = response.json() logger.debug(response_data) return response_data return None def get_management_token() -> 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 = _execute_request( url=url, data=body, method='POST', token_required=False) return token_data['access_token'] def get_job(job_id: str) -> dict: """Get job info. Args: job_id (str): Job ID. Returns: dict: Response body. """ url = 'https://{}/api/v2/jobs/{}'.format(config.Auth0.DOMAIN, job_id) response = _execute_request(url=url) return response def get_job_errors(job_id: str) -> dict: """Get job errors. Args: job_id (str): Job ID. Returns: dict: Response body. """ url = 'https://{}/api/v2/jobs/{}/errors'.format(config.Auth0.DOMAIN, job_id) response = _execute_request(url=url) return response def import_users_job(users: List[dict]) -> dict: """Run import users job. Args: users (List[dict]): List of users data. Returns: dict: Response body. """ url = 'https://{}/api/v2/jobs/users-imports'.format(config.Auth0.DOMAIN) body = {'upsert': 'true', 'send_completion_email': 'false', 'connection_id': config.Auth0.DB_CONNECTION_ID} data = {'users': json.dumps(users)} response = _execute_request(url=url, data=body, files=data, method='POST') return response def export_users_job(fields: List[str] or None = None) -> dict: """Run export users job. Args: fields (List[str] or None): Fields to export. Returns: dict: Response body. """ if not fields: fields = ['user_id', 'email'] fields = [{'name': f} for f in fields] url = 'https://{}/api/v2/jobs/users-exports'.format(config.Auth0.DOMAIN) body = { 'connection_id': config.Auth0.DB_CONNECTION_ID, 'format': 'json', 'fields': fields } response = _execute_request(url=url, json_body=body, method='POST') return response def run_job(func: Callable, *args, **kwargs) -> dict: """Run import/export job. Args: func (Callable): Function to call. Returns: dict: Job data. """ response = func(*args, **kwargs) job_id = response['id'] retry_count = 1 while response['status'] == 'pending': time.sleep(retry_count * config.Retry.WAIT_RATE) response = get_job(job_id) if retry_count < config.Retry.COUNT: retry_count = retry_count + 1 return response @execute_chunks(config.IMPORT_USERS_CHUNK_SIZE) def import_users(users: List[dict]) -> List[str]: """Import users. Args: users (List[dict]): List of users data. Returns: List[str]: Failed user IDs. """ response = run_job(import_users_job, users) job_id = response['id'] failed_users = [] if response['summary']['failed']: import_errors = get_job_errors(job_id) for error in import_errors: user_id = error['user']['user_id'] message = ','.join([e['message'] for e in error['errors']]) logger.error('{} import failed: {}'.format(user_id, message)) failed_users.append(user_id) return failed_users def export_users(fields: List[str] or None = None) -> Dict[str, str or dict]: """Export all users email and user_id fields. Args: fields (List[str] or None): Fields to export. Returns: Dict[str, str or dict]: User email to user_id or user data mapping. """ response = run_job(export_users_job, fields) response = requests.get(response['location']) decompressed_data = zlib.decompress(response.content, 16 + zlib.MAX_WBITS) text_data = decompressed_data.decode('utf-8') user_json_list = text_data.split('\n') users = {} for user_json in user_json_list: if user_json: user = json.loads(user_json) user['user_id'] = get_internal_id(user['user_id']) users[user['email']] = user if fields else user['user_id'] return users def get_role(name: str) -> str or None: """Get role by name. Args: name (str): Role name. Returns: str or None: Role ID or None if not found. """ url = 'https://{}/api/v2/roles?name_filter={}'.format( config.Auth0.DOMAIN, name) response = _execute_request(url=url) for role in response: if role['name'] == name: return role['id'] return None def create_role(name: str) -> str: """Create role. Args: name (str): Role name. Returns: str: Role ID. """ url = 'https://{}/api/v2/roles'.format(config.Auth0.DOMAIN) body = {'name': name} response = _execute_request(url=url, data=body, method='POST') return response['id'] @execute_chunks(config.ASSIGN_ROLE_CHUNK_SIZE) def assign_users_role(user_ids: List[str], role_id: str): """Assign users to role. Args: user_ids (List[str]): Set of user IDs. role_id (str): Role ID. """ url = 'https://{}/api/v2/roles/{}/users'.format(config.Auth0.DOMAIN, role_id) body = {'users': [get_auth0_id(user_id) for user_id in user_ids]} _execute_request(url=url, data=body, method='POST') def delete_user(user_id: str): """Delete user. Args: user_id (str): User ID. """ url = 'https://{}/api/v2/users/{}'.format( config.Auth0.DOMAIN, get_auth0_id(user_id)) _execute_request(url=url, method='DELETE') def get_user_roles(user_id: str) -> List[dict]: """Get user roles. Args: user_id (str): User ID. Returns: List[dict]: User roles. """ url = 'https://{}/api/v2/users/{}/roles'.format( config.Auth0.DOMAIN, get_auth0_id(user_id)) response = _execute_request(url=url) return response def get_role_assignments(role_id: str) -> List[dict]: """Get role assignments. Args: role_id (str): Role ID. Returns: List[str]: User IDs. """ url = 'https://{}/api/v2/roles/{}/users'.format(config.Auth0.DOMAIN, role_id) users = [] _handle_pagination( url, lambda response: users.extend( [get_internal_id(u['user_id']) for u in response['users']]) ) return users def get_all_roles() -> Dict[str, str]: """Get all roles. Returns: Dict[str, str]: Roles mapping name to ID. """ url = 'https://{}/api/v2/roles'.format(config.Auth0.DOMAIN) roles = {} _handle_pagination( url, lambda response: roles.update( {r['name']: r['id'] for r in response['roles']}) ) return roles def delete_role(role_id: str): """Delete role. Args: role_id (str): Role ID. """ url = 'https://{}/api/v2/roles/{}'.format(config.Auth0.DOMAIN, role_id) _execute_request(url=url, method='DELETE')