"""Functional tests for the split level hierarchy logic using the in-memory SQLite DB. These tests differ from test_split_levels.py in that they use real DB data for the RDS layer and mock only the Snowflake read (get_track_tuids_by_subaccount) and upload operations. The key scenario exercised throughout: collaborator 10 has both a TRACK split for tuid-A and a SUBACCOUNT split for sub-1 (see tests/seed/04_split.sql). When sub-1 maps to ["tuid-A", "tuid-NEW"], the hierarchy must use the track split for tuid-A and only generate a subaccount-derived split for tuid-NEW. """ from itertools import count from unittest.mock import patch import pytest from trigger.constants import SplitType from trigger.schemas import CoveredTrackKey from trigger.split_levels import SubaccountSplitLevel, TrackSplitLevel from trigger.trigger import sync_splits @pytest.fixture() def all_splits(): """Run sync_splits end-to-end and return the split list passed to write_temporary_csv.""" with ( patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") as mock_tuids, patch("trigger.trigger.snowflake.create_splits_table"), patch("trigger.trigger.snowflake.upload_file_to_table_stage"), patch("trigger.trigger.snowflake.copy_into_table_from_stage"), patch("trigger.trigger.csv.write_temporary_csv") as mock_csv, patch("os.remove"), ): mock_tuids.return_value = {"sub-1": ["tuid-A", "tuid-NEW"]} mock_csv.return_value = "/tmp/test.csv" sync_splits("run-auto-1", []) return mock_csv.call_args[0][0] def _run_track_level(report_run_uuid="run-auto-1", invalid=()): level = TrackSplitLevel() splits = level.get_splits(report_run_uuid, list(invalid)) return level.to_track_splits(splits, set(), count(-1, -1)) def test_covers_all_collaborator_tuid_pairs(): _, covered = _run_track_level() assert CoveredTrackKey("10", "tuid-A") in covered assert CoveredTrackKey("11", "tuid-B") in covered assert CoveredTrackKey("12", "tuid-C") in covered def test_does_not_cover_subaccount_identifiers(): _, covered = _run_track_level() assert CoveredTrackKey("10", "sub-1") not in covered def test_excludes_invalid_collaborators(): track_splits, covered = _run_track_level(invalid=[10]) assert all(r.collaborator_id != 10 for r in track_splits) assert CoveredTrackKey("10", "tuid-A") not in covered def test_tuid_already_covered_by_track_is_excluded(): _, covered = _run_track_level() sub_level = SubaccountSplitLevel() sub_splits = sub_level.get_splits("run-auto-1", []) with patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") as mock: mock.return_value = {"sub-1": ["tuid-A", "tuid-NEW"]} resolved, _ = sub_level.to_track_splits(sub_splits, covered, count(-1, -1)) collab_10_tuids = {r.identifier for r in resolved if r.collaborator_id == 10} assert "tuid-A" not in collab_10_tuids assert "tuid-NEW" in collab_10_tuids def test_tuid_not_covered_by_track_is_included(): _, covered = _run_track_level() sub_level = SubaccountSplitLevel() sub_splits = sub_level.get_splits("run-auto-1", []) with patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") as mock: mock.return_value = {"sub-1": ["tuid-ONLY-SUB"]} resolved, _ = sub_level.to_track_splits(sub_splits, covered, count(-1, -1)) assert any(r.identifier == "tuid-ONLY-SUB" and r.collaborator_id == 10 for r in resolved) def test_resolved_subaccount_rows_carry_subaccount_split_type(): _, covered = _run_track_level() sub_level = SubaccountSplitLevel() sub_splits = sub_level.get_splits("run-auto-1", []) with patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") as mock: mock.return_value = {"sub-1": ["tuid-X"]} resolved, _ = sub_level.to_track_splits(sub_splits, covered, count(-1, -1)) assert resolved assert all(r.split_type_id == SplitType.SUBACCOUNT for r in resolved) def test_invalid_collaborators_excluded_from_subaccount_splits(): sub_splits = SubaccountSplitLevel().get_splits("run-auto-1", [10]) assert all(r.collaborator_id != 10 for r in sub_splits) def test_track_split_wins_over_subaccount_for_same_tuid(all_splits): tuid_a_for_10 = [r for r in all_splits if r.identifier == "tuid-A" and r.collaborator_id == 10] assert len(tuid_a_for_10) == 1 assert tuid_a_for_10[0].split_type_id == SplitType.TRACK def test_subaccount_expansion_fills_uncovered_tuids(all_splits): tuid_new_for_10 = [ r for r in all_splits if r.identifier == "tuid-NEW" and r.collaborator_id == 10 ] assert len(tuid_new_for_10) == 1 assert tuid_new_for_10[0].split_type_id == SplitType.SUBACCOUNT def test_no_duplicate_collaborator_tuid_pairs(all_splits): keys = [(r.collaborator_id, r.identifier) for r in all_splits] assert len(keys) == len(set(keys)) def test_invalid_collaborators_absent_from_all_splits(): with ( patch("trigger.split_levels.snowflake.get_track_tuids_by_subaccount") as mock_tuids, patch("trigger.trigger.snowflake.create_splits_table"), patch("trigger.trigger.snowflake.upload_file_to_table_stage"), patch("trigger.trigger.snowflake.copy_into_table_from_stage"), patch("trigger.trigger.csv.write_temporary_csv") as mock_csv, patch("os.remove"), ): mock_tuids.return_value = {"sub-1": ["tuid-A", "tuid-NEW"]} mock_csv.return_value = "/tmp/test.csv" sync_splits("run-auto-1", [10]) result = mock_csv.call_args[0][0] assert all(r.collaborator_id != 10 for r in result)