"""Bulk split ingest logic — web API entry point for bulk split ingestion.""" from collections import defaultdict from datetime import datetime, timezone import logging from typing import Literal from pydantic import Field, model_validator from collaborator.models.rds.collaborator_persister import CollaboratorPersister from collaborator.models.rds.recipient import Recipient # noqa: F401 from collaborator.models.rds.split_persister import SplitPersister from collaborator.models.rds.split_type import SplitType # noqa: F401 from collaborator.models.snowflake.account_payment_term_persister import ( AccountPaymentTermPersister, ) from collaborator.models.snowflake.product_persister import ProductPersister from collaborator.schemas import BaseSchema from collaborator.schemas.split import ( BulkSplitRow, IngestSummarySchema, OverwrittenSplitSchema, ) log = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Internal models # --------------------------------------------------------------------------- class SplitIngestionRow(BaseSchema): """A single split entry in the bulk ingestion pipeline.""" vendor_id: int collaborator_id: int | None = None collaborator_name: str | None = None split_rate: float = Field(gt=0, le=1) split_type: Literal["NET", "GROSS"] tuid: str | None = None product_id: int | None = None @model_validator(mode="after") def must_have_identifier(self) -> "SplitIngestionRow": """Require at least one of tuid or product_id.""" if not self.tuid and not self.product_id: raise ValueError("Each row must have either a tuid or a Product ID") return self def validate_single_vendor(rows: list[SplitIngestionRow]) -> int: """Ensure all rows share a single vendor_id and return it.""" vendor_ids = {row.vendor_id for row in rows} if len(vendor_ids) != 1: raise ValueError( f"File must contain exactly one vendor ID; found: {sorted(vendor_ids)}" ) return vendor_ids.pop() # --------------------------------------------------------------------------- # DB-backed validation + product ID expansion # --------------------------------------------------------------------------- def validate_collaborators(rows: list[SplitIngestionRow], vendor_id: int) -> None: """Verify all collaborator IDs in the file exist and belong to the vendor.""" collaborator_ids = { r.collaborator_id for r in rows if r.collaborator_id is not None } if not collaborator_ids: return found = CollaboratorPersister.get_vendor_map_by_ids(collaborator_ids) errors = [] for cid in sorted(collaborator_ids): if cid not in found: errors.append(f"Collaborator ID {cid} not found") elif found[cid] != vendor_id: errors.append( f"Collaborator ID {cid} belongs to vendor {found[cid]}, not {vendor_id}" ) if errors: raise ValueError("Collaborator validation failed:\n" + "\n".join(errors)) def resolve_product_ids( rows: list[SplitIngestionRow], vendor_id: int ) -> tuple[list[SplitIngestionRow], int]: """Expand product-ID-only rows to TUID rows; validate product IDs belong to the vendor. Returns (rows, skipped_count) where rows replaces all product-ID-only rows with one row per track under that product, and skipped_count is the number of product-ID rows dropped because the product had no eligible tracks. Rows that already have a tuid are unchanged. """ product_id_rows = [r for r in rows if r.product_id and not r.tuid] if not product_id_rows: return rows, 0 product_ids = {r.product_id for r in product_id_rows} valid_ids = ProductPersister.get_ids_for_vendor(product_ids, vendor_id) missing = product_ids - valid_ids if missing: raise ValueError( f"Product IDs not found for vendor {vendor_id}: {sorted(missing)}" ) track_pairs = ProductPersister.get_tuid_product_pairs(valid_ids) tuids_by_product: dict[int, list[str]] = defaultdict(list) for tuid_str, product_id in track_pairs: tuids_by_product[product_id].append(tuid_str) result = [r for r in rows if r.tuid] # preserve tuid-only rows unchanged skipped_ids: list[int] = [] for row in product_id_rows: pid = row.product_id assert pid is not None tuid_list = tuids_by_product.get(pid, []) if not tuid_list: skipped_ids.append(pid) continue for tuid_str in tuid_list: result.append(row.model_copy(update={"tuid": tuid_str})) if skipped_ids: log.warning( "%d product ID(s) had no eligible tracks and were skipped: %s", len(skipped_ids), sorted(skipped_ids), ) return result, len(skipped_ids) def validate_tracks(rows: list[SplitIngestionRow], vendor_id: int) -> None: """Verify all TUIDs in the file exist in Snowflake and belong to the vendor.""" tuids = {r.tuid for r in rows if r.tuid} if not tuids: return tuid_vendor_map = ProductPersister.get_vendor_map_by_tuids(tuids) missing = sorted(tuids - tuid_vendor_map.keys()) if missing: raise ValueError(f"TUIDs not found in Snowflake: {missing}") invalid = sorted(t for t, vid in tuid_vendor_map.items() if vid != vendor_id) if invalid: raise ValueError(f"TUIDs do not belong to vendor {vendor_id}: {invalid}") # --------------------------------------------------------------------------- # Change computation # --------------------------------------------------------------------------- def compute_changes( rows: list[SplitIngestionRow], vendor_id: int, input_rows: int, vendor_currency: str, products_skipped_no_tracks: int = 0, newly_created_collabs: int = 0, dry_run: bool = True, ) -> IngestSummarySchema: """Compute what the ingest would do by comparing file rows against existing splits. All rows must already have tuid set (call resolve_product_ids before this). """ collaborator_ids = { r.collaborator_id for r in rows if r.collaborator_id is not None } new_collaborator_names = { r.collaborator_name for r in rows if r.collaborator_id is None and r.collaborator_name } tuids = {r.tuid for r in rows if r.tuid} # {tuid: {collaborator_id: split_rate}} for all existing SPLIT_TYPE_TRACK splits existing_by_tuid = SplitPersister.get_track_splits_with_rates(tuids) tuid_to_product_id = ProductPersister.get_product_ids_for_tuids(tuids) # (tuid, collaborator_id) → new split rate from the ingest file ingest_rate: dict[tuple[str, int], float] = { (r.tuid, r.collaborator_id): r.split_rate for r in rows if r.tuid and r.collaborator_id is not None } # collaborator_id → name from the ingest file collab_name_by_id: dict[int, str | None] = { r.collaborator_id: r.collaborator_name for r in rows if r.collaborator_id is not None } # Group known-collaborator rows by TUID → set of collaborator IDs ingest_by_tuid: dict[str, set[int]] = defaultdict(set) for r in rows: if r.tuid and r.collaborator_id is not None: ingest_by_tuid[r.tuid].add(r.collaborator_id) # Group new-collaborator rows by TUID → set of names (deduplicates same name per TUID) new_collab_by_tuid: dict[str, set[str]] = defaultdict(set) for r in rows: if r.tuid and r.collaborator_id is None and r.collaborator_name: new_collab_by_tuid[r.tuid].add(r.collaborator_name) tracks_with_existing_updated = 0 tracks_with_new = 0 tracks_with_unaffected = 0 existing_splits_updated = 0 new_splits_created = 0 impacted_tuids: set[str] = set() overwritten_splits: list[OverwrittenSplitSchema] = [] collabs_with_new_splits: set[int] = set() for tuid in tuids: existing_collab_rates = existing_by_tuid.get(tuid, {}) existing_collaborator_ids = set(existing_collab_rates.keys()) tuid_collaborator_ids = ingest_by_tuid.get(tuid, set()) tuid_new_collab_names = new_collab_by_tuid.get(tuid, set()) overlapping = tuid_collaborator_ids & existing_collaborator_ids new_for_track = tuid_collaborator_ids - existing_collaborator_ids unaffected = existing_collaborator_ids - tuid_collaborator_ids existing_splits_updated += len(overlapping) new_splits_created += len(new_for_track) + len(tuid_new_collab_names) if overlapping: tracks_with_existing_updated += 1 if new_for_track or tuid_new_collab_names: tracks_with_new += 1 if unaffected: tracks_with_unaffected += 1 if overlapping or new_for_track or tuid_new_collab_names: impacted_tuids.add(tuid) for collab_id in sorted(overlapping): overwritten_splits.append( OverwrittenSplitSchema( product_id=tuid_to_product_id.get(tuid), tuid=tuid, collaborator_name=collab_name_by_id.get(collab_id), old_split_rate=existing_collab_rates[collab_id], new_split_rate=ingest_rate.get((tuid, collab_id), 0.0), ) ) collabs_with_new_splits.update(new_for_track) impacted_product_ids = { tuid_to_product_id[tuid] for tuid in impacted_tuids if tuid in tuid_to_product_id } return IngestSummarySchema( vendor_id=vendor_id, input_rows=input_rows, expanded_rows=len(rows), existing_collaborators_matched=len(collaborator_ids) - newly_created_collabs, new_collaborators=len(new_collaborator_names) + newly_created_collabs, tracks_with_existing_splits_updated=tracks_with_existing_updated, tracks_with_new_splits=tracks_with_new, tracks_with_unaffected_splits=tracks_with_unaffected, existing_splits_updated=existing_splits_updated, new_splits_created=new_splits_created, products_skipped_no_tracks=products_skipped_no_tracks, vendor_currency=vendor_currency, dry_run=dry_run, splits_ingested=existing_splits_updated + new_splits_created, tracks_impacted=tracks_with_existing_updated + tracks_with_new, products_impacted=len(impacted_product_ids), collaborators_receiving_new_splits=len(collabs_with_new_splits), overwritten_splits=overwritten_splits, ) # --------------------------------------------------------------------------- # API entry point # --------------------------------------------------------------------------- def run_bulk_ingest( splits_config: list[BulkSplitRow], dry_run: bool = True, ticket_id: str | None = None, ) -> IngestSummarySchema: """Orchestrate the bulk ingest pipeline for the web API. Args: splits_config: List of validated BulkSplitRow instances from the handler. dry_run: If True, validate and compute changes without writing to DB. ticket_id: Required when dry_run is False; used as created_by/updated_by. Returns: IngestSummarySchema with counts of what was (or would be) changed. """ rows = [ SplitIngestionRow( vendor_id=row.vendor_id, collaborator_id=row.collaborator_id, collaborator_name=row.collaborator_name, split_rate=row.split_rate, split_type=row.rate_type, tuid=row.tuid, product_id=row.product_id, ) for row in splits_config ] vendor_id = validate_single_vendor(rows) if not dry_run and not ticket_id: date_str = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M UTC") ticket_id = f"From UI {date_str} - vendor {vendor_id}" input_rows = len(rows) validate_collaborators(rows, vendor_id) rows, products_skipped_no_tracks = resolve_product_ids(rows, vendor_id) validate_tracks(rows, vendor_id) currency = AccountPaymentTermPersister.get_currency_for_vendor(vendor_id) newly_created_collabs = 0 if not dry_run and any(r.collaborator_id is None for r in rows): names_needed = { r.collaborator_name for r in rows if r.collaborator_id is None and r.collaborator_name } name_to_id, newly_created_collabs = ( CollaboratorPersister.get_or_create_by_names( names_needed, vendor_id, ticket_id, currency ) ) rows = [ ( row.model_copy( update={"collaborator_id": name_to_id[row.collaborator_name]} ) if row.collaborator_id is None and row.collaborator_name in name_to_id else row ) for row in rows ] summary = compute_changes( rows, vendor_id, input_rows, currency, products_skipped_no_tracks, newly_created_collabs=newly_created_collabs, dry_run=dry_run, ) if not dry_run: split_data = [ { "tuid": r.tuid, "collaborator_id": r.collaborator_id, "split_rate": r.split_rate, "split_type": r.split_type, } for r in rows if r.tuid and r.collaborator_id is not None ] SplitPersister.upsert_track_splits(split_data, ticket_id) return summary