"""Tests for split utils.""" import pytest from collaborator.constants import error from collaborator.constants.split import RateType from collaborator.constants.split_type import SplitTypeId from collaborator.models.rds.collaborator_persister import CollaboratorPersister from collaborator.schemas.split import ReplacementSplitSchema, ReplacementSplitsSchema from collaborator.utils import logging from collaborator.utils import split as split_util from collaborator.utils.error import OwsError def _replacement(identifier, split_type_id, splits): """Build a ReplacementSplitsSchema from plain dicts.""" return ReplacementSplitsSchema( identifier=identifier, split_type_id=split_type_id, splits=[ReplacementSplitSchema(**split) for split in splits], ) class _FakeSplit: """Minimal stand-in for an ORM Split, exposing split_id and to_dict.""" def __init__(self, split_id, data): self.split_id = split_id self._data = data def to_dict(self): return dict(self._data) # --------------------------------------------------------------------------- # filter_replacement_splits_by_type # --------------------------------------------------------------------------- def test_filter_replacement_splits_by_type_returns_only_matching_type(): """Only replacements whose split_type_id matches are returned.""" track = _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) subaccount = _replacement( "sub-1", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 2, "split_rate": 0.5, "rate_type": RateType.NET}], ) assert split_util.filter_replacement_splits_by_type( [track, subaccount], SplitTypeId.TRACK ) == [track] assert split_util.filter_replacement_splits_by_type( [track, subaccount], SplitTypeId.SUBACCOUNT ) == [subaccount] def test_filter_replacement_splits_by_type_empty_when_none_match(): """Returns an empty list when no replacement matches the type.""" track = _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) assert ( split_util.filter_replacement_splits_by_type([track], SplitTypeId.SUBACCOUNT) == [] ) # --------------------------------------------------------------------------- # validate_split_rates # --------------------------------------------------------------------------- def test_validate_split_rates_passes_for_valid_rates(): """Valid track and subaccount rates do not raise.""" replacements = [ _replacement( "tuid-1", SplitTypeId.TRACK, [ {"collaborator_id": 1, "split_rate": 0.6, "rate_type": RateType.NET}, {"collaborator_id": 2, "split_rate": 0.4, "rate_type": RateType.NET}, ], ), _replacement( "sub-1", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 3, "split_rate": 1.0, "rate_type": RateType.NET}], ), ] # Should not raise. split_util.validate_split_rates(replacements) def test_validate_split_rates_raises_when_subaccount_over_100(): """A single subaccount split over 100% is rejected.""" replacements = [ _replacement( "sub-1", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 1, "split_rate": 1.5, "rate_type": RateType.NET}], ), ] with pytest.raises(OwsError) as err: split_util.validate_split_rates(replacements) assert err.value.code == error.ERROR_CODE_BAD_PARAMS def test_validate_split_rates_raises_when_track_total_over_100(): """Track splits for a single tuid summing over 100% are rejected.""" replacements = [ _replacement( "tuid-1", SplitTypeId.TRACK, [ {"collaborator_id": 1, "split_rate": 0.6, "rate_type": RateType.NET}, {"collaborator_id": 2, "split_rate": 0.6, "rate_type": RateType.NET}, ], ), ] with pytest.raises(OwsError) as err: split_util.validate_split_rates(replacements) assert err.value.code == error.ERROR_CODE_BAD_PARAMS def test_validate_split_rates_track_total_is_per_tuid(): """Totals are computed per tuid, so two tuids each <=100% are allowed.""" replacements = [ _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.6, "rate_type": RateType.NET}], ), _replacement( "tuid-2", SplitTypeId.TRACK, [{"collaborator_id": 2, "split_rate": 0.6, "rate_type": RateType.NET}], ), ] # Combined they exceed 100% but per-tuid they don't, so no error. split_util.validate_split_rates(replacements) def test_validate_split_rates_subaccount_exactly_100_allowed(): """A subaccount split of exactly 100% is allowed (only > 100% fails).""" replacements = [ _replacement( "sub-1", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 1, "split_rate": 1.0, "rate_type": RateType.NET}], ), ] split_util.validate_split_rates(replacements) # --------------------------------------------------------------------------- # get_collabs_needing_tcs_agreement # --------------------------------------------------------------------------- def test_get_collabs_needing_tcs_agreement_returns_only_dp_enabled(mocker): """Only collaborators with a dp_enabled_date are returned.""" mocker.patch.object( CollaboratorPersister, "get_by_ids", return_value=[ {"id": 1, "dp_enabled_date": None}, {"id": 2, "dp_enabled_date": "2020-01-01T00:00:00"}, {"id": 3, "dp_enabled_date": None}, {"id": 4, "dp_enabled_date": "2021-06-06T00:00:00"}, ], ) result = split_util.get_collabs_needing_tcs_agreement([1, 2, 3, 4]) assert sorted(result) == [2, 4] def test_get_collabs_needing_tcs_agreement_empty_when_none_dp_enabled(mocker): """Returns an empty list when no collaborator has DP enabled.""" mocker.patch.object( CollaboratorPersister, "get_by_ids", return_value=[ {"id": 1, "dp_enabled_date": None}, {"id": 2, "dp_enabled_date": None}, ], ) assert split_util.get_collabs_needing_tcs_agreement([1, 2]) == [] # --------------------------------------------------------------------------- # log_upsert_events # --------------------------------------------------------------------------- def test_log_upsert_events_logs_an_update_event_per_split(mocker, mock_user): """Logs a single bulk update event covering each provided split.""" bulk_log_events_mock = mocker.patch.object(logging, "bulk_log_events") splits = [ _FakeSplit(1, {"id": 1, "split_rate": 0.5}), _FakeSplit(2, {"id": 2, "split_rate": 0.5}), ] split_util.log_upsert_events(splits, mock_user) bulk_log_events_mock.assert_called_once() event, entity, event_data, user = bulk_log_events_mock.call_args[0] assert event == logging.LOG_EVENT_UPDATE assert entity == "split" assert len(event_data) == 2 assert {item["id"] for item in event_data} == {1, 2} assert user == mock_user def test_log_upsert_events_with_no_splits(mocker, mock_user): """An empty split list still logs an (empty) bulk event.""" bulk_log_events_mock = mocker.patch.object(logging, "bulk_log_events") split_util.log_upsert_events([], mock_user) bulk_log_events_mock.assert_called_once() assert bulk_log_events_mock.call_args[0][2] == []