"""Tests for snowflake SplitPersister.""" from unittest.mock import MagicMock import pytest from collaborator.models.snowflake.split_persister import ( _SNOWFLAKE_IN_CLAUSE_LIMIT, SplitPersister, ) def _make_mock_session(rows=None): session = MagicMock() session.execute.return_value.all.return_value = rows or [] return session @pytest.mark.parametrize( "tuids, expected", [ ([], []), (["abc", "xyz"], []), # non-digit tuids are skipped ], ) def test_get_splits_for_template_empty(tuids, expected): """Returns empty list for empty input or all-non-digit tuids.""" session = _make_mock_session() result = SplitPersister.get_splits_for_template(tuids, session=session) assert result == expected session.execute.assert_not_called() def test_get_splits_for_template_returns_rows(): """Returns rows from session.execute for a small TUID list.""" fake_rows = [("111", 1, 0.5, "NET"), ("222", 2, 0.25, "GROSS")] session = _make_mock_session(rows=fake_rows) result = SplitPersister.get_splits_for_template(["111", "222"], session=session) assert result == fake_rows assert session.execute.call_count == 1 def test_get_splits_for_template_chunks_large_tuid_list(): """Splits a TUID list larger than _SNOWFLAKE_IN_CLAUSE_LIMIT into multiple queries.""" total = _SNOWFLAKE_IN_CLAUSE_LIMIT + 100 tuids = [str(i) for i in range(total)] chunk1_rows = [("1", 10, 0.5, "NET")] chunk2_rows = [("2", 20, 0.3, "GROSS")] session = MagicMock() first_call = MagicMock() first_call.all.return_value = chunk1_rows second_call = MagicMock() second_call.all.return_value = chunk2_rows session.execute.side_effect = [first_call, second_call] result = SplitPersister.get_splits_for_template(tuids, session=session) assert session.execute.call_count == 2 assert result == chunk1_rows + chunk2_rows def test_get_splits_for_template_exactly_at_limit(): """A list of exactly _SNOWFLAKE_IN_CLAUSE_LIMIT TUIDs uses a single query.""" tuids = [str(i) for i in range(_SNOWFLAKE_IN_CLAUSE_LIMIT)] session = _make_mock_session(rows=[]) SplitPersister.get_splits_for_template(tuids, session=session) assert session.execute.call_count == 1