import time from datetime import datetime, timedelta from auth0.v3.authentication import GetToken from auth0.v3.management import Auth0 from auth0.v3.management.rest import RestClientOptions import backoff from lucenequerybuilder import Q import requests from atlas_um.logs import logger from .models import User, UserList, Role from .schemas import UserSchema class Auth0Client: DEFAULT_ACCESS_TOKEN_EXPIRATION = 600 TIMEOUT = 30 RETRIES = 5 def __init__( self, app=None, domain=None, client_id=None, client_secret=None, access_token=None, token_expiration=None, ): self._client = None self.app = app if app is not None: self.init_app(app) self.domain = domain self.client_id = client_id self.client_secret = client_secret self.access_token = access_token self.access_token_generated_at = None self.token_expiration = ( token_expiration or self.DEFAULT_ACCESS_TOKEN_EXPIRATION ) def init_app(self, app): self.domain = self.domain or app.config.get("AUTH0_DOMAIN") self.client_id = self.client_id or app.config.get("AUTH0_CLIENT_ID") self.client_secret = self.client_secret or app.config.get( "AUTH0_CLIENT_SECRET" ) self.access_token = self.access_token or app.config.get( "AUTH0_ACCESS_TOKEN" ) self.token_expiration = self.token_expiration or app.config.get( "AUTH0_ACCESS_TOKEN_EXPIRATION" ) app.extensions["auth0"] = self @property def client(self): expired = self.should_fetch_access_token() if expired: self.fetch_access_token() if expired or not self._client: self._client = Auth0( self.domain, self.access_token, RestClientOptions(False, self.TIMEOUT, self.RETRIES), ) return self._client def fetch_access_token(self): get_token = GetToken(self.domain) management_api_root = "https://{}/api/v2/".format(self.domain) token = get_token.client_credentials( self.client_id, self.client_secret, management_api_root ) self.access_token_generated_at = datetime.now() self.access_token = token["access_token"] return self.access_token def should_fetch_access_token(self): if not self.access_token: return True return self.access_token_past_expiry() def access_token_past_expiry(self): return ( self.access_token_generated_at is not None and self.access_token_generated_at + timedelta(seconds=self.token_expiration) <= datetime.now() ) def list_users_ids(self, role_name=None, batch_from=None, batch_to=None): """ Using this method to iterate over the users using roles instead of the standard Auth0 method to list users as this method allows to use point pagination without limitations. """ for role in self.list_roles(): if role_name and role.get("name") != role_name: continue from_param = None page = 0 while True: resp = self.client.roles.list_users( id=role["id"], take=100, from_param=from_param ) users = resp.get("users", []) from_param = resp.get("next") page += 1 if batch_from and page < batch_from: time.sleep(0.1) continue for user in users: yield user.get("user_id") if not from_param or (batch_to and page >= batch_to): break def search_users(self, search_text): query_str = "*" + search_text.lower() + "*" q = str( Q("name", query_str, wildcard=True) | Q("email", query_str, wildcard=True) ) logger.bind(query=q).debug("User search performed with Lucene query") return UserList(self.client.users.list(q=q)) @backoff.on_exception( backoff.expo, requests.exceptions.RequestException, max_tries=3 ) def get_user(self, user_id: str) -> User: user_data = self.client.users.get(user_id) deserialized_user_data = UserSchema().load(user_data) return User(**deserialized_user_data) @backoff.on_exception( backoff.expo, requests.exceptions.RequestException, max_tries=3 ) def get_user_roles(self, user_id: str) -> list[Role]: roles_data = self.client.users.list_roles(user_id) return [Role(**r) for r in roles_data.get("roles", [])] @backoff.on_exception( backoff.expo, requests.exceptions.RequestException, max_tries=3 ) def get_user_by_email(self, email): users = self.client.users_by_email.search_users_by_email(email) if len(users) > 1: logger.bind( email=email, auth0_user_ids=[_user.get("user_id") for _user in users], ).error("Auth0 user has more than one email") return None if len(users) == 1: deserialized_user_data = UserSchema().load(users[0]) return User(**deserialized_user_data) def list_roles(self): resp = self.client.roles.list(per_page=100) for role in resp.get("roles", []): yield role