import datetime from unittest import mock import pytest from dmp_workflows import config from dmp_workflows.hooks.jwt import OwsJwtHook from pytest_mock import MockerFixture MOCK_TOKEN_VALUE = "mock_token" @pytest.fixture def airflow_variable(mocker: MockerFixture) -> mock.MagicMock: return mocker.patch("dmp_workflows.hooks.jwt.Variable", autospec=True) @pytest.fixture def secrets_manager_hook(mocker: MockerFixture) -> mock.MagicMock: def get_secret(secret_name: str) -> str: if secret_name == config.M2M_TOKEN_SECRET_KEY_NAME: return MOCK_TOKEN_VALUE elif secret_name == config.M2M_TOKEN_SECRET_EXPIRY_KEY_NAME: # https://github.com/theorchard/lambda-jwt-refresh/blob/master/lambda/populate_m2m_jwt/index.py#L77-L79 return str(datetime.datetime.utcnow() + datetime.timedelta(hours=1)) else: raise ValueError(secret_name) hook_mock = mock.MagicMock() hook_mock.get_secret = get_secret return mocker.patch( "dmp_workflows.hooks.jwt.SecretsManagerHook", return_value=hook_mock ) def test_get_jwt_from_cache(airflow_variable: mock.MagicMock) -> None: token_expires_at_iso = datetime.datetime.utcnow().isoformat() airflow_variable.get.return_value = { "token": MOCK_TOKEN_VALUE, "expires_at": token_expires_at_iso, } jwt = OwsJwtHook._get_jwt_from_cache() assert jwt is not None assert jwt.token == MOCK_TOKEN_VALUE assert jwt.expires_at.isoformat() == token_expires_at_iso def test_get_jwt_from_cache_miss(airflow_variable: mock.MagicMock) -> None: airflow_variable.get.return_value = None jwt = OwsJwtHook._get_jwt_from_cache() assert jwt is None def test_get_jwt_from_secrets_manager(secrets_manager_hook: mock.MagicMock) -> None: jwt = OwsJwtHook._get_jwt_from_secrets_manager() assert jwt.token is not None assert jwt.expires_at > datetime.datetime.utcnow() def test_get_jwt_token_cache_not_expired(airflow_variable: mock.MagicMock) -> None: airflow_variable.get.return_value = { "token": MOCK_TOKEN_VALUE, "expires_at": ( datetime.datetime.utcnow() + datetime.timedelta(hours=1) ).isoformat(), } jwt_token = OwsJwtHook.get_jwt_token() assert jwt_token == MOCK_TOKEN_VALUE def test_get_jwt_token_cache_expired( airflow_variable: mock.MagicMock, secrets_manager_hook: mock.MagicMock ) -> None: airflow_variable.get.return_value = { "token": MOCK_TOKEN_VALUE, "expires_at": ( datetime.datetime.utcnow() - datetime.timedelta(hours=1) ).isoformat(), } jwt_token = OwsJwtHook.get_jwt_token() assert jwt_token == MOCK_TOKEN_VALUE