"""Tests for jwtauth.testing utils function.""" import json from typing import Any from unittest.mock import MagicMock, patch import pytest from pydantic import SecretStr from jwtauth import JWTAuth from jwtauth.testing.config import AUTH_AUDIENCE, AUTH_URL from jwtauth.testing.constants import AUTH0_MFA_GRANT_TYPE, AUTH0_MFA_REQUIRED_MSG from jwtauth.testing.schemas import Auth0Creds, SecretLookupInfo, UserCreds from jwtauth.testing.secretsmanager import JwtAuthSecretsManager from jwtauth.testing.utils import ( CredentialLookupError, MfaTokenLookupError, SecretsManager, decode_token, generate_bearer_jwt_token, generate_bearer_jwt_token_mfa, get_auth0_creds, get_bearer_token_identity_uuid, get_user_creds, login_from_secrets_manager, ) @pytest.mark.parametrize( "secret_name,environment,get_secret_return_value,expected,expected_full_secret_name", [ pytest.param( "ANOTHER_AUTH0_CREDENTIALS", "prod", { "auth0_client_id": "1234565656", "auth0_client_secret": "secret", }, Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="secret", ), "prod/test-service/ANOTHER_AUTH0_CREDENTIALS", id="get_secret returns a dictionary is OK", ), pytest.param( "YET_ANOTHER_AUTH0_CREDENTIALS", "uat", json.dumps( { "auth0_client_id": "1234565656", "auth0_client_secret": "secret", } ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="secret", ), "uat/test-service/YET_ANOTHER_AUTH0_CREDENTIALS", id="get_secret returns a string is OK", ), pytest.param( "YET_ANOTHER_AUTH0_CREDENTIALS", None, { "auth0_client_id": "1234565656", "auth0_client_secret": "secret", }, Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="secret", ), "qa/test-service/YET_ANOTHER_AUTH0_CREDENTIALS", id="Uses default values for environment", ), pytest.param( None, "uat", { "auth0_client_id": "1234565656", "auth0_client_secret": "secret", }, Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="secret", ), "uat/test-service/AUTH0_CREDENTIALS", id="Uses default value for secret_name", ), ], ) def test_get_auth0_creds( secret_name: str | None, environment: str | None, get_secret_return_value: dict[str, Any] | None, expected_full_secret_name: str, expected: Auth0Creds, ) -> None: """Test get_auth0_creds.""" secret_manager = MagicMock(spec=SecretsManager) secret_manager.get_secret.return_value = get_secret_return_value kwargs = { "secrets_manager": secret_manager, "service_name": "test-service", } if secret_name: kwargs["secret_name"] = secret_name if environment: kwargs["environment"] = environment ret = get_auth0_creds( **kwargs, # type: ignore[arg-type] ) assert ret == expected secret_manager.get_secret.assert_called_with( expected_full_secret_name, ) def test_get_auth0_creds__error() -> None: """Test a get_auth0_creds error case.""" secret_manager = MagicMock(spec=SecretsManager) secret_manager.get_secret.return_value = None with pytest.raises(CredentialLookupError) as excinfo: _ = get_auth0_creds( secrets_manager=secret_manager, service_name="test-service", secret_name="AUTH0_CREDS", ) assert "Could not find secret 'qa/test-service/AUTH0_CREDS'" in str(excinfo.value) secret_manager.get_secret.assert_called_with( "qa/test-service/AUTH0_CREDS", ) @pytest.mark.parametrize( "secret_name,environment,get_secret_return_value,expected,expected_full_secret_name", [ pytest.param( "ANOTHER_USER_CREDENTIALS", "prod", { "email": "test@sonymusic-pde.com", "password": "secret", "otp_secret_key": "1234", }, UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), "prod/test-service/ANOTHER_USER_CREDENTIALS", id="get_secret returns a dictionary is OK", ), pytest.param( "YET_ANOTHER_USER_CREDENTIALS", "uat", json.dumps( { "email": "test@sonymusic-pde.com", "password": "secret", "otp_secret_key": "1234", } ), UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), "uat/test-service/YET_ANOTHER_USER_CREDENTIALS", id="get_secret returns a string is OK", ), pytest.param( "YET_ANOTHER_USER_CREDENTIALS", None, { "email": "test@sonymusic-pde.com", "password": "secret", "otp_secret_key": "1234", }, UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), "qa/test-service/YET_ANOTHER_USER_CREDENTIALS", id="Uses default values for environment", ), pytest.param( None, "uat", { "email": "test@sonymusic-pde.com", "password": "secret", "otp_secret_key": "1234", }, UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), "uat/test-service/USER_CREDENTIALS", id="Uses default value for secret_name", ), ], ) def test_get_user_creds( secret_name: str | None, environment: str | None, get_secret_return_value: dict[str, Any] | None, expected_full_secret_name: str, expected: UserCreds, ) -> None: """Test get_user_creds.""" secret_manager = MagicMock(spec=SecretsManager) secret_manager.get_secret.return_value = get_secret_return_value kwargs = { "secrets_manager": secret_manager, "service_name": "test-service", } if secret_name: kwargs["secret_name"] = secret_name if environment: kwargs["environment"] = environment ret = get_user_creds( **kwargs, # type: ignore[arg-type] ) assert ret == expected secret_manager.get_secret.assert_called_with( expected_full_secret_name, ) def test_get_user_creds__error() -> None: """Test a get_user_creds error case.""" secret_manager = MagicMock(spec=SecretsManager) secret_manager.get_secret.return_value = None with pytest.raises(CredentialLookupError) as excinfo: _ = get_user_creds( secrets_manager=secret_manager, service_name="test-service", secret_name="USER_CREDENTIALS", environment="qa", ) assert "Could not find secret 'qa/test-service/USER_CREDENTIALS'" in str( excinfo.value ) secret_manager.get_secret.assert_called_with( "qa/test-service/USER_CREDENTIALS", ) @pytest.mark.parametrize( "user_creds, auth0_creds, mock_post_response, expected_token, mfa_required", [ pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ), { "access_token": "valid_token_123", "token_type": "Bearer", "expires_in": 3600, }, "valid_token_123", False, id="Valid credentials", ), pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ), { "error": AUTH0_MFA_REQUIRED_MSG, "mfa_token": "test_mfa_token_string", }, "mfa_valid_token_456", True, id="MFA required", ), ], ) @patch("jwtauth.testing.utils.generate_bearer_jwt_token_mfa") @patch("jwtauth.testing.utils.requests.post") def test_generate_bearer_jwt_token( mock_post: MagicMock, mock_generate_bearer_jwt_token_mfa: MagicMock, user_creds: UserCreds, auth0_creds: Auth0Creds, mock_post_response: dict[str, Any], expected_token: str, mfa_required: bool, ) -> None: """Test generate_auth_token.""" mock_post.return_value.json.return_value = mock_post_response if mfa_required: mock_generate_bearer_jwt_token_mfa.return_value = "mfa_valid_token_456" else: mock_generate_bearer_jwt_token_mfa.return_value = None ret = generate_bearer_jwt_token( user_creds=user_creds, auth0_creds=auth0_creds, ) assert ret == expected_token mock_post.assert_called_once_with( AUTH_URL, 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, }, ) if mfa_required: mock_generate_bearer_jwt_token_mfa.assert_called_once_with( user_creds=user_creds, auth0_creds=auth0_creds, mfa_token=mock_post_response["mfa_token"], ) else: mock_generate_bearer_jwt_token_mfa.assert_not_called() @pytest.mark.parametrize( "user_creds, auth0_creds, mock_post_response", [ pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), Auth0Creds( auth0_client_id="invalid-client-id", auth0_client_secret="abcdef123456", ), { "error": "Invalid client credentials", }, id="Invalid auth0 client id", ), ], ) @patch("jwtauth.testing.utils.requests.post") def test_generate_bearer_jwt_token__error( mock_post: MagicMock, user_creds: UserCreds, auth0_creds: Auth0Creds, mock_post_response: dict[str, Any], ) -> None: """Test generate_auth_token error case.""" mock_post.return_value.json.return_value = mock_post_response with pytest.raises(ValueError) as excinfo: _ = generate_bearer_jwt_token( user_creds=user_creds, auth0_creds=auth0_creds, ) assert "Auth0 error occurred" in str(excinfo.value) @patch("jwtauth.testing.utils.generate_bearer_jwt_token_mfa") @patch("jwtauth.testing.utils.requests.post") def test_generate_bearer_jwt_token__mfa_exception_propagated( mock_post: MagicMock, mock_generate_bearer_jwt_token_mfa: MagicMock, ) -> None: """Test generate_bearer_jwt_token propagates MFA exceptions.""" user_creds = UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ) auth0_creds = Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ) # Mock initial response requiring MFA mock_post.return_value.json.return_value = { "error": AUTH0_MFA_REQUIRED_MSG, "mfa_token": "test_mfa_token_string", } # Mock MFA function to raise an exception mock_generate_bearer_jwt_token_mfa.side_effect = MfaTokenLookupError( "Failed to generate MFA token" ) with pytest.raises(MfaTokenLookupError) as excinfo: _ = generate_bearer_jwt_token( user_creds=user_creds, auth0_creds=auth0_creds, ) assert "Failed to generate MFA token" in str(excinfo.value) mock_generate_bearer_jwt_token_mfa.assert_called_once_with( user_creds=user_creds, auth0_creds=auth0_creds, mfa_token="test_mfa_token_string", ) @patch("jwtauth.testing.utils.generate_bearer_jwt_token") @patch("jwtauth.testing.utils.get_user_creds") @patch("jwtauth.testing.utils.get_auth0_creds") def test_login_from_secrets_manager( mock_get_auth0_creds: MagicMock, mock_get_user_creds: MagicMock, mock_generate_bearer_jwt_token: MagicMock, ) -> None: """Test login_from_secrets_manager""" mock_generate_bearer_jwt_token.return_value = "test.token" jwtauth_secrets_manager = MagicMock(spec=JwtAuthSecretsManager) mock_get_auth0_creds.return_value = Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="secret", ) mock_get_user_creds.return_value = UserCreds( email="test@sonymusic-pde.com", password=SecretStr("password"), otp_secret_key=None, ) bearer_token = login_from_secrets_manager( get_user_creds_args=SecretLookupInfo( environment="qa", service_name="pdp-integration-test", secret_name="PDP_TEST_USER_CREDENTIALS", ), get_auth0_creds_args=SecretLookupInfo( environment="qa", service_name="pdp-integration-test", secret_name="PDP_TEST_APP_AUTH0_CREDENTIALS", ), secrets_manager=jwtauth_secrets_manager, ) assert bearer_token == "test.token" mock_get_auth0_creds.assert_called_once_with( secrets_manager=jwtauth_secrets_manager, service_name="pdp-integration-test", secret_name="PDP_TEST_APP_AUTH0_CREDENTIALS", environment="qa", ) mock_get_user_creds.assert_called_once_with( secrets_manager=jwtauth_secrets_manager, service_name="pdp-integration-test", secret_name="PDP_TEST_USER_CREDENTIALS", environment="qa", ) @pytest.mark.parametrize( "mock_get_auth0_creds_return_value, mock_get_user_creds_return_value, expected_error", [ pytest.param( Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="secret", ), None, "Failed to get user creds", id="get_user_creds returns null", ) ], ) @patch("jwtauth.testing.utils.generate_bearer_jwt_token") @patch("jwtauth.testing.utils.get_user_creds") @patch("jwtauth.testing.utils.get_auth0_creds") def test_login_from_secrets_manager__error( mock_get_auth0_creds: MagicMock, mock_get_user_creds: MagicMock, mock_generate_bearer_jwt_token: MagicMock, mock_get_auth0_creds_return_value: Auth0Creds | None, mock_get_user_creds_return_value: UserCreds | None, expected_error: str, ) -> None: """Test login_from_secrets_manager""" mock_generate_bearer_jwt_token.return_value = "test.token" mock_get_auth0_creds.return_value = mock_get_auth0_creds_return_value mock_get_user_creds.return_value = mock_get_user_creds_return_value jwtauth_secrets_manager = MagicMock(spec=JwtAuthSecretsManager) with pytest.raises(AssertionError) as excinfo: _ = login_from_secrets_manager( get_user_creds_args=SecretLookupInfo( environment="qa", service_name="test-application", secret_name="USER_CREDENTIALS", ), get_auth0_creds_args=SecretLookupInfo( environment="qa", service_name="test-application", secret_name="APP_AUTH0_CREDENTIALS", ), secrets_manager=jwtauth_secrets_manager, ) assert expected_error in str(excinfo.value) mock_generate_bearer_jwt_token.assert_not_called() @patch("jwtauth.testing.utils.jwt_auth_from_environment") def test_decode_token(mock_jwt_auth_from_environment: MagicMock) -> None: """Test decode token.""" mock_jwt_auth = MagicMock(spec=JWTAuth) mock_jwt_auth_from_environment.return_value = mock_jwt_auth expected = {"claim1": "value1", "claim2": "value2"} mock_jwt_auth.get_token.return_value = expected assert decode_token("test-token", environment="qa") == expected mock_jwt_auth_from_environment.assert_called_once_with(environment="qa") mock_jwt_auth.get_token.assert_called_once_with("test-token") @pytest.mark.parametrize( "token_claims, expected_identity_uuid", [ pytest.param( { "https://grass.theorchard.com/user_metadata": { "orchardIdentityId": "my_identity_uuid" }, }, "my_identity_uuid", id="The claims have a valid orchardIdentityId", ), pytest.param( { "https://grass.theorchard.com/user_metadata": {"other": "claim"}, }, None, id="The claims do not have orchardIdentityId", ), pytest.param( { "https://grass.theorchard.com/other_metadata": {"other": "claim"}, }, None, id="The claims do not have user_metadata", ), ], ) @patch("jwtauth.testing.utils.decode_token") def test_get_bearer_token_identity_uuid( mock_decode_token: MagicMock, token_claims: dict[str, Any], expected_identity_uuid: str | None, ) -> None: """Test get_bearer_token_identity_uuid.""" mock_decode_token.return_value = token_claims assert ( get_bearer_token_identity_uuid("test.token", environment="qa") == expected_identity_uuid ) mock_decode_token.assert_called_once_with( "test.token", environment="qa", ) @patch("jwtauth.testing.utils.decode_token") def test_get_bearer_token_identity_uuid__error( mock_decode_token: MagicMock, ) -> None: """Test when get_bearer_token_identity_uuid raises an assertion error.""" mock_decode_token.return_value = { "https://grass.theorchard.com/user_metadata": { "orchardIdentityId": { "unexpected": "dictionary", } }, } with pytest.raises(AssertionError) as excinfo: get_bearer_token_identity_uuid("test.token", environment="qa") assert "Found unexpected type for orchardIdentityId: " in str( excinfo.value ) @pytest.mark.parametrize( "user_creds, auth0_creds, mfa_token, mock_post_response, expected_token, expected_otp", [ pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ), "test_mfa_token_string", { "access_token": "valid_token_123", "token_type": "Bearer", "expires_in": 3600, }, "valid_token_123", "123456", id="Valid credentials", ), ], ) @patch("jwtauth.testing.utils.TOTP") @patch("jwtauth.testing.utils.requests.post") def test_generate_bearer_jwt_token_mfa( mock_post: MagicMock, mock_totp: MagicMock, user_creds: UserCreds, auth0_creds: Auth0Creds, mfa_token: str, mock_post_response: dict[str, Any], expected_token: str, expected_otp: str, ) -> None: """Test generate_bearer_jwt_token_mfa.""" mock_response = MagicMock( status_code=200, json=MagicMock(return_value=mock_post_response) ) mock_post.return_value = mock_response mock_totp_instance = MagicMock() mock_totp_instance.now.return_value = expected_otp mock_totp.return_value = mock_totp_instance ret = generate_bearer_jwt_token_mfa( user_creds=user_creds, auth0_creds=auth0_creds, mfa_token=mfa_token, ) assert ret == expected_token mock_totp.assert_called_once_with(user_creds.otp_secret_key) mock_totp_instance.now.assert_called_once() mock_post.assert_called_once_with( AUTH_URL, 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": expected_otp, }, ) @patch("jwtauth.testing.utils.TOTP") def test_generate_bearer_jwt_token_mfa__otp_error( mock_totp: MagicMock, ) -> None: """Test generate_bearer_jwt_token_mfa when TOTP raises an exception.""" user_creds = UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="invalid_otp_key", ) auth0_creds = Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ) mfa_token = "test_mfa_token_string" # Mock TOTP to raise an exception mock_totp.side_effect = Exception("Invalid base32 string") with pytest.raises(MfaTokenLookupError) as exc_info: _ = generate_bearer_jwt_token_mfa( user_creds=user_creds, auth0_creds=auth0_creds, mfa_token=mfa_token, ) assert "Failed to generate one-time-password" in str(exc_info.value) mock_totp.assert_called_once_with(user_creds.otp_secret_key) @pytest.mark.parametrize( "user_creds, auth0_creds, mfa_token, mock_post_response, expected_error_msg, should_call_totp", [ pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key=None, ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ), "test_mfa_token_string", { "status_code": 400, "json": {"error": "Should raise before the POST"}, "text": "Should raise before the POST", }, "UserCreds.otp_secret_key is empty", False, id="Missing UserCreds.otp_secret_key", ), pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="null", ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ), "test_mfa_token_string", { "status_code": 400, "json": {"error": "Should raise before the POST"}, "text": "Should raise before the POST", }, "UserCreds.otp_secret_key is empty", False, id="UserCreds.otp_secret_key equals the literal string 'null'", ), pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="asdfbdsdfs", ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ), "", { "status_code": 400, "json": {"error": "Should raise before the POST"}, "text": "Should raise before the POST", }, "mfa_token is empty", False, id="Missing mfa_token", ), pytest.param( UserCreds( email="test@sonymusic-pde.com", password=SecretStr("secret"), otp_secret_key="1234", ), Auth0Creds( auth0_client_id="1234565656", auth0_client_secret="abcdef123456", ), "test_mfa_token_string", { "status_code": 400, "json": {"error": "Invalid OTP or MFA TOKEN."}, "text": "Invalid OTP or MFA TOKEN.", }, "Invalid OTP or MFA TOKEN.", True, id="POST returns a non-200 http status code", ), ], ) @patch("jwtauth.testing.utils.TOTP") @patch("jwtauth.testing.utils.requests.post") def test_generate_bearer_jwt_token_mfa__errors( mock_post: MagicMock, mock_totp: MagicMock, user_creds: UserCreds, auth0_creds: Auth0Creds, mfa_token: str, mock_post_response: dict[str, Any], expected_error_msg: str, should_call_totp: bool, ) -> None: """Test generate_bearer_jwt_token_mfa with errors.""" mock_response = MagicMock( status_code=mock_post_response["status_code"], json=MagicMock(return_value=mock_post_response["json"]), text=mock_post_response["text"], ) mock_post.return_value = mock_response mock_totp_instance = MagicMock() mock_totp_instance.now.return_value = "123456" mock_totp.return_value = mock_totp_instance with pytest.raises(MfaTokenLookupError) as exc_info: _ = generate_bearer_jwt_token_mfa( user_creds=user_creds, auth0_creds=auth0_creds, mfa_token=mfa_token, ) assert expected_error_msg in str(exc_info.value) if should_call_totp: mock_totp.assert_called_once_with(user_creds.otp_secret_key) mock_totp_instance.now.assert_called_once() else: mock_totp.assert_not_called() mock_totp_instance.now.assert_not_called()