"""Unit tests for OwsClient and AsyncOwsClient.""" import json import logging from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import fakeredis import pytest from owsclient.m2m.base import ( DEFAULT_CACHE_KEY, M2M_JWT_ACCESS_TOKEN_SECRET_NAME, AsyncM2MTokenManager, BaseM2MTokenManager, M2MToken, M2MTokenManager, ) from owsclient.protocols import AsyncCache, Cache, SecretsManager TEST_M2M_JWT_ACCESS_TOKEN_STR = json.dumps( { "expires_at": "2099-07-22 16:55:30.13455+0000", "token": "TEST_TOKEN", "client_credentials_secret_arn": "arn:aws:secretsmanager:us-east-1:1234:secret:qa/lambda-test-m2m-client/M2M_AUTH0_CLIENT_CREDENTIALS-H3WvVs", # noqa: E501 } ) TEST_EXPIRED_M2M_JWT_ACCESS_TOKEN_STR = json.dumps( { "expires_at": "2020-01-01 00:00:00.00000+0000", "token": "EXPIRED_TEST_TOKEN", "client_credentials_secret_arn": "arn:aws:secretsmanager:us-east-1:1234:secret:qa/lambda-test-m2m-client/M2M_AUTH0_CLIENT_CREDENTIALS-H3WvVs", # noqa: E501 } ) @pytest.fixture() def mock_secrets_manager() -> MagicMock: """Create a mock secrets manager instance.""" instance = MagicMock(spec=SecretsManager) instance.get_secret = MagicMock(return_value=TEST_M2M_JWT_ACCESS_TOKEN_STR) return instance def test_base_m2m_token_manager(mock_secrets_manager: MagicMock) -> None: """Test SecretsManagerTokenManager.""" obj = BaseM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", ) secret_name = obj.generate_secret_name() assert obj.service_name == "my-app" assert obj.environment == "qa" assert secret_name == f"qa/my-app/{M2M_JWT_ACCESS_TOKEN_SECRET_NAME}" @patch.object(M2MTokenManager, "_get_token_from_cache") @patch.object(M2MTokenManager, "_get_token_payload_from_secret_manager") def test_m2m_token_manager_get_token_string_empty_cache( mock__get_token_payload_from_secret_manager: MagicMock, mock__get_token_from_cache: MagicMock, mock_secrets_manager: MagicMock, ) -> None: """Test M2MTokenManager.get_token_string() when cache is empty.""" mock_cache = MagicMock(spec=Cache) mock__get_token_from_cache.return_value = None mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_cache, ) token_string = obj.get_token_string() assert token_string == "TEST_TOKEN" mock__get_token_from_cache.assert_called_once() mock_cache.set.assert_called_once_with( "_m2m_token", value=TEST_M2M_JWT_ACCESS_TOKEN_STR, timeout=0 ) mock__get_token_payload_from_secret_manager.assert_called_once() @patch.object(M2MTokenManager, "_get_token_from_cache") @patch.object(M2MTokenManager, "_get_token_payload_from_secret_manager") def test_m2m_token_manager_get_token_string_empty_cache_uses_custom_cache_key( mock__get_token_payload_from_secret_manager: MagicMock, mock__get_token_from_cache: MagicMock, mock_secrets_manager: MagicMock, ) -> None: """Test M2MTokenManager.get_token_string() uses custom cache key.""" mock_cache = MagicMock(spec=Cache) cache_key = "custom_cache_key_overrides" mock__get_token_from_cache.return_value = None mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_cache, cache_key=cache_key, ) token_string = obj.get_token_string() assert token_string == "TEST_TOKEN" mock__get_token_from_cache.assert_called_once() mock_cache.set.assert_called_once_with( cache_key, value=TEST_M2M_JWT_ACCESS_TOKEN_STR, timeout=0 ) mock__get_token_payload_from_secret_manager.assert_called_once() @patch.object(M2MTokenManager, "_get_token_from_cache") @patch.object(M2MTokenManager, "_get_token_payload_from_secret_manager") def test_m2m_token_manager_get_token_string_valid_from_cache( mock__get_token_payload_from_secret_manager: MagicMock, mock__get_token_from_cache: MagicMock, mock_secrets_manager: MagicMock, ) -> None: """Test M2MTokenManager.get_token_string() when cache has valid token.""" mock_cache = MagicMock(spec=Cache) expected_jwt_token = "this is an unexpired jwt" mock_token = MagicMock(spec=M2MToken) mock_token.token = expected_jwt_token mock_token.is_expired.return_value = False mock__get_token_from_cache.return_value = mock_token mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_cache, ) actual = obj.get_token_string() assert actual == expected_jwt_token mock__get_token_from_cache.assert_called_once() mock_token.is_expired.assert_called_once_with(leeway_seconds=60) mock_cache.set.assert_not_called() mock__get_token_payload_from_secret_manager.assert_not_called() @patch.object(M2MTokenManager, "_get_token_from_cache") @patch.object(M2MTokenManager, "_get_token_payload_from_secret_manager") def test_m2m_token_manager_get_token_string_expired_from_cache( mock__get_token_payload_from_secret_manager: MagicMock, mock__get_token_from_cache: MagicMock, mock_secrets_manager: MagicMock, ) -> None: """Test M2MTokenManager.get_token_string() when cache has expired token.""" mock_cache = MagicMock(spec=Cache) unexpected_jwt_token = "this is an expired jwt" mock_token = MagicMock(spec=M2MToken) mock_token.token = unexpected_jwt_token mock_token.is_expired.return_value = True mock__get_token_from_cache.return_value = mock_token mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_cache, ) actual = obj.get_token_string() assert actual == "TEST_TOKEN" mock__get_token_from_cache.assert_called_once() mock_token.is_expired.assert_called_once_with(leeway_seconds=60) mock_cache.set.assert_called_once_with( "_m2m_token", value=TEST_M2M_JWT_ACCESS_TOKEN_STR, timeout=0 ) mock__get_token_payload_from_secret_manager.assert_called_once() @patch.object(M2MTokenManager, "_get_token_from_cache") @patch.object(M2MTokenManager, "_get_token_payload_from_secret_manager") def test_m2m_token_manager_get_token_string_warns_on_expired_from_secrets_manager( mock__get_token_payload_from_secret_manager: MagicMock, mock__get_token_from_cache: MagicMock, mock_secrets_manager: MagicMock, caplog: pytest.LogCaptureFixture, ) -> None: """Test M2MTokenManager.get_token_string() warns when secrets manager has expired token.""" mock_cache = MagicMock(spec=Cache) mock__get_token_from_cache.return_value = None mock__get_token_payload_from_secret_manager.return_value = ( TEST_EXPIRED_M2M_JWT_ACCESS_TOKEN_STR ) obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_cache, ) with caplog.at_level(logging.WARNING): token_string = obj.get_token_string() # Should still return the token even though it's expired assert token_string == "EXPIRED_TEST_TOKEN" # Verify warning was logged assert len(caplog.records) == 1 assert caplog.records[0].levelname == "WARNING" assert "Fetched an expired token from secrets manager" in caplog.records[0].message assert "qa/my-app/M2M_JWT_ACCESS_TOKEN" in caplog.records[0].message mock__get_token_from_cache.assert_called_once() # Expired token should NOT be cached mock_cache.set.assert_not_called() mock__get_token_payload_from_secret_manager.assert_called_once() def test_m2m_token_manager__get_token_payload_from_secret_manager( mock_secrets_manager: MagicMock, ) -> None: """Test M2MTokenManager._get_token_payload_from_secret_manager().""" obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", ) token_string = obj._get_token_payload_from_secret_manager() assert token_string == TEST_M2M_JWT_ACCESS_TOKEN_STR mock_secrets_manager.get_secret.assert_called_once_with( secret_name=f"qa/my-app/{M2M_JWT_ACCESS_TOKEN_SECRET_NAME}" ) def test_m2m_token_manager__get_token_payload_dict_from_secret_manager( mock_secrets_manager: MagicMock, ) -> None: """Test M2MTokenManager._get_token_payload_from_secret_manager() returns str. SecretsManager.get_secret can return str or dict, this asserts that _get_token_payload_from_secret_manager always returns str. """ obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", ) mock_secrets_manager.get_secret.return_value = json.loads( TEST_M2M_JWT_ACCESS_TOKEN_STR ) token_string = obj._get_token_payload_from_secret_manager() assert token_string == TEST_M2M_JWT_ACCESS_TOKEN_STR mock_secrets_manager.get_secret.assert_called_once_with( secret_name=f"qa/my-app/{M2M_JWT_ACCESS_TOKEN_SECRET_NAME}" ) class MockCache: """MockCache.""" def __init__(self) -> None: """Initialize fakeredis.""" self._cache = fakeredis.FakeRedis() def get(self, key: str) -> Any: """Get key value from cache.""" return self._cache.get(key) def set(self, key: str, value: Any, *, timeout: int | None = None) -> bool | None: """Set key with value to cache.""" self._cache.set(key, value) return True def get_mock_cache(key: str | None = None, val: Any = None) -> Cache: """Return a MockCache. Use key and val arguments as helpers to set the cache. """ cache = MockCache() if key and val: cache.set(key, val) return cache @pytest.mark.parametrize( "mock_cache, expected", [ pytest.param({}, None, id="An empty cache dict returns None"), pytest.param( {"not_the_key": 123}, None, id="An cache dict without the key returns none" ), pytest.param( {"_m2m_token": None}, None, id="When the cached value is None, None is returned", ), pytest.param( {"_m2m_token": 123}, None, id="When the cached value is an integer, None is returned", ), pytest.param( {"_m2m_token": True}, None, id="When the cached value is a boolean, None is returned", ), pytest.param( {"_m2m_token": "some string"}, None, id="When the cached value cannot be validated as an M2MToken, None is returned", ), pytest.param( {"_m2m_token": M2MToken.model_validate_json(TEST_M2M_JWT_ACCESS_TOKEN_STR)}, None, id="When the cached value is an M2MToken, None is returned", ), pytest.param( {"_m2m_token": TEST_M2M_JWT_ACCESS_TOKEN_STR}, M2MToken.model_validate_json(TEST_M2M_JWT_ACCESS_TOKEN_STR), id="The cached value is validated as an M2MToken", ), pytest.param( get_mock_cache(), None, id="A cache obj can be used, even if it does not contain the cached value", ), pytest.param( get_mock_cache("_m2m_token", TEST_M2M_JWT_ACCESS_TOKEN_STR), M2MToken.model_validate_json(TEST_M2M_JWT_ACCESS_TOKEN_STR), id="A cache obj can be used and will validate the value as an M2MToken", ), ], ) def test_m2m_token_manager__get_token_from_cache( mock_secrets_manager: MagicMock, mock_cache: Cache | dict[str, Any], expected: M2MToken | None, ) -> None: """Test M2MTokenManager._get_token_from_cache().""" obj = M2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_cache, ) actual = obj._get_token_from_cache() assert actual == expected @pytest.mark.anyio @patch.object(AsyncM2MTokenManager, "_get_token_from_cache") @patch.object(AsyncM2MTokenManager, "_get_token_payload_from_secret_manager") async def test_async_m2m_token_manager_get_token_string_empty_cache( mock__get_token_payload_from_secret_manager: AsyncMock, mock__get_token_from_cache: AsyncMock, mock_secrets_manager: MagicMock, ) -> None: """Test AsyncM2MTokenManager.get_token_string() when cache is empty.""" mock_async_cache = MagicMock(spec=AsyncCache) mock__get_token_from_cache.return_value = None mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_async_cache, ) token_string = await obj.get_token_string() assert token_string == "TEST_TOKEN" mock__get_token_from_cache.assert_awaited_once() mock__get_token_payload_from_secret_manager.assert_awaited_once() mock_async_cache.set.assert_awaited_once_with( DEFAULT_CACHE_KEY, value=TEST_M2M_JWT_ACCESS_TOKEN_STR, ttl=0 ) @pytest.mark.anyio @patch.object(AsyncM2MTokenManager, "_get_token_from_cache") @patch.object(AsyncM2MTokenManager, "_get_token_payload_from_secret_manager") async def test_async_m2m_token_manager_get_token_string_empty_cache_uses_custom_cache_key( mock__get_token_payload_from_secret_manager: AsyncMock, mock__get_token_from_cache: AsyncMock, mock_secrets_manager: MagicMock, ) -> None: """Test AsyncM2MTokenManager.get_token_string() uses custom cache key.""" mock_async_cache = MagicMock(spec=AsyncCache) cache_key = "custom_cache_key" mock__get_token_from_cache.return_value = None mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_async_cache, cache_key=cache_key, ) token_string = await obj.get_token_string() assert token_string == "TEST_TOKEN" mock__get_token_from_cache.assert_awaited_once_with() mock__get_token_payload_from_secret_manager.assert_awaited_once() mock_async_cache.set.assert_awaited_once_with( cache_key, value=TEST_M2M_JWT_ACCESS_TOKEN_STR, ttl=0 ) @pytest.mark.anyio @patch.object(AsyncM2MTokenManager, "_get_token_from_cache") @patch.object(AsyncM2MTokenManager, "_get_token_payload_from_secret_manager") async def test_async_m2m_token_manager_get_token_string_valid_from_cache( mock__get_token_payload_from_secret_manager: AsyncMock, mock__get_token_from_cache: AsyncMock, mock_secrets_manager: MagicMock, ) -> None: """Test AsyncM2MTokenManager.get_token_string() when cache has valid token.""" mock_async_cache = MagicMock(spec=AsyncCache) expected_jwt_token = "this is an unexpired jwt" mock_token = MagicMock(spec=M2MToken) mock_token.token = expected_jwt_token mock_token.is_expired.return_value = False mock__get_token_from_cache.return_value = mock_token mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_async_cache, ) token_string = await obj.get_token_string() assert token_string == expected_jwt_token mock__get_token_from_cache.assert_awaited_once() mock_token.is_expired.assert_called_once_with(leeway_seconds=60) mock_async_cache.set.assert_not_awaited() mock__get_token_payload_from_secret_manager.assert_not_awaited() @pytest.mark.anyio @patch.object(AsyncM2MTokenManager, "_get_token_from_cache") @patch.object(AsyncM2MTokenManager, "_get_token_payload_from_secret_manager") async def test_async_m2m_token_manager_get_token_string_expired_from_cache( mock__get_token_payload_from_secret_manager: AsyncMock, mock__get_token_from_cache: AsyncMock, mock_secrets_manager: MagicMock, ) -> None: """Test AsyncM2MTokenManager.get_token_string() when cache has expired token.""" mock_async_cache = MagicMock(spec=AsyncCache) expected_jwt_token = "this is an unexpired jwt" mock_token = MagicMock(spec=M2MToken) mock_token.token = expected_jwt_token mock_token.is_expired.return_value = True mock__get_token_from_cache.return_value = mock_token mock__get_token_payload_from_secret_manager.return_value = ( TEST_M2M_JWT_ACCESS_TOKEN_STR ) obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_async_cache, ) token_string = await obj.get_token_string() assert token_string == "TEST_TOKEN" mock__get_token_from_cache.assert_awaited_once() mock_token.is_expired.assert_called_once_with(leeway_seconds=60) mock_async_cache.set.assert_awaited_once_with( "_m2m_token", value=TEST_M2M_JWT_ACCESS_TOKEN_STR, ttl=0 ) mock__get_token_payload_from_secret_manager.assert_awaited_once() @pytest.mark.anyio @patch.object(AsyncM2MTokenManager, "_get_token_from_cache") @patch.object(AsyncM2MTokenManager, "_get_token_payload_from_secret_manager") async def test_async_m2m_token_manager_get_token_string_warns_on_expired_from_secrets_manager( mock__get_token_payload_from_secret_manager: AsyncMock, mock__get_token_from_cache: AsyncMock, mock_secrets_manager: MagicMock, caplog: pytest.LogCaptureFixture, ) -> None: """Test AsyncM2MTokenManager.get_token_string() warns when secrets manager has expired token.""" mock_async_cache = MagicMock(spec=AsyncCache) mock__get_token_from_cache.return_value = None mock__get_token_payload_from_secret_manager.return_value = ( TEST_EXPIRED_M2M_JWT_ACCESS_TOKEN_STR ) obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_async_cache, ) with caplog.at_level(logging.WARNING): token_string = await obj.get_token_string() # Should still return the token even though it's expired assert token_string == "EXPIRED_TEST_TOKEN" # Verify warning was logged assert len(caplog.records) == 1 assert caplog.records[0].levelname == "WARNING" assert "Fetched an expired token from secrets manager" in caplog.records[0].message assert "qa/my-app/M2M_JWT_ACCESS_TOKEN" in caplog.records[0].message mock__get_token_from_cache.assert_awaited_once() # Expired token should NOT be cached mock_async_cache.set.assert_not_awaited() mock__get_token_payload_from_secret_manager.assert_awaited_once() @pytest.mark.anyio async def test_async_m2m_token_manager__get_token_payload_from_secret_manager( mock_secrets_manager: MagicMock, ) -> None: """Test AsyncM2MTokenManager._get_token_payload_from_secret_manager().""" obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", ) token_string = await obj._get_token_payload_from_secret_manager() assert token_string == TEST_M2M_JWT_ACCESS_TOKEN_STR mock_secrets_manager.get_secret.assert_called_once_with( f"qa/my-app/{M2M_JWT_ACCESS_TOKEN_SECRET_NAME}" ) @pytest.mark.anyio async def test_async_m2m_token_manager__get_token_payload_dict_from_secret_manager( mock_secrets_manager: MagicMock, ) -> None: """Test AsyncM2MTokenManager._get_token_payload_from_secret_manager() returns str. SecretsManager.get_secret can return str or dict, this asserts that _get_token_payload_from_secret_manager always returns str. """ obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", ) mock_secrets_manager.get_secret.return_value = json.loads( TEST_M2M_JWT_ACCESS_TOKEN_STR ) token_string = await obj._get_token_payload_from_secret_manager() assert token_string == TEST_M2M_JWT_ACCESS_TOKEN_STR mock_secrets_manager.get_secret.assert_called_once_with( f"qa/my-app/{M2M_JWT_ACCESS_TOKEN_SECRET_NAME}" ) @pytest.mark.anyio @pytest.mark.parametrize( "mock_cache, expected", [ pytest.param({}, None, id="An empty cache dict returns None"), pytest.param( {"not_the_key": 123}, None, id="An cache dict without the key returns none" ), pytest.param( {"_m2m_token": None}, None, id="When the cached value is None, None is returned", ), pytest.param( {"_m2m_token": 123}, None, id="When the cached value is an integer, None is returned", ), pytest.param( {"_m2m_token": True}, None, id="When the cached value is a boolean, None is returned", ), pytest.param( {"_m2m_token": "some string"}, None, id="When the cached value cannot be validated as an M2MToken, None is returned", ), pytest.param( {"_m2m_token": M2MToken.model_validate_json(TEST_M2M_JWT_ACCESS_TOKEN_STR)}, None, id="When the cached value is an M2MToken, None is returned", ), pytest.param( {"_m2m_token": TEST_M2M_JWT_ACCESS_TOKEN_STR}, M2MToken.model_validate_json(TEST_M2M_JWT_ACCESS_TOKEN_STR), id="The cached value is validated as an M2MToken", ), ], ) async def test_async_m2m_token_manager__get_token_from_cache_dict( mock_secrets_manager: MagicMock, mock_cache: dict[str, Any], expected: M2MToken | None, ) -> None: """Test AsyncM2MTokenManager._get_token_from_cache() where cache is a dict.""" obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_cache, ) actual = await obj._get_token_from_cache() assert actual == expected class MockAsyncCache: """MockAsyncCache.""" def __init__(self) -> None: """Initialize fakeredis.""" self._cache = fakeredis.FakeRedis() async def get(self, key: str) -> Any: """Get key value from cache.""" return self._cache.get(key) async def set( self, key: str, value: Any, ttl: Any = None, **kwargs: Any ) -> bool | None: """Set key with value to cache.""" self._cache.set(key, value) return True @pytest.fixture() def mock_async_cache() -> AsyncCache: """Return fixture for AsyncCache.""" return MockAsyncCache() @pytest.mark.anyio @pytest.mark.parametrize( "value, expected", [ pytest.param( 123, None, id="When the cached value is an integer, None is returned", ), pytest.param( "some string", None, id="When the cached value cannot be validated as an M2MToken, None is returned", ), pytest.param( TEST_M2M_JWT_ACCESS_TOKEN_STR, M2MToken.model_validate_json(TEST_M2M_JWT_ACCESS_TOKEN_STR), id="The cached value is validated as an M2MToken", ), ], ) async def test_async_m2m_token_manager__get_token_from_cache( mock_secrets_manager: MagicMock, mock_async_cache: AsyncCache, value: Any, expected: M2MToken | None, ) -> None: """Test AsyncM2MTokenManager._get_token_from_cache() where cache is an AsyncCache.""" await mock_async_cache.set("_m2m_token", value) obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_async_cache, ) actual = await obj._get_token_from_cache() assert actual == expected @pytest.mark.anyio async def test_async_m2m_token_manager__get_token_from_cache_empty( mock_secrets_manager: MagicMock, mock_async_cache: AsyncCache, ) -> None: """Test AsyncM2MTokenManager._get_token_from_cache() where cache is empty AsyncCache.""" obj = AsyncM2MTokenManager( secrets_manager=mock_secrets_manager, environment="qa", service_name="my-app", cache=mock_async_cache, ) actual = await obj._get_token_from_cache() assert not actual