"""Tests for split logic used by the PUT /splits endpoint.""" from unittest.mock import ANY import pytest from collaborator.constants import error from collaborator.constants.split import RateType from collaborator.constants.split_type import SplitTypeId from collaborator.logic import split as split_logic from collaborator.schemas.split import ( ReplacementSplitSchema, ReplacementSplitsSchema, ReplaceSplitsRequestSchema, ) from collaborator.utils.error import OwsError from collaborator.utils.typing import AuthorizedResources, User 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 to_dict.""" def __init__(self, data): self._data = data def to_dict(self): return dict(self._data) # --------------------------------------------------------------------------- # replace_splits # --------------------------------------------------------------------------- def test_replace_splits_calls_persister_and_returns_final_splits(mocker): """Authorises, persists, spools to Snowflake and returns the final splits.""" body = ReplaceSplitsRequestSchema( vendor_id=24601, dp_splits_agreed=True, replacements=[ _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) ], ) user = User(type="oa", id=99) authorized_resources = AuthorizedResources() check_mock = mocker.patch.object( split_logic, "_check_replace_can_modify_split_data", return_value=24601, ) has_dp_mock = mocker.patch.object( split_logic.ows_account, "has_direct_payments", return_value=True ) final_split = _FakeSplit({"id": 7, "identifier": "tuid-1"}) persister_mock = mocker.patch.object( split_logic.SplitPersister, "replace_splits", return_value=([final_split], [4, 5]), ) snowflake_mock = mocker.patch.object( split_logic.SnowflakeSplitPersister, "replace_splits" ) mocker.patch.object( split_logic, "uwsgi_spool_task", side_effect=lambda func, *a, **k: func(*a, **k) ) result = split_logic.replace_splits(authorized_resources, body, user) assert result == [final_split.to_dict()] check_mock.assert_called_once_with( authorized_resources, body.replacements, body.vendor_id ) has_dp_mock.assert_called_once_with("24601") persister_mock.assert_called_once_with( replacements=body.replacements, dp_splits_agreed=True, vendor_id=24601, has_direct_payments=True, user=user, ) snowflake_mock.assert_called_once_with([final_split.to_dict()], [4, 5]) def test_replace_splits_defaults_dp_splits_agreed_to_false(mocker): """A missing dp_splits_agreed is passed to the persister as False.""" body = ReplaceSplitsRequestSchema( vendor_id=24601, replacements=[ _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) ], ) user = User(type="oa", id=99) mocker.patch.object( split_logic, "_check_replace_can_modify_split_data", return_value=24601 ) mocker.patch.object( split_logic.ows_account, "has_direct_payments", return_value=False ) persister_mock = mocker.patch.object( split_logic.SplitPersister, "replace_splits", return_value=([], []) ) mocker.patch.object(split_logic.SnowflakeSplitPersister, "replace_splits") mocker.patch.object( split_logic, "uwsgi_spool_task", side_effect=lambda func, *a, **k: func(*a, **k) ) split_logic.replace_splits(AuthorizedResources(), body, user) assert persister_mock.call_args.kwargs["dp_splits_agreed"] is False # --------------------------------------------------------------------------- # _check_replace_can_modify_split_data # --------------------------------------------------------------------------- def test_check_replace_authorises_single_vendor(mocker): """A single derived vendor is authorised and returned.""" mocker.patch.object( split_logic, "_verify_track_splits_belong_to_vendor", return_value={24601} ) mocker.patch.object( split_logic, "_verify_subaccount_splits_belong_to_vendor", return_value=set() ) auth_mock = mocker.patch.object(split_logic, "check_vendors_authorization") split_logic._check_replace_can_modify_split_data(AuthorizedResources(), [], 24601) auth_mock.assert_called_once_with(ANY, [24601]) def test_check_replace_combines_track_and_subaccount_vendors(mocker): """Track and subaccount vendor sets are combined before the single check.""" mocker.patch.object( split_logic, "_verify_track_splits_belong_to_vendor", return_value={24601} ) mocker.patch.object( split_logic, "_verify_subaccount_splits_belong_to_vendor", return_value={24601} ) auth_mock = mocker.patch.object(split_logic, "check_vendors_authorization") split_logic._check_replace_can_modify_split_data(AuthorizedResources(), [], 24601) auth_mock.assert_called_once_with(ANY, [24601]) # --------------------------------------------------------------------------- # _verify_track_splits_belong_to_vendor # --------------------------------------------------------------------------- def test_verify_track_splits_returns_empty_when_no_track_replacements(mocker): """Subaccount-only input yields an empty vendor set and no lookups.""" track_mock = mocker.patch.object(split_logic.ows_track, "get_tracks_batched") subaccount = _replacement( "sub-1", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) assert split_logic._verify_track_splits_belong_to_vendor([subaccount], {24601}) == { 24601 } track_mock.assert_not_called() def test_verify_track_splits_single_vendor(mocker): """Track and collaborator on the same vendor returns that vendor.""" track = _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) mocker.patch.object( split_logic.ows_track, "get_tracks_batched", return_value=[{"identifier": "tuid-1", "upc": "UPC-1"}], ) mocker.patch.object( split_logic.ows_product, "get_products_by_upc", return_value={"items": [{"upc": "UPC-1", "vendor_id": 24601}]}, ) mocker.patch.object( split_logic.CollaboratorPersister, "get_by_ids", return_value=[{"vendor_id": 24601}], ) assert split_logic._verify_track_splits_belong_to_vendor([track], {24601}) == { 24601 } def test_verify_track_splits_raises_when_track_and_collaborator_differ(mocker): """Track vendor and collaborator vendor differing is rejected.""" track = _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) mocker.patch.object( split_logic.ows_track, "get_tracks_batched", return_value=[{"identifier": "tuid-1", "upc": "UPC-1"}], ) mocker.patch.object( split_logic.ows_product, "get_products_by_upc", return_value={"items": [{"upc": "UPC-1", "vendor_id": 24601}]}, ) mocker.patch.object( split_logic.CollaboratorPersister, "get_by_ids", return_value=[{"vendor_id": 90210}], ) with pytest.raises(OwsError) as err: split_logic._verify_track_splits_belong_to_vendor([track], {24601}) assert err.value.code == error.ERROR_CODE_SPLIT_VENDOR_MISMATCH # --------------------------------------------------------------------------- # _verify_subaccount_splits_belong_to_vendor # --------------------------------------------------------------------------- def test_verify_subaccount_splits_returns_empty_when_no_subaccount_replacements(mocker): """Track-only input yields an empty vendor set and no subaccount lookups.""" subaccount_mock = mocker.patch.object(split_logic.ows_account, "get_subaccount") track = _replacement( "tuid-1", SplitTypeId.TRACK, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ) assert split_logic._verify_subaccount_splits_belong_to_vendor([track], {24601}) == { 24601 } subaccount_mock.assert_not_called() def test_verify_subaccount_splits_single_vendor(mocker): """Subaccounts resolving to one vendor returns that vendor.""" subaccounts = [ _replacement( "sub-1", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ), _replacement( "sub-2", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 2, "split_rate": 0.5, "rate_type": RateType.NET}], ), ] mocker.patch.object( split_logic.ows_account, "get_subaccount", return_value={"vendor_id": 24601}, ) assert split_logic._verify_subaccount_splits_belong_to_vendor( subaccounts, {24601} ) == {24601} def test_verify_subaccount_splits_raises_when_vendors_differ(mocker): """Subaccounts on different vendors are rejected.""" subaccounts = [ _replacement( "sub-1", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 1, "split_rate": 0.5, "rate_type": RateType.NET}], ), _replacement( "sub-2", SplitTypeId.SUBACCOUNT, [{"collaborator_id": 2, "split_rate": 0.5, "rate_type": RateType.NET}], ), ] mocker.patch.object( split_logic.ows_account, "get_subaccount", side_effect=[{"vendor_id": 24601}, {"vendor_id": 90210}], ) with pytest.raises(OwsError) as err: split_logic._verify_subaccount_splits_belong_to_vendor(subaccounts, {24601}) assert err.value.code == error.ERROR_CODE_SPLIT_VENDOR_MISMATCH