"""Tests datasources lifespan methods.""" from unittest.mock import MagicMock, patch import pytest import spotipy from fastapi import FastAPI from owsclient import OwsClient from python_pdp_sdk.backends.authorization_backend import AuthorizationBackend from product_staging import config from product_staging.api.datasources import ( AUTHORIZATION_BACKEND_KEY, OWS_CLIENT_KEY, REDIS_CONNECTOR_KEY, S3_CLIENT_KEY, SFN_CLIENT_KEY, SPOTIFY_CLIENT_KEY, datasources_lifespan, get_authorization_backend, get_ows_client, get_redis_client, get_s3_client, get_sfn_client, get_spotify_client, setup_authorization_backend, ) from product_staging.connectors.redis import RedisConnector @pytest.fixture() def mock_sfn_client() -> MagicMock: """Pytest fixture to create a mock SFN client.""" mock_sfn_client = MagicMock() return mock_sfn_client def test_get_s3_client(mock_s3_client: MagicMock): """Verify method returns the dictionary item at key:S3_CLIENT_KEY.""" with patch.dict( "product_staging.api.datasources.DATA_SOURCES", {S3_CLIENT_KEY: mock_s3_client}, ): assert get_s3_client() == mock_s3_client def test_get_sfn_client(mock_sfn_client: MagicMock): """Verify method returns the dictionary item at key:SFN_CLIENT_KEY.""" with patch.dict( "product_staging.api.datasources.DATA_SOURCES", {SFN_CLIENT_KEY: mock_sfn_client}, ): assert get_sfn_client() == mock_sfn_client @pytest.fixture def mock_redis_connector() -> RedisConnector: """Return a mocked redis connector.""" return MagicMock(spec=RedisConnector) def test_get_redis_client(mock_redis_connector: RedisConnector): """Verify method returns the dictionary item at key:REDIS_CONNECTOR_KEY.""" with patch.dict( "product_staging.api.datasources.DATA_SOURCES", {REDIS_CONNECTOR_KEY: mock_redis_connector}, ): assert get_redis_client() == mock_redis_connector def test_get_ows_client(mock_ows_client: OwsClient) -> None: """Verify method returns the dicitionary item at key:OWS_CLIENT_KEY.""" with patch.dict( "product_staging.api.datasources.DATA_SOURCES", {OWS_CLIENT_KEY: mock_ows_client}, ): assert get_ows_client() == mock_ows_client def test_get_authorization_backend( mock_authorization_backend: AuthorizationBackend, ) -> None: """Verify method returns the dictionary item at key:AUTHORIZATION_BACKEND_KEY.""" with patch.dict( "product_staging.api.datasources.DATA_SOURCES", {AUTHORIZATION_BACKEND_KEY: mock_authorization_backend}, ): assert get_authorization_backend() == mock_authorization_backend def test_setup_authorization_backend( mock_ows_client: OwsClient, mock_authorization_backend ) -> None: """Verify method composes OwsPdpClient into PdpAuthorizationBackend.""" mock_ows_pdp_client = MagicMock() with ( patch( "product_staging.api.datasources.OwsPdpClient", return_value=mock_ows_pdp_client, ) as ows_pdp_client_cls, patch( "product_staging.api.datasources.PdpAuthorizationBackend", return_value=mock_authorization_backend, ) as authorization_backend_cls, ): result = setup_authorization_backend(mock_ows_client) ows_pdp_client_cls.assert_called_once_with(mock_ows_client) authorization_backend_cls.assert_called_once_with(mock_ows_pdp_client) assert result == mock_authorization_backend async def test_datasources_lifespan_sets_authorization_backend( mock_s3_client, mock_ows_client: OwsClient, mock_authorization_backend, ) -> None: """Verify lifespan registers authorization backend in DATA_SOURCES.""" mock_s3_client_cm = MagicMock() mock_s3_client_cm.__aenter__.return_value = mock_s3_client mock_redis_connector = MagicMock(spec=RedisConnector) mock_sfn_client = MagicMock() mock_lambda_client = MagicMock() mock_ecs_client = MagicMock() mock_ec2_client = MagicMock() mock_sqs_client = MagicMock() with ( patch("product_staging.api.datasources.Session") as session_cls, patch( "product_staging.api.datasources.RedisConnector", return_value=mock_redis_connector, ) as redis_connector_cls, patch( "product_staging.api.datasources.boto3.client", side_effect=[ mock_sfn_client, mock_lambda_client, mock_ecs_client, mock_ec2_client, mock_sqs_client, ], ), patch( "product_staging.api.datasources.OwsClient", return_value=mock_ows_client, ), patch( "product_staging.api.datasources.setup_authorization_backend", return_value=mock_authorization_backend, ) as setup_authorization_backend_mock, patch("product_staging.api.datasources.shutdown_s3_client") as shutdown_s3_mock, patch( "product_staging.api.datasources.shutdown_redis_connector" ) as shutdown_redis_mock, patch( "product_staging.api.datasources.shutdown_sfn_client" ) as shutdown_sfn_mock, patch( "product_staging.api.datasources.shutdown_ecs_client" ) as shutdown_ecs_mock, patch( "product_staging.api.datasources.shutdown_ec2_client" ) as shutdown_ec2_mock, ): session_cls.return_value.client.return_value = mock_s3_client_cm async with datasources_lifespan(FastAPI()) as sources: assert sources[AUTHORIZATION_BACKEND_KEY] == mock_authorization_backend assert get_authorization_backend() == mock_authorization_backend setup_authorization_backend_mock.assert_called_once_with(mock_ows_client) redis_connector_cls.assert_called_once_with( redis_host=config.REDIS_HOST, redis_port=config.REDIS_PORT, use_redis_cache=config.CACHE_USE_REDIS, ) shutdown_s3_mock.assert_called_once_with(mock_s3_client_cm) shutdown_redis_mock.assert_called_once_with(redis_connector=mock_redis_connector) shutdown_sfn_mock.assert_called_once_with(mock_sfn_client) shutdown_ecs_mock.assert_called_once_with(mock_ecs_client) shutdown_ec2_mock.assert_called_once_with(mock_ec2_client) def test_get_spotify_client(): """Verify method returns the dictionary item at key:SPOTIFY_CLIENT_KEY.""" mock_spotify_client = MagicMock(spec=spotipy.Spotify) with patch.dict( "product_staging.api.datasources.DATA_SOURCES", {SPOTIFY_CLIENT_KEY: mock_spotify_client}, ): assert get_spotify_client() == mock_spotify_client