"""Utility function for jwtauth testing.""" import json from typing import Any, Protocol import requests from pyotp import TOTP from jwtauth.testing import CLAIM_ORCHARD_IDENTITY_ID, JWT_USER_METADATA from jwtauth.testing.config import AUTH_AUDIENCE, AUTH_URL from jwtauth.testing.constants import ( AUTH0_ERROR_FIELD, AUTH0_MFA_GRANT_TYPE, AUTH0_MFA_REQUIRED_MSG, AUTH0_MFA_TOKEN_FIELD, ) from jwtauth.testing.schemas import Auth0Creds, SecretLookupInfo, UserCreds from jwtauth.utils import jwt_auth_from_environment class CredentialLookupError(Exception): """ Raised when a Secrets Manager lookup fails when fetching Auth0 App credentials or User credentials. """ pass class MfaTokenLookupError(Exception): """ Raised when multifactor authentication fails due to missing credentials, failed OTP, or a 400 error returned from the `/oauth/token` Auth0 endpoint. """ pass # Matches owsclient.protocols.SecretsManager class SecretsManager(Protocol): """SecretsManager gets a secret string.""" def get_secret(self, secret_name: str) -> str | dict[str, Any]: """Return a string.""" ... def _get_secret( secrets_manager: SecretsManager, service_name: str, secret_name: str = "AUTH0_CREDS", environment: str = "qa", ) -> dict[str, Any] | None: """Fetch the secret and return a dict. This method should be reusable by methods to fetch auth0 secrets and user secrets. """ full_secret_name = f"{environment}/{service_name}/{secret_name}" secret_data = secrets_manager.get_secret(full_secret_name) if not secret_data: raise CredentialLookupError(f"Could not find secret '{full_secret_name}'") if isinstance(secret_data, dict): return secret_data return json.loads(secret_data) # type: ignore[no-any-return] def get_auth0_creds( secrets_manager: SecretsManager, service_name: str, secret_name: str = "AUTH0_CREDENTIALS", environment: str = "qa", ) -> Auth0Creds: """Get credentials.""" auth0_creds = _get_secret( secrets_manager=secrets_manager, service_name=service_name, secret_name=secret_name, environment=environment, ) return Auth0Creds.model_validate(auth0_creds) def get_user_creds( secrets_manager: SecretsManager, service_name: str, secret_name: str = "USER_CREDENTIALS", environment: str = "qa", ) -> UserCreds: """Get user credentials.""" user_creds = _get_secret( secrets_manager=secrets_manager, service_name=service_name, secret_name=secret_name, environment=environment, ) return UserCreds.model_validate(user_creds) def generate_bearer_jwt_token(user_creds: UserCreds, auth0_creds: Auth0Creds) -> str: """Generate an auth token for the given UserCreds and Auth0Creds.""" data = { "grant_type": "password", "username": user_creds.email, "password": user_creds.password.get_secret_value(), "audience": AUTH_AUDIENCE, "scope": "", "client_id": auth0_creds.auth0_client_id, "client_secret": auth0_creds.auth0_client_secret, } r = requests.post(AUTH_URL, data=data) resp = r.json() if AUTH0_ERROR_FIELD in resp and resp[AUTH0_ERROR_FIELD] == AUTH0_MFA_REQUIRED_MSG: # The MFA function will raise `MfaTokenLookupError` if resp does # not contain `AUTH0_MFA_TOKEN_FIELD` return generate_bearer_jwt_token_mfa( user_creds=user_creds, auth0_creds=auth0_creds, mfa_token=resp.get(AUTH0_MFA_TOKEN_FIELD), ) elif AUTH0_ERROR_FIELD in resp: raise ValueError(f"Auth0 error occurred: {resp}") return str(resp["access_token"]) def generate_bearer_jwt_token_mfa( user_creds: UserCreds, auth0_creds: Auth0Creds, mfa_token: str, ) -> str: """Generate an auth token for the given UserCreds, Auth0Creds and mfa_token.""" # Raise an error when `otp_secret_key` is falsy or is the string "null". if (not user_creds.otp_secret_key) or (user_creds.otp_secret_key == "null"): raise MfaTokenLookupError("UserCreds.otp_secret_key is empty") if not mfa_token: raise MfaTokenLookupError("mfa_token is empty") try: # Generate one-time-passcode one_time_password = TOTP(user_creds.otp_secret_key).now() except Exception as ex: raise MfaTokenLookupError("Failed to generate one-time-password") from ex # Omitting optional `client_assertion` and `client_assertion_type` body parameters data = { "grant_type": AUTH0_MFA_GRANT_TYPE, "client_id": auth0_creds.auth0_client_id, "client_secret": auth0_creds.auth0_client_secret, "mfa_token": mfa_token, "otp": one_time_password, } r = requests.post(AUTH_URL, data=data) # Expected status codes: # 200 - OTP verification successful. # 400 - Invalid request, such as missing parameters or invalid OTP. if r.status_code != 200: raise MfaTokenLookupError( f"Auth0 error occurred (status {r.status_code}): {r.text}" ) resp = r.json() return str(resp["access_token"]) def login_from_secrets_manager( get_user_creds_args: SecretLookupInfo, get_auth0_creds_args: SecretLookupInfo, secrets_manager: SecretsManager, ) -> str: """Fetch secrets then generate auth0 credentials.""" user_creds = get_user_creds( secrets_manager=secrets_manager, service_name=get_user_creds_args.service_name, secret_name=get_user_creds_args.secret_name, environment=get_user_creds_args.environment, ) assert user_creds, "Failed to get user creds" auth0_creds = get_auth0_creds( secrets_manager=secrets_manager, service_name=get_auth0_creds_args.service_name, secret_name=get_auth0_creds_args.secret_name, environment=get_auth0_creds_args.environment, ) assert auth0_creds, "Failed to get auth0 creds" return generate_bearer_jwt_token(user_creds=user_creds, auth0_creds=auth0_creds) def decode_token(token: str, environment: str) -> dict[str, Any]: """Extract claims from token.""" auth = jwt_auth_from_environment(environment=environment) return auth.get_token(token) def get_bearer_token_identity_uuid(bearer_token: str, environment: str) -> str | None: """Extract identity_uuid from bearer_token.""" token_claims = decode_token(bearer_token, environment=environment) identity_uuid = token_claims.get(JWT_USER_METADATA, {}).get( CLAIM_ORCHARD_IDENTITY_ID ) assert ( isinstance(identity_uuid, str) or identity_uuid is None ), f"Found unexpected type for {CLAIM_ORCHARD_IDENTITY_ID}: {type(identity_uuid)}" return identity_uuid