from unittest.mock import AsyncMock import pandas as pd import pytest from src.backend.constants import SnowFlakeColumns from src.backend.logic.look_data import look_data from tests_backend.conftest import TEST_LABEL_ID LOOK_DATA_FETCH_FUNCTIONS = [ look_data.fetch_sr, look_data.fetch_sr_ugc1, look_data.fetch_sr_ugc2, look_data.fetch_sr_ugc3, look_data.fetch_mv, look_data.fetch_at, ] # noinspection PyUnresolvedReferences @pytest.mark.asyncio @pytest.mark.parametrize("include_mv", [True, False]) @pytest.mark.parametrize("include_at", [True, False]) async def test_fetch(mocker, include_mv, include_at): """Test fetch function using mocks.""" for task in LOOK_DATA_FETCH_FUNCTIONS: mock = AsyncMock() mock.__name__ = task.__name__ mocker.patch.object(look_data, task.__name__, mock) result = await look_data.fetch( TEST_LABEL_ID, include_mv=include_mv, include_at=include_at ) assert isinstance(result, dict) # Assert that SR functions are always called once. sr_functions = [ look_data.fetch_sr, look_data.fetch_sr_ugc1, look_data.fetch_sr_ugc2, look_data.fetch_sr_ugc3, ] for function in sr_functions: assert function.call_count == 1 # Assert that the MV and AT functions are called only if requested. assert look_data.fetch_mv.call_count == int(include_mv) assert look_data.fetch_at.call_count == int(include_at) def test_preprocess_at3(): sample_df = pd.DataFrame( { SnowFlakeColumns.UPC: [ "prefix.mock_upc1", "prefix.mock_upc2", ], SnowFlakeColumns.VTR_COUNTRIES: ["RU,CN,IT", "NO"], SnowFlakeColumns.RTR_COUNTRIES: ["ES,IT", None], SnowFlakeColumns.STR_COUNTRIES: [None, "NO,IT,CL"], } ) result = look_data._preprocess_at3(sample_df) assert sorted(result.loc[0, SnowFlakeColumns.ABBRIVATION]) == sorted( ["RU", "CN", "IT", "ES"] ) assert sorted(result.loc[1, SnowFlakeColumns.ABBRIVATION]) == sorted( ["IT", "NO", "CL"] ) def test_col_preprocessor(): df = pd.DataFrame( columns=["prefix.a", "prefix.B", "prefix.C", "prefix.D", "e", "F"] ) expected_columns = ["a", "b", "c", "d", "e", "f"] assert list(look_data._col_preprocessor(df).columns) == expected_columns def test_looktask_has_default_preprocessor(): task = look_data.LookTask("test", "test") assert task.preprocessor class TestFormatters: @pytest.mark.parametrize( "value,expected", [ (" a , b, c,c", {"a", "b", "c"}), ("", None), (None, None), ], ) def test_formatter_1(self, value, expected): assert look_data.Formatters.formatter_1(value) == expected @pytest.mark.parametrize( "value,expected", [ (" a b c c ", {"a", "b", "c"}), ("", None), (None, None), ], ) def test_formatter_2(self, value, expected): assert look_data.Formatters.formatter_2(value) == expected def test_snowflake_request_factory(): result = look_data._snowflake_request_factory() assert isinstance(result, look_data.SnowflakeRequest)