"""Split level hierarchy for report generation.""" from abc import ABC, abstractmethod from collections.abc import Iterator from .schemas import CoveredTrackKey, SplitRow from .utils import rds, snowflake class SplitLevel(ABC): @abstractmethod def get_splits(self, report_run_uuid: str, invalid_collaborators: list) -> list[SplitRow]: """Fetch raw splits from RDS.""" @abstractmethod def to_track_splits( self, splits: list[SplitRow], covered: set[CoveredTrackKey], synthetic_ids: Iterator[int], ) -> tuple[list[SplitRow], set[CoveredTrackKey]]: """Return TempSplitRow instances, excluding covered (collaborator_id, tuid) pairs.""" class TrackSplitLevel(SplitLevel): def get_splits(self, report_run_uuid: str, invalid_collaborators: list): return rds.get_track_splits_for_report_run(report_run_uuid, invalid_collaborators) def to_track_splits( self, splits: list[SplitRow], _covered: set[CoveredTrackKey], _synthetic_ids: Iterator[int], ): return splits, { CoveredTrackKey(str(split.collaborator_id), str(split.identifier)) for split in splits } class SubaccountSplitLevel(SplitLevel): def get_splits(self, report_run_uuid: str, invalid_collaborators: list): return rds.get_subaccount_splits_for_report_run(report_run_uuid, invalid_collaborators) def to_track_splits( self, splits: list[SplitRow], covered: set[CoveredTrackKey], synthetic_ids: Iterator[int], ): subaccount_ids = list({s.identifier for s in splits}) tracks_by_subaccount = snowflake.get_track_tuids_by_subaccount(subaccount_ids) new_covered = set(covered) track_splits = [] for split in splits: for tuid in tracks_by_subaccount.get(str(split.identifier), []): key = CoveredTrackKey(str(split.collaborator_id), str(tuid)) if key in new_covered: continue new_covered.add(key) track_splits.append( SplitRow( next(synthetic_ids), split.split_type_id, tuid, split.split_rate, split.collaborator_id, split.rate_type, ) ) return track_splits, new_covered