from typing import Any, Coroutine import pytest import pytest_asyncio from fastapi.testclient import TestClient from sqlalchemy import text from syrupy.assertion import SnapshotAssertion from delivery_metadata.api.app import app from delivery_metadata.clients.art_relations import get_distribution_features from delivery_metadata.clients.art_relations.product import ArtRelationsProduct from delivery_metadata.constants import DistributionTypeId @pytest_asyncio.fixture async def create_customer_master_master_distribution_type_table( art_relations_product_mock: ArtRelationsProduct, test_client: TestClient ) -> None: async with app.state.art_relations_connector.db_session( transaction=True, turn_off_foreign_key_checks=True ) as session: await session.execute( text("TRUNCATE TABLE customer_master_master_distribution_type") ) await session.execute( text(""" INSERT INTO customer_master_master_distribution_type ( customer_master_master_id, distribution_type_id, distribution_features_ids ) VALUES ( :customer_master_master_id, :distribution_type_id, :distribution_features_ids ) """), [ { "customer_master_master_id": 123, "distribution_type_id": DistributionTypeId.FULL_TRACK.value, "distribution_features_ids": "1,5,6,8,7,2,4,3", }, { "customer_master_master_id": 123, "distribution_type_id": DistributionTypeId.TONE.value, "distribution_features_ids": "10,9", }, { "customer_master_master_id": 123, "distribution_type_id": 3, # unused distribution type "distribution_features_ids": "11", }, { "customer_master_master_id": 456, "distribution_type_id": DistributionTypeId.FULL_TRACK.value, "distribution_features_ids": "3", }, { "customer_master_master_id": 456, "distribution_type_id": 3, # unused distribution type "distribution_features_ids": "11", }, ], ) @pytest.mark.asyncio async def test_get_distribution_features( create_customer_master_master_distribution_type_table: Coroutine[Any, Any, None], snapshot: SnapshotAssertion, ) -> None: """Test get_distribution_features.""" distribution_features = await get_distribution_features(123) distribution_feature = distribution_features.pop() assert distribution_feature.distribution_feature_ids == snapshot assert distribution_feature.distribution_type_id == DistributionTypeId.FULL_TRACK