"""Unit tests for the demographics logic layer.""" import datetime from unittest.mock import MagicMock, patch import pytest from analytics.api import app from analytics.connectors import redis as redis_connector from analytics.logic import demographics QUERY_MOCK_PATH = "analytics.logic.demographics.TrackDemographics" GP_QUERY_MOCK_PATH = "analytics.logic.demographics.GlobalParticipantDemographics" @pytest.fixture(autouse=True) def _flush_cache(): """Flush the fake redis cache between tests — the logic function is wrapped with @cache_in_redis.""" redis_connector.client.flushall() yield redis_connector.client.flushall() @pytest.fixture def mock_row(): """Mock row returned by the query class (dict-like).""" return { "_17": 5214236, "_18_22": 18016044, "_23_27": 11165474, "_28_34": 5969109, "_35_44": 2942418, "_45_59": 1606414, "_60_": 436636, "UA": 77677, "M": 8273548, "F": 34247706, "UG": 2906754, } @pytest.fixture def mock_row_zero_streams(): """Mock row with zero values.""" return { "_17": 0, "_18_22": 0, "_23_27": 0, "_28_34": 0, "_35_44": 0, "_45_59": 0, "_60_": 0, "UA": 0, "M": 0, "F": 0, "UG": 0, } @pytest.fixture def expected_payload(): """Return expected payload from logic.""" return { "isrc": "TEST", "demographics": { "age": { "_17": 5214236, "_18_22": 18016044, "_23_27": 11165474, "_28_34": 5969109, "_35_44": 2942418, "_45_59": 1606414, "_60_": 436636, "UA": 77677, }, "gender": {"M": 8273548, "F": 34247706, "UG": 2906754}, }, "sources": [{"id": 286, "name": "Spotify"}, {"id": 1, "name": "Apple Music"}], } @pytest.fixture def expected_payload_zero_streams(): """Return expected payload for zero-data rows.""" return { "isrc": "TEST", "demographics": { "age": { "_17": 0, "_18_22": 0, "_23_27": 0, "_28_34": 0, "_35_44": 0, "_45_59": 0, "_60_": 0, "UA": 0, }, "gender": {"M": 0, "F": 0, "UG": 0}, }, "sources": [{"id": 286, "name": "Spotify"}, {"id": 1, "name": "Apple Music"}], } @pytest.fixture def mock_add_outage_error_to_stores(): """Mock add_outage_error_to_stores.""" with patch( "analytics.logic.demographics.add_outage_error_to_stores" ) as add_outage_error_to_stores: add_outage_error_to_stores.return_value = [ {"id": 286, "name": "Spotify"}, {"id": 1, "name": "Apple Music"}, ] yield add_outage_error_to_stores @pytest.fixture def mock_store_availability(): """Mock store_availability.""" with patch("analytics.logic.demographics.store_availability") as store_availability: store_availability.get_demographic_store_ids.return_value = [1, 286] store_availability.get_demographic_sources.return_value = [ {"id": 286, "name": "Spotify"}, {"id": 1, "name": "Apple Music"}, ] yield store_availability @pytest.fixture def mock_get_max_available_date_success(): """Mock successful get_max_available_date.""" with patch( "analytics.utils.data_availability.get_max_available_date" ) as get_max_available_date: get_max_available_date.return_value = "2018-07-06" yield get_max_available_date @pytest.fixture def mock_query_success(mock_row): """Mock the TrackDemographics query class to return one row.""" with patch(QUERY_MOCK_PATH) as query_cls: query_cls.return_value.execute.return_value = [mock_row] yield query_cls @pytest.fixture def mock_query_zero(mock_row_zero_streams): """Mock the TrackDemographics query class to return a zero-valued row.""" with patch(QUERY_MOCK_PATH) as query_cls: query_cls.return_value.execute.return_value = [mock_row_zero_streams] yield query_cls @pytest.fixture def mock_query_empty(): """Mock the TrackDemographics query class to return no rows.""" with patch(QUERY_MOCK_PATH) as query_cls: query_cls.return_value.execute.return_value = [] yield query_cls @pytest.fixture def mock_gp_query_success(mock_row): """Mock the GlobalParticipantDemographics query class to return one row.""" with patch(GP_QUERY_MOCK_PATH) as query_cls: query_cls.return_value.execute.return_value = [mock_row] yield query_cls @pytest.fixture def permissions_for_label(): """Fixture with a representative permissions dict.""" return { "permission_label_ids": [1234], "permission_subaccount_ids": [], "permission_artist_ids": [], "permission_label_participant_ids": [], "permission_feed_ids": [], } def _isrc_query_params(**overrides): params = { "isrc": "TEST", "query_type": "isrc", "countries": [], "store_ids": [], "start_date": None, "end_date": None, "distributors": ["theorchard", "sme", "awal"], } params.update(overrides) return params def test_get_demographics_calls_query_with_isrc_params( mock_query_success, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, permissions_for_label, ): """ISRC query passes expected params to TrackDemographics.""" with app.test_request_context(): demographics.get_demographics(_isrc_query_params(), permissions_for_label) kwargs = mock_query_success.call_args.args[0] assert kwargs["query_type"] == "isrc" assert kwargs["isrc"] == "TEST" assert kwargs["table"] == ( "V_STREAMS_DEMOGRAPHICS_BY_TRACK_COUNTRY_FEED_DISTRIBUTOR_DAILY" ) # date fallback fills in start_date / end_date assert kwargs["start_date"] is not None assert kwargs["end_date"] is not None # store_ids expanded to demographic store ids when caller passed empty assert kwargs["store_ids"] == [1, 286] assert kwargs["distributors"] == ["theorchard", "sme", "awal"] def test_get_demographics_calls_gp_query( mock_gp_query_success, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, permissions_for_label, ): """global_participant_id query uses the GP class and GP view.""" params = _isrc_query_params( query_type="global_participant_id", global_participant_id="gp-test", ) params.pop("isrc") with app.test_request_context(): demographics.get_demographics(params, permissions_for_label) kwargs = mock_gp_query_success.call_args.args[0] assert kwargs["query_type"] == "global_participant_id" assert kwargs["global_participant_id"] == "gp-test" assert kwargs["table"] == ( "V_STREAMS_DEMOGRAPHICS_BY_PARTICIPANT_FEED_DISTRIBUTOR_DAILY" ) def test_get_demographics_gp_switches_to_country_view_when_filtered( mock_gp_query_success, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, permissions_for_label, ): """With countries filter, GP query targets the COUNTRY view.""" params = _isrc_query_params( query_type="global_participant_id", global_participant_id="gp-test", countries=["US", "MX"], ) params.pop("isrc") with app.test_request_context(): demographics.get_demographics(params, permissions_for_label) kwargs = mock_gp_query_success.call_args.args[0] assert kwargs["table"] == ( "V_STREAMS_DEMOGRAPHICS_BY_PARTICIPANT_COUNTRY_FEED_DISTRIBUTOR_DAILY" ) def test_get_demographics_uses_given_dates( mock_query_success, mock_store_availability, mock_add_outage_error_to_stores, permissions_for_label, ): """When dates are provided, they are not overridden.""" params = _isrc_query_params( start_date="2019-11-01", end_date="2019-12-01", ) with app.test_request_context(): demographics.get_demographics(params, permissions_for_label) kwargs = mock_query_success.call_args.args[0] assert kwargs["start_date"] == "2019-11-01" assert kwargs["end_date"] == "2019-12-01" def test_get_demographics_drops_start_date_on_all_time( mock_query_success, mock_store_availability, mock_add_outage_error_to_stores, permissions_for_label, ): """start_date='ALL_TIME' is dropped so the query has no lower bound. The graphql layer sends ?days=0&start_date=HIGHWATERMARK for the all-time case; `_get_global_filters` resolves that to ``start_date='ALL_TIME', end_date=``. Marshmallow's ``fields.Date`` would reject the string, so the logic must strip it before handing params to the query class. """ end_date = datetime.date(2025, 1, 1) params = _isrc_query_params( start_date="ALL_TIME", end_date=end_date, ) with app.test_request_context(): demographics.get_demographics(params, permissions_for_label) kwargs = mock_query_success.call_args.args[0] assert kwargs["start_date"] is None assert kwargs["end_date"] == end_date def test_get_demographics_intersects_store_ids( mock_query_success, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, permissions_for_label, ): """User-supplied store_ids intersect with demographic store ids.""" with app.test_request_context(): demographics.get_demographics( _isrc_query_params(store_ids=[1, 999]), permissions_for_label, ) kwargs = mock_query_success.call_args.args[0] assert kwargs["store_ids"] == [1] def test_get_demographics_success_payload( mock_query_success, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, expected_payload, permissions_for_label, ): """Full payload shape when query returns real data.""" with app.test_request_context(): result = demographics.get_demographics( _isrc_query_params(), permissions_for_label ) assert result == expected_payload def test_get_demographics_zero_streams( mock_query_zero, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, expected_payload_zero_streams, permissions_for_label, ): """Zero-valued row is propagated untouched to the response body.""" with app.test_request_context(): result = demographics.get_demographics( _isrc_query_params(), permissions_for_label ) assert result == expected_payload_zero_streams def test_get_demographics_empty_rows( mock_query_empty, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, permissions_for_label, ): """When query returns no rows, demographics stays empty.""" with app.test_request_context(): result = demographics.get_demographics( _isrc_query_params(), permissions_for_label ) assert result["demographics"] == {"age": {}, "gender": {}} def test_get_demographics_short_circuits_on_empty_store_intersection( mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, permissions_for_label, ): """Store intersection empty -> skip query entirely.""" with patch(QUERY_MOCK_PATH) as query_cls: with app.test_request_context(): result = demographics.get_demographics( _isrc_query_params(store_ids=[9999]), permissions_for_label, ) query_cls.assert_not_called() assert result["demographics"] == {"age": {}, "gender": {}} def test_get_demographics_uses_row_mapping_from_sqlalchemy( mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, expected_payload, permissions_for_label, mock_row, ): """SQLAlchemy Row objects are unpacked via ._mapping.""" row_obj = MagicMock() row_obj._mapping = mock_row with patch(QUERY_MOCK_PATH) as query_cls: query_cls.return_value.execute.return_value = [row_obj] with app.test_request_context(): result = demographics.get_demographics( _isrc_query_params(), permissions_for_label ) assert result == expected_payload @pytest.mark.parametrize( "query_type,expected_flag", [ ("isrc", True), ("global_participant_id", False), ], ) def test_transfer_product_ownership_flag_only_for_isrc_path( mock_query_success, mock_gp_query_success, mock_store_availability, mock_add_outage_error_to_stores, mock_get_max_available_date_success, permissions_for_label, query_type, expected_flag, ): """Flag is only threaded for the Song-page (isrc) demographics path.""" if query_type == "isrc": params = _isrc_query_params() query_cls_mock = mock_query_success else: params = _isrc_query_params( query_type="global_participant_id", global_participant_id="gp-test", ) params.pop("isrc") query_cls_mock = mock_gp_query_success params["transfer_product_ownership_enabled"] = True with app.test_request_context(): demographics.get_demographics(params, permissions_for_label) kwargs = query_cls_mock.call_args.args[0] assert kwargs["transfer_product_ownership_enabled"] is expected_flag