"""Tests for SplitLevel hierarchy classes.""" from itertools import count from unittest.mock import patch from trigger.constants import SplitType from trigger.schemas import SplitRow from trigger.split_levels import SubaccountSplitLevel, TrackSplitLevel @patch("trigger.split_levels.rds.get_track_splits_for_report_run") def test_get_splits_delegates_to_rds(mock_get_splits): rows = [SplitRow(1, SplitType.TRACK, "tuid-A", 0.5, 42, "NET")] mock_get_splits.return_value = rows result = TrackSplitLevel().get_splits("uuid", []) mock_get_splits.assert_called_with("uuid", []) assert result == rows def test_to_track_splits_passes_through(): splits = [ SplitRow(1, SplitType.TRACK, "tuid-A", 0.5, 42, "NET"), SplitRow(2, SplitType.TRACK, "tuid-B", 0.3, 99, "NET"), ] result, _ = TrackSplitLevel().to_track_splits(splits, set(), count(-1, -1)) assert result == splits @patch("trigger.split_levels.rds.get_subaccount_splits_for_report_run") def test_get_subaccount_splits_delegates_to_rds(mock_get_splits): rows = [SplitRow(10, SplitType.SUBACCOUNT, "sub-1", 0.5, 42, "NET")] mock_get_splits.return_value = rows result = SubaccountSplitLevel().get_splits("uuid", []) mock_get_splits.assert_called_with("uuid", []) assert result == rows @patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") def test_to_track_splits_empty_input_returns_empty(mock_get_tuids): mock_get_tuids.return_value = {} result, _ = SubaccountSplitLevel().to_track_splits([], set(), count(-1, -1)) assert result == [] @patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") def test_to_track_splits_excludes_covered(mock_get_tuids): mock_get_tuids.return_value = {"sub-1": ["tuid-A", "tuid-B"]} result, _ = SubaccountSplitLevel().to_track_splits( [SplitRow(10, SplitType.SUBACCOUNT, "sub-1", 0.5, 42, "NET")], {("42", "tuid-A")}, count(-1, -1), ) assert result == [SplitRow(-1, SplitType.SUBACCOUNT, "tuid-B", 0.5, 42, "NET")] @patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") def test_to_track_splits_assigns_sequential_negative_ids(mock_get_tuids): mock_get_tuids.return_value = {"sub-1": ["tuid-A", "tuid-B"]} result, _ = SubaccountSplitLevel().to_track_splits( [SplitRow(10, SplitType.SUBACCOUNT, "sub-1", 0.5, 42, "NET")], set(), count(-1, -1) ) assert [s.split_id for s in result] == [-1, -2] @patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") def test_to_track_splits_ids_continue_from_supplied_counter(mock_get_tuids): mock_get_tuids.return_value = {"sub-1": ["tuid-A"], "sub-2": ["tuid-B"]} splits = [ SplitRow(10, SplitType.SUBACCOUNT, "sub-1", 0.5, 42, "NET"), SplitRow(11, SplitType.SUBACCOUNT, "sub-2", 0.5, 99, "NET"), ] shared_counter = count(-5, -1) result, _ = SubaccountSplitLevel().to_track_splits(splits, set(), shared_counter) assert [s.split_id for s in result] == [-5, -6] @patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") def test_to_track_splits_multiple_collaborators(mock_get_tuids): mock_get_tuids.return_value = {"sub-1": ["tuid-X"]} splits = [ SplitRow(10, SplitType.SUBACCOUNT, "sub-1", 0.5, 42, "NET"), SplitRow(11, SplitType.SUBACCOUNT, "sub-1", 0.25, 99, "NET"), ] result, _ = SubaccountSplitLevel().to_track_splits(splits, set(), count(-1, -1)) assert len(result) == 2 assert {s.collaborator_id for s in result} == {42, 99} @patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") def test_to_track_splits_deduplicates_within_output(mock_get_tuids): mock_get_tuids.return_value = {"sub-1": ["tuid-X"], "sub-2": ["tuid-X"]} splits = [ SplitRow(10, SplitType.SUBACCOUNT, "sub-1", 0.5, 42, "NET"), SplitRow(11, SplitType.SUBACCOUNT, "sub-2", 0.5, 42, "NET"), ] result, _ = SubaccountSplitLevel().to_track_splits(splits, set(), count(-1, -1)) assert len(result) == 1 assert result[0].identifier == "tuid-X"