"""Unit tests for bulk_ingest_splits.""" from contextlib import ExitStack from unittest.mock import patch from pydantic import ValidationError import pytest from collaborator.logic.bulk_ingest_splits import ( SplitIngestionRow, compute_changes, resolve_product_ids, run_bulk_ingest, validate_collaborators, validate_single_vendor, validate_tracks, ) from collaborator.schemas.split import BulkSplitRow MODULE = "collaborator.logic.bulk_ingest_splits" def _row( tuid, collaborator_id, split_rate=0.1, collaborator_name=None, product_id=None ): return SplitIngestionRow( vendor_id=24601, collaborator_id=collaborator_id, collaborator_name=collaborator_name or f"Collab {collaborator_id}", split_rate=split_rate, split_type="NET", tuid=tuid, product_id=product_id, ) def _run_compute(rows, existing_by_tuid=None, tuid_to_product_id=None): """Run compute_changes with mocked DB calls.""" with ( patch( "collaborator.models.rds.split_persister.SplitPersister.get_track_splits_with_rates", return_value=existing_by_tuid or {}, ), patch( "collaborator.models.snowflake.product_persister.ProductPersister.get_product_ids_for_tuids", return_value=tuid_to_product_id or {}, ), ): return compute_changes( rows=rows, vendor_id=24601, input_rows=len(rows), vendor_currency="USD", ) # --------------------------------------------------------------------------- # Core counts # --------------------------------------------------------------------------- def test_new_splits_only(): """All rows have no existing splits → all new.""" rows = [_row("111", 1, 0.2), _row("222", 2, 0.3)] summary = _run_compute(rows) assert summary.new_splits_created == 2 assert summary.existing_splits_updated == 0 assert summary.splits_ingested == 2 assert summary.tracks_impacted == 2 assert summary.overwritten_splits == [] assert summary.collaborators_receiving_new_splits == 2 def test_all_overwritten(): """All rows overwrite existing splits.""" rows = [_row("111", 1, 0.5)] existing = {"111": {1: 0.2}} tuid_to_product = {"111": 10} summary = _run_compute( rows, existing_by_tuid=existing, tuid_to_product_id=tuid_to_product ) assert summary.existing_splits_updated == 1 assert summary.new_splits_created == 0 assert summary.splits_ingested == 1 assert summary.collaborators_receiving_new_splits == 0 assert len(summary.overwritten_splits) == 1 ow = summary.overwritten_splits[0] assert ow.tuid == "111" assert ow.product_id == 10 assert ow.old_split_rate == pytest.approx(0.2) assert ow.new_split_rate == pytest.approx(0.5) assert ow.collaborator_name == "Collab 1" def test_mixed_new_and_overwritten(): """Some rows are new, some overwrite existing splits.""" rows = [ _row("111", 1, 0.2), # new _row("111", 2, 0.3), # overwrite ] existing = {"111": {2: 0.1}} tuid_to_product = {"111": 5} summary = _run_compute( rows, existing_by_tuid=existing, tuid_to_product_id=tuid_to_product ) assert summary.new_splits_created == 1 assert summary.existing_splits_updated == 1 assert summary.splits_ingested == 2 assert ( summary.collaborators_receiving_new_splits == 1 ) # only collab 1 got a new split assert len(summary.overwritten_splits) == 1 # --------------------------------------------------------------------------- # products_impacted # --------------------------------------------------------------------------- def test_products_impacted_counts_unique_products(): """Two tracks from the same product → products_impacted == 1.""" rows = [_row("111", 1, 0.1), _row("222", 1, 0.1)] tuid_to_product = {"111": 10, "222": 10} summary = _run_compute(rows, tuid_to_product_id=tuid_to_product) assert summary.products_impacted == 1 def test_products_impacted_multiple_products(): """Tracks from two different products → products_impacted == 2.""" rows = [_row("111", 1, 0.1), _row("222", 1, 0.1)] tuid_to_product = {"111": 10, "222": 20} summary = _run_compute(rows, tuid_to_product_id=tuid_to_product) assert summary.products_impacted == 2 # --------------------------------------------------------------------------- # tracks_impacted # --------------------------------------------------------------------------- def test_tracks_impacted_excludes_unaffected(): """A track where only unaffected splits exist (no change) is not impacted.""" rows = [_row("111", 1, 0.1)] # new split for collab 1 existing = {"111": {2: 0.5}} # collab 2 is unaffected summary = _run_compute(rows, existing_by_tuid=existing) assert summary.tracks_impacted == 1 assert summary.tracks_with_unaffected_splits == 1 # --------------------------------------------------------------------------- # OverwrittenSplit details # --------------------------------------------------------------------------- def test_overwritten_split_missing_product_id(): """When tuid has no product_id mapping, product_id is None.""" rows = [_row("999", 3, 0.4)] existing = {"999": {3: 0.25}} summary = _run_compute(rows, existing_by_tuid=existing, tuid_to_product_id={}) assert summary.overwritten_splits[0].product_id is None def test_empty_rows_returns_zeroes(): """Empty row list produces a zeroed-out summary.""" summary = _run_compute([]) assert summary.splits_ingested == 0 assert summary.tracks_impacted == 0 assert summary.products_impacted == 0 assert summary.overwritten_splits == [] def test_new_collab_by_name_counted_as_new_split(): """A row with collaborator_name but no collaborator_id counts as a new split.""" rows = [ SplitIngestionRow( vendor_id=24601, collaborator_id=None, collaborator_name="New Artist", split_rate=0.1, split_type="NET", tuid="111", ) ] summary = _run_compute(rows) assert summary.new_splits_created == 1 assert summary.existing_splits_updated == 0 assert summary.new_collaborators == 1 # --------------------------------------------------------------------------- # SplitIngestionRow validator # --------------------------------------------------------------------------- def test_split_ingestion_row_requires_tuid_or_product_id(): """Constructing a row with neither tuid nor product_id raises ValidationError.""" with pytest.raises(ValidationError, match="tuid or a Product ID"): SplitIngestionRow( vendor_id=1, collaborator_id=1, split_rate=0.1, split_type="NET", tuid=None, product_id=None, ) def test_split_ingestion_row_accepts_tuid_only(): """A row with only tuid is valid.""" row = SplitIngestionRow( vendor_id=1, collaborator_id=1, split_rate=0.1, split_type="NET", tuid="abc" ) assert row.tuid == "abc" assert row.product_id is None def test_split_ingestion_row_accepts_product_id_only(): """A row with only product_id is valid.""" row = SplitIngestionRow( vendor_id=1, collaborator_id=1, split_rate=0.1, split_type="NET", product_id=5 ) assert row.product_id == 5 assert row.tuid is None # --------------------------------------------------------------------------- # validate_single_vendor # --------------------------------------------------------------------------- def test_validate_single_vendor_returns_vendor_id(): """Returns the vendor_id when all rows share the same vendor.""" rows = [_row("111", 1), _row("222", 2)] assert validate_single_vendor(rows) == 24601 def test_validate_single_vendor_raises_on_multiple_vendors(): """Raises when rows contain more than one distinct vendor_id.""" rows = [ _row("111", 1), SplitIngestionRow( vendor_id=99999, collaborator_id=2, split_rate=0.1, split_type="NET", tuid="222", ), ] with pytest.raises(ValueError, match="exactly one vendor ID"): validate_single_vendor(rows) # --------------------------------------------------------------------------- # validate_collaborators # --------------------------------------------------------------------------- MODULE_RDS_COLLAB = ( "collaborator.models.rds.collaborator_persister.CollaboratorPersister" ) def test_validate_collaborators_passes_when_all_valid(): """No exception raised when all collaborator IDs belong to the vendor.""" rows = [_row("111", 1), _row("222", 2)] with patch( f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids", return_value={1: 24601, 2: 24601} ): validate_collaborators(rows, vendor_id=24601) # no exception def test_validate_collaborators_raises_on_missing_collaborator(): """Raises when a collaborator ID is not found in the DB.""" rows = [_row("111", 99)] with patch(f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids", return_value={}): with pytest.raises(ValueError, match="Collaborator ID 99 not found"): validate_collaborators(rows, vendor_id=24601) def test_validate_collaborators_raises_on_wrong_vendor(): """Raises when a collaborator belongs to a different vendor.""" rows = [_row("111", 5)] with patch(f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids", return_value={5: 11111}): with pytest.raises(ValueError, match="belongs to vendor 11111"): validate_collaborators(rows, vendor_id=24601) def test_validate_collaborators_skips_rows_without_collaborator_id(): """Rows with collaborator_id=None (new collaborators) are not validated.""" rows = [ SplitIngestionRow( vendor_id=24601, collaborator_name="New Artist", split_rate=0.1, split_type="NET", tuid="111", ) ] with patch(f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids") as mock_get: validate_collaborators(rows, vendor_id=24601) mock_get.assert_not_called() # --------------------------------------------------------------------------- # validate_tracks # --------------------------------------------------------------------------- MODULE_SF_PRODUCT = "collaborator.models.snowflake.product_persister.ProductPersister" def test_validate_tracks_passes_when_all_valid(): """No exception raised when all TUIDs exist and belong to the vendor.""" rows = [_row("111", 1), _row("222", 2)] with patch( f"{MODULE_SF_PRODUCT}.get_vendor_map_by_tuids", return_value={"111": 24601, "222": 24601}, ): validate_tracks(rows, vendor_id=24601) # no exception def test_validate_tracks_raises_on_missing_tuid(): """Raises when a TUID is not found in Snowflake.""" rows = [_row("missing-tuid", 1)] with patch(f"{MODULE_SF_PRODUCT}.get_vendor_map_by_tuids", return_value={}): with pytest.raises(ValueError, match="TUIDs not found in Snowflake"): validate_tracks(rows, vendor_id=24601) def test_validate_tracks_raises_on_wrong_vendor(): """Raises when a TUID belongs to a different vendor.""" rows = [_row("111", 1)] with patch( f"{MODULE_SF_PRODUCT}.get_vendor_map_by_tuids", return_value={"111": 99999} ): with pytest.raises(ValueError, match="do not belong to vendor"): validate_tracks(rows, vendor_id=24601) # --------------------------------------------------------------------------- # resolve_product_ids # --------------------------------------------------------------------------- def _product_row(product_id, collaborator_id=1, split_rate=0.1): return SplitIngestionRow( vendor_id=24601, collaborator_id=collaborator_id, collaborator_name=f"Collab {collaborator_id}", split_rate=split_rate, split_type="NET", product_id=product_id, ) def test_resolve_product_ids_expands_product_to_tuids(): """A product-only row is expanded into one row per track under that product.""" rows = [_product_row(10)] with ( patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value={10}), patch( f"{MODULE_SF_PRODUCT}.get_tuid_product_pairs", return_value=[("abc", 10), ("def", 10)], ), ): expanded, skipped = resolve_product_ids(rows, vendor_id=24601) assert skipped == 0 tuids = {r.tuid for r in expanded} assert tuids == {"abc", "def"} assert all(r.product_id == 10 for r in expanded) def test_resolve_product_ids_raises_on_invalid_product(): """Raises when a product ID does not belong to the vendor.""" rows = [_product_row(99)] with patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value=set()): with pytest.raises(ValueError, match="Product IDs not found for vendor"): resolve_product_ids(rows, vendor_id=24601) def test_resolve_product_ids_skips_product_with_no_tracks(): """Products with no eligible tracks are skipped and counted.""" rows = [_product_row(10)] with ( patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value={10}), patch(f"{MODULE_SF_PRODUCT}.get_tuid_product_pairs", return_value=[]), ): expanded, skipped = resolve_product_ids(rows, vendor_id=24601) assert skipped == 1 assert expanded == [] def test_resolve_product_ids_preserves_tuid_only_rows(): """Rows that already have a tuid pass through unchanged.""" tuid_row = _row("existing-tuid", 1) product_row = _product_row(10) with ( patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value={10}), patch( f"{MODULE_SF_PRODUCT}.get_tuid_product_pairs", return_value=[("new-tuid", 10)], ), ): expanded, skipped = resolve_product_ids( [tuid_row, product_row], vendor_id=24601 ) assert skipped == 0 tuids = {r.tuid for r in expanded} assert "existing-tuid" in tuids assert "new-tuid" in tuids def test_resolve_product_ids_noop_when_no_product_rows(): """If all rows already have tuids, the DB is never called.""" rows = [_row("111", 1), _row("222", 2)] with patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor") as mock_get: expanded, skipped = resolve_product_ids(rows, vendor_id=24601) mock_get.assert_not_called() assert skipped == 0 assert expanded == rows # --------------------------------------------------------------------------- # run_bulk_ingest # --------------------------------------------------------------------------- MODULE_RDS_SPLIT = "collaborator.models.rds.split_persister.SplitPersister" MODULE_PAYMENT = ( "collaborator.models.snowflake.account_payment_term_persister" ".AccountPaymentTermPersister" ) MODULE_RDS_COLLAB_CREATE = ( "collaborator.models.rds.collaborator_persister.CollaboratorPersister" ) def _bulk_row(tuid, collaborator_id=1, vendor_id=24601, collaborator_name=None): return BulkSplitRow( vendor_id=vendor_id, collaborator_id=collaborator_id, collaborator_name=collaborator_name or f"Collab {collaborator_id}", split_rate=0.2, rate_type="NET", tuid=tuid, ) def _patch_pipeline( collab_vendor_map=None, tuid_vendor_map=None, existing_by_tuid=None, tuid_to_product_id=None, currency="USD", ): """Context manager stack that stubs all DB calls in run_bulk_ingest.""" return ( patch( f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids", return_value=collab_vendor_map or {1: 24601}, ), patch( f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value=set(), ), patch( f"{MODULE_SF_PRODUCT}.get_vendor_map_by_tuids", return_value=tuid_vendor_map or {"tuid-a": 24601}, ), patch(f"{MODULE_PAYMENT}.get_currency_for_vendor", return_value=currency), patch( f"{MODULE_RDS_SPLIT}.get_track_splits_with_rates", return_value=existing_by_tuid or {}, ), patch( f"{MODULE_SF_PRODUCT}.get_product_ids_for_tuids", return_value=tuid_to_product_id or {}, ), ) def test_run_bulk_ingest_dry_run_returns_summary(): """dry_run=True returns a summary without calling upsert.""" rows = [_bulk_row("tuid-a")] with ExitStack() as stack: for p in _patch_pipeline(): stack.enter_context(p) mock_upsert = stack.enter_context( patch(f"{MODULE_RDS_SPLIT}.upsert_track_splits") ) summary = run_bulk_ingest(rows, dry_run=True) assert summary.dry_run is True assert summary.splits_ingested == 1 mock_upsert.assert_not_called() def test_run_bulk_ingest_dry_run_does_not_write(): """dry_run=True skips upsert even when there are rows to write.""" rows = [_bulk_row("tuid-a"), _bulk_row("tuid-b")] with ExitStack() as stack: for p in _patch_pipeline(tuid_vendor_map={"tuid-a": 24601, "tuid-b": 24601}): stack.enter_context(p) mock_upsert = stack.enter_context( patch(f"{MODULE_RDS_SPLIT}.upsert_track_splits") ) run_bulk_ingest(rows, dry_run=True) mock_upsert.assert_not_called() def test_run_bulk_ingest_write_calls_upsert(): """dry_run=False calls upsert_track_splits with the correct rows.""" rows = [_bulk_row("tuid-a")] with ExitStack() as stack: for p in _patch_pipeline(): stack.enter_context(p) mock_upsert = stack.enter_context( patch(f"{MODULE_RDS_SPLIT}.upsert_track_splits") ) summary = run_bulk_ingest(rows, dry_run=False, ticket_id="TICKET-1") assert summary.dry_run is False mock_upsert.assert_called_once() split_data, ticket = mock_upsert.call_args.args assert len(split_data) == 1 assert split_data[0]["tuid"] == "tuid-a" assert split_data[0]["collaborator_id"] == 1 assert ticket == "TICKET-1" def test_run_bulk_ingest_auto_generates_ticket_id(): """When dry_run=False and no ticket_id is given, one is auto-generated.""" rows = [_bulk_row("tuid-a")] with ExitStack() as stack: for p in _patch_pipeline(): stack.enter_context(p) mock_upsert = stack.enter_context( patch(f"{MODULE_RDS_SPLIT}.upsert_track_splits") ) run_bulk_ingest(rows, dry_run=False, ticket_id=None) _, auto_ticket = mock_upsert.call_args.args assert auto_ticket.startswith("From UI") assert "vendor 24601" in auto_ticket def test_run_bulk_ingest_creates_new_collaborators_when_writing(): """When dry_run=False and a row has no collaborator_id, get_or_create_by_names is called.""" new_collab_row = BulkSplitRow( vendor_id=24601, collaborator_id=None, collaborator_name="Brand New Artist", split_rate=0.1, rate_type="NET", tuid="tuid-a", ) with ( patch(f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids", return_value={}), patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value=set()), patch( f"{MODULE_SF_PRODUCT}.get_vendor_map_by_tuids", return_value={"tuid-a": 24601}, ), patch(f"{MODULE_PAYMENT}.get_currency_for_vendor", return_value="USD"), patch( f"{MODULE_RDS_COLLAB_CREATE}.get_or_create_by_names", return_value=({"Brand New Artist": 999}, 1), ) as mock_create, patch(f"{MODULE_RDS_SPLIT}.get_track_splits_with_rates", return_value={}), patch(f"{MODULE_SF_PRODUCT}.get_product_ids_for_tuids", return_value={}), patch(f"{MODULE_RDS_SPLIT}.upsert_track_splits"), ): summary = run_bulk_ingest([new_collab_row], dry_run=False, ticket_id="T-1") mock_create.assert_called_once_with({"Brand New Artist"}, 24601, "T-1", "USD") assert summary.new_collaborators == 1 def test_run_bulk_ingest_skips_new_collaborator_creation_on_dry_run(): """dry_run=True must never call get_or_create_by_names.""" new_collab_row = BulkSplitRow( vendor_id=24601, collaborator_id=None, collaborator_name="Brand New Artist", split_rate=0.1, rate_type="NET", tuid="tuid-a", ) with ( patch(f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids", return_value={}), patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value=set()), patch( f"{MODULE_SF_PRODUCT}.get_vendor_map_by_tuids", return_value={"tuid-a": 24601}, ), patch(f"{MODULE_PAYMENT}.get_currency_for_vendor", return_value="USD"), patch(f"{MODULE_RDS_COLLAB_CREATE}.get_or_create_by_names") as mock_create, patch(f"{MODULE_RDS_SPLIT}.get_track_splits_with_rates", return_value={}), patch(f"{MODULE_SF_PRODUCT}.get_product_ids_for_tuids", return_value={}), ): run_bulk_ingest([new_collab_row], dry_run=True) mock_create.assert_not_called() def test_run_bulk_ingest_propagates_validation_error(): """A ValueError from any validation step propagates to the caller.""" rows = [_bulk_row("bad-tuid")] with ( patch(f"{MODULE_RDS_COLLAB}.get_vendor_map_by_ids", return_value={1: 24601}), patch(f"{MODULE_SF_PRODUCT}.get_ids_for_vendor", return_value=set()), patch( f"{MODULE_SF_PRODUCT}.get_vendor_map_by_tuids", return_value={} ), # tuid not found patch(f"{MODULE_PAYMENT}.get_currency_for_vendor", return_value="USD"), ): with pytest.raises(ValueError, match="TUIDs not found"): run_bulk_ingest(rows, dry_run=True)