import math import time from typing import Generator, List from auth0.v3 import Auth0Error from auth0.v3.authentication import GetToken from auth0.v3.management import Auth0 from copy_auth0_logs_to_cloudwatch import config _auth0_client: Auth0 or None = None def _get_token() -> str: """Get access token. Returns: str: Access token. """ get_token_obj = GetToken(config.AUTH0_DOMAIN) token = get_token_obj.client_credentials( config.AUTH0_CLIENT_ID, config.AUTH0_CLIENT_SECRET, 'https://{}/api/v2/'.format(config.AUTH0_DOMAIN) ) access_token = token['access_token'] return access_token def _update_auth0_client(): """Update auth0 client object. """ global _auth0_client _auth0_client = Auth0(config.AUTH0_DOMAIN, _get_token()) def _get_auth0_client() -> Auth0: """Get auth0 client. """ if not _auth0_client: _update_auth0_client() return _auth0_client def retry_error(f): """Handle rate limits or any other possible errors. """ def wrapped(*args, **kwargs): for i in range(1, config.RATE_LIMIT_RETRY_COUNT + 1): try: return f(*args, **kwargs) except Auth0Error: if i == config.RATE_LIMIT_RETRY_COUNT: raise time.sleep(i * config.RATE_LIMIT_WATE_MULTIPLIER) return wrapped def token_expiration(f): """Handle token expiration. """ def wrapped(*args, **kwargs): try: return f(*args, **kwargs) except Auth0Error as e: if e.status_code == 401: _update_auth0_client() return f(*args, **kwargs) else: raise return wrapped @retry_error @token_expiration def _execute_search(page_number: int, search_query: str) -> List[dict]: """Execute Auth0 logs search with rate limit handling. Args: page_number (int): Search page number. search_query (str): Search query. Returns: List[dict]: Log events. """ result = _get_auth0_client().logs.search( page=page_number, q=search_query, per_page=config.LOG_EVENTS_COUNT_PER_PAGE, fields=config.LOG_FIELDS, include_fields=config.LOG_INCLUDE_FIELDS, sort=config.LOG_SORT_ORDER ) return result['logs'] def get_logs(log_id: str) -> Generator[dict, None, None]: """Get log events from Auth0 management API. Args: log_id (str): Last processed log ID. Returns: Generator[dict, None, None]: Log events generator. """ page_count = math.ceil(config.LOG_EVENTS_COUNT_PER_SEARCH / config.LOG_EVENTS_COUNT_PER_PAGE) search_query = f'log_id:{{{log_id} TO *] AND date:[{config.LOG_MIN_DATE} TO *]' if config.LOG_EVENTS_FILTER_CONDITION: search_query = f'{search_query} AND {config.LOG_EVENTS_FILTER_CONDITION}' for i in range(page_count): events = _execute_search(i, search_query) if not events: return for entry in sorted(events, key=lambda x: x.get('log_id')): yield entry if len(events) < config.LOG_EVENTS_COUNT_PER_PAGE: return