"""Split Persister. Handles doing CRUD operations on the split table. """ from collections import defaultdict from typing import Optional import sqlalchemy from sqlalchemy import and_, or_, select, text, tuple_ from sqlalchemy.dialects.mysql import insert as mysql_insert from sqlalchemy.dialects.sqlite import insert as sqlite_insert from sqlalchemy.orm.session import Session from collaborator.connectors import mysql from collaborator.constants import error from collaborator.constants.split_type import SplitTypeId from collaborator.models.rds.collaborator import Collaborator from collaborator.models.rds.split import Split 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 from collaborator.utils.needs_tcs_agreement import needs_ts_and_cs_agreement from collaborator.utils.typing import User class SplitPersister: """Handles high level operations for splits.""" @classmethod @mysql.db_session def get_for_identifiers(cls, split_type_id, identifiers, session): """Get splits for identifiers. Args: split_type_id (int): type of split. identifiers (List[str]): identifiers of the resources. session (sqlalchemy.orm.session.Session): database session. Returns: list: list of splits """ filters = [ (Split.identifier.in_(identifiers)), (Split.split_type_id == split_type_id), (Split.collaborator_id == Collaborator.collaborator_id), ] results = session.query(Split, Collaborator).filter(*filters).all() return [item[0].to_dict() for item in results] @classmethod @mysql.db_session def get_for_collaborator( cls, collaborator_id: int, split_types: Optional[list[SplitTypeId]] = None, *, session: Session, ) -> list[Split]: """Get splits based on a collaborator ID, optionally filtered by type. Args: collaborator_id (int): ID of the collaborator to get splits for. split_types (list[SplitTypeId]): optional split types to filter by. session (sqlalchemy.orm.session.Session): Database session. Returns: list: list of splits """ filters = [ Split.collaborator_id == collaborator_id, ] if split_types: filters.append(Split.split_type_id.in_(split_types)) results = session.query(Split).filter(*filters).all() return results @classmethod @mysql.db_session def get_for_ids(cls, ids, session): """Get splits based on a list of IDs. Args: ids (list): list of IDs for splits to get. session (sqlalchemy.orm.session.Session): database session. Returns: list: list of splits """ results = session.query(Split).filter(Split.split_id.in_(ids)).all() return [item.to_dict() for item in results] @classmethod @mysql.db_session def get_for_replacement_splits( cls, replacements: list[ReplacementSplitsSchema], vendor_id: int, session, ) -> list[Split]: """Get existing splits from replacement split identifer/split type pairs. Args: replacements (list[IdentifierSplitsSchema]): List of identifier splits. vendor_id (int): The vendor ID to filter by. session (sqlalchemy.orm.session.Session): database session. """ or_filters = [] for ident_split in replacements: and_filters = [ Split.identifier == ident_split.identifier, Split.split_type_id == ident_split.split_type_id, ] or_filters.append(and_(*and_filters)) query = ( session.query(Split) .join(Collaborator, Split.collaborator_id == Collaborator.collaborator_id) .where(Collaborator.vendor_id == vendor_id) .where(or_(*or_filters)) ) return query.all() @classmethod @mysql.db_session def replace_track_splits( cls, split_data: dict, user: User, has_direct_payments: bool, session: Session, ) -> tuple: """Insert, update and/or delete splits based on the provided data. Args: split_data (list): list of dictionaries with split data. user (User): the user making the change session (sqlalchemy.orm.session.Session): database session. skip_checks (bool): whether to skip T&Cs agreement check and DP validations. Returns: tuple: (list of final splits, list of deleted split IDs) """ track_ids = [track["tuid"] for track in split_data["tracks"]] previous_splits = cls.get_for_identifiers(SplitTypeId.TRACK, track_ids) previous_split_ids = [split["id"] for split in previous_splits] needs_tcs_agreement = needs_ts_and_cs_agreement(split_data, previous_splits) if needs_tcs_agreement and not split_data.get("dp_splits_agreed"): raise OwsError( code=error.ERROR_CODE_BAD_PARAMS, message=error.ERROR_MESSAGE_TERMS_AND_CONDITIONS_MUST_BE_AGREED, ) update_dp_agreement = needs_tcs_agreement and split_data.get("dp_splits_agreed") collaborator_ids, tuid_to_collab_id_map = _map_tuids_to_collab_ids( split_data, has_direct_payments ) if update_dp_agreement: session.query(Collaborator).filter( Collaborator.collaborator_id.in_(collaborator_ids), Collaborator.dp_enabled_date != None, # noqa ).update( {Collaborator.dp_splits_agreed_date: sqlalchemy.func.now()}, synchronize_session=False, ) inserted_or_updated_splits = _insert_track_splits( split_data, tuid_to_collab_id_map, user.id, session ) inserted_split_ids = [item["id"] for item in inserted_or_updated_splits] deleted_split_ids = _delete_track_splits( inserted_split_ids, previous_split_ids, user.id, session ) session.commit() final_splits = cls.get_for_identifiers(SplitTypeId.TRACK, track_ids) _log_events(final_splits, user, session) return (final_splits, deleted_split_ids) @classmethod @mysql.db_session def replace_splits( cls, replacements: list[ReplacementSplitsSchema], vendor_id: int, has_direct_payments: bool, dp_splits_agreed: bool, user: User, session: Session, ) -> tuple: """Insert, update and/or delete splits based on the provided data. Args: replacements (list): list of replacement split data. vendor_id (int): the vendor ID. has_direct_payments (bool): whether the vendor has direct payments enabled. dp_splits_agreed (bool): whether the user has agreed to the DP terms. user (User): the user making the change session (sqlalchemy.orm.session.Session): database session. Returns: tuple: (list of final splits, list of deleted split IDs) """ current_splits = cls.get_for_replacement_splits( replacements, vendor_id, session=session ) current_collab_ids = {split.collaborator_id for split in current_splits} replacement_collab_ids = set( [ split.collaborator_id for replacement in replacements for split in replacement.splits ] ) modified_collab_ids = list(current_collab_ids | replacement_collab_ids) if has_direct_payments: split_util.validate_split_rates(replacements) collab_ids_needing_agreement = split_util.get_collabs_needing_tcs_agreement( modified_collab_ids ) if collab_ids_needing_agreement and not dp_splits_agreed: raise OwsError( code=error.ERROR_CODE_BAD_PARAMS, message=error.ERROR_MESSAGE_TERMS_AND_CONDITIONS_MUST_BE_AGREED, ) session.query(Collaborator).filter( Collaborator.collaborator_id.in_(collab_ids_needing_agreement), Collaborator.dp_enabled_date != None, # noqa ).update( {Collaborator.dp_splits_agreed_date: sqlalchemy.func.now()}, synchronize_session=False, ) replacements_by_unique_key = { ( replacement.identifier, replacement.split_type_id, split.collaborator_id, ): (replacement, split) for replacement in replacements for split in replacement.splits } deleted_split_ids = _delete_splits_not_being_replaced( current_splits, replacements_by_unique_key, user.id, session ) _insert_or_update_replacement_splits( replacements_by_unique_key, user.id, session, ) session.commit() session.expire_all() final_splits = cls.get_for_replacement_splits( replacements, vendor_id, session=session ) split_util.log_upsert_events(final_splits, user) return (final_splits, deleted_split_ids) @classmethod @mysql.db_session def get_track_splits_with_rates( cls, tuids: set[str], session: Session ) -> dict[str, dict[int, float]]: """Return {tuid: {collaborator_id: split_rate}} for existing track splits.""" if not tuids: return {} rows = session.execute( select(Split.identifier, Split.collaborator_id, Split.split_rate).where( Split.identifier.in_(tuids), Split.split_type_id == SplitTypeId.TRACK, ) ).all() result: dict[str, dict[int, float]] = defaultdict(dict) for identifier, collaborator_id, split_rate in rows: result[identifier][collaborator_id] = split_rate return dict(result) @classmethod @mysql.db_session def upsert_track_splits( cls, split_data: list[dict], ticket_id: str, session: Session, ) -> None: """Upsert SplitTypeId.TRACK splits. Each dict in split_data must have keys: tuid, collaborator_id, split_rate, split_type. Existing splits are updated in-place; new splits are inserted. Splits for collaborators not present in split_data are left untouched. """ tuids = {row["tuid"] for row in split_data} collaborator_ids = {row["collaborator_id"] for row in split_data} existing: dict[tuple[str, int], Split] = { (s.identifier, s.collaborator_id): s for s in session.execute( select(Split).where( Split.identifier.in_(tuids), Split.collaborator_id.in_(collaborator_ids), Split.split_type_id == SplitTypeId.TRACK, ) ) .scalars() .all() } for row in split_data: key = (row["tuid"], row["collaborator_id"]) if key in existing: s = existing[key] s.split_rate = row["split_rate"] s.rate_type = row["split_type"] s.updated_by = ticket_id else: session.add( Split( identifier=row["tuid"], collaborator_id=row["collaborator_id"], split_rate=row["split_rate"], split_type_id=SplitTypeId.TRACK, rate_type=row["split_type"], created_by=ticket_id, ) ) session.commit() def _map_tuids_to_collab_ids(split_data: dict, has_direct_payments: bool) -> tuple: collaborator_ids_set = set() tuid_collabid_map = [] for track in split_data["tracks"]: track_splits_total_percentage = 0 for split in track["splits"]: collaborator_ids_set.add(split.get("collaborator_id")) tuid_collabid_map.append((track["tuid"], split.get("collaborator_id"))) track_splits_total_percentage += split.get("split_rate") if has_direct_payments and track_splits_total_percentage > 1: raise OwsError( code=error.ERROR_CODE_BAD_PARAMS, message=error.ERROR_MESSAGE_SPLIT_PERCENTAGE_OVER_100, ) collaborator_ids = list(collaborator_ids_set) return (collaborator_ids, tuid_collabid_map) def _log_events(inserted_or_updated_splits: list, user: User, session: Session) -> None: event_data = [ { "id": split["id"], "original": next( s for s in inserted_or_updated_splits if s["id"] == split["id"] ), "updated": split, } for split in inserted_or_updated_splits ] logging.bulk_log_events(logging.LOG_EVENT_UPDATE, "split", event_data, user) def _insert_track_splits( split_data: dict, uid_collabid_map: list, user_id: str, session: Session ) -> list: split_data_to_insert = [ { "id": split.get("id"), "identifier": track["tuid"], "collaborator_id": split["collaborator_id"], "split_rate": split["split_rate"], "split_type_id": split["split_type_id"], "rate_type": split["rate_type"], "created_by": user_id, } for track in split_data["tracks"] for split in track["splits"] ] if len(split_data_to_insert) != 0: updated_splits_tuid_collabid_map = _get_track_splits_to_update( split_data_to_insert, uid_collabid_map, session ) updated_by_expr = sqlalchemy.case( ( tuple_(Split.identifier, Split.collaborator_id).in_( updated_splits_tuid_collabid_map ), user_id, ), else_=Split.updated_by, ) dialect_name = session.connection().dialect.name if dialect_name == "mysql": mysql_stmt = mysql_insert(Split).values(split_data_to_insert) session.execute( mysql_stmt.on_duplicate_key_update( split_rate=mysql_stmt.inserted.split_rate, rate_type=mysql_stmt.inserted.rate_type, updated_by=updated_by_expr, ) ) elif dialect_name == "sqlite": sqlite_stmt = sqlite_insert(Split).values(split_data_to_insert) session.execute( sqlite_stmt.on_conflict_do_update( index_elements=["identifier", "collaborator_id", "split_type_id"], set_={ "split_rate": sqlite_stmt.excluded.split_rate, "rate_type": sqlite_stmt.excluded.rate_type, "updated_by": updated_by_expr, }, ) ) else: raise NotImplementedError( f"Conflict resolution not implemented for {dialect_name}" ) inserted_or_updated_splits_query = session.query(Split).filter( tuple_(Split.identifier, Split.collaborator_id).in_(uid_collabid_map) ) inserted_or_updated_splits = [ row.to_dict() for row in inserted_or_updated_splits_query.all() ] return inserted_or_updated_splits def _get_track_splits_to_update( split_data_to_insert: list, uid_collabid_map: list, session: Session ) -> list: existing_splits = ( session.query(Split) .filter(tuple_(Split.identifier, Split.collaborator_id).in_(uid_collabid_map)) .all() ) existing_splits_map = {} for split in existing_splits: key = (split.identifier, split.collaborator_id) existing_splits_map[key] = split splits_to_update_tuid_collabid_map = [] for split_to_insert in split_data_to_insert: key = (split_to_insert["identifier"], split_to_insert["collaborator_id"]) existing_split = existing_splits_map.get(key) rate_changed = ( existing_split and split_to_insert and existing_split.split_rate != split_to_insert["split_rate"] ) type_changed = ( existing_split and split_to_insert and existing_split.rate_type != split_to_insert["rate_type"] ) if rate_changed or type_changed: splits_to_update_tuid_collabid_map.append(key) return splits_to_update_tuid_collabid_map def _delete_track_splits( splits_ids_after_insert: list, previous_split_ids: list, user_id: str, session: Session, ) -> list: split_ids_to_delete = set(previous_split_ids) - set(splits_ids_after_insert) is_mysql = ( bool(split_ids_to_delete) and session.connection().dialect.name == "mysql" ) if is_mysql: session.execute(text("SET @deleted_by = :user_id"), {"user_id": user_id}) session.query(Split).filter(Split.split_id.in_(split_ids_to_delete)).delete( synchronize_session=False ) if is_mysql: session.execute(text("SET @deleted_by = NULL")) return list(split_ids_to_delete) def _insert_or_update_replacement_splits( replacements_by_unique_key: dict[ tuple, tuple[ReplacementSplitsSchema, ReplacementSplitSchema] ], user_id: str, session: Session, ) -> None: split_data_to_upsert = [ { "identifier": replacement.identifier, "split_type_id": replacement.split_type_id, "collaborator_id": split.collaborator_id, "split_rate": split.split_rate, "rate_type": split.rate_type, "created_by": user_id, } for replacement, split in replacements_by_unique_key.values() ] if not split_data_to_upsert: return dialect_name = session.connection().dialect.name if dialect_name == "mysql": mysql_stmt = mysql_insert(Split).values(split_data_to_upsert) session.execute( mysql_stmt.on_duplicate_key_update( split_rate=mysql_stmt.inserted.split_rate, rate_type=mysql_stmt.inserted.rate_type, updated_by=user_id, ) ) elif dialect_name == "sqlite": sqlite_stmt = sqlite_insert(Split).values(split_data_to_upsert) session.execute( sqlite_stmt.on_conflict_do_update( index_elements=["identifier", "collaborator_id", "split_type_id"], set_={ "split_rate": sqlite_stmt.excluded.split_rate, "rate_type": sqlite_stmt.excluded.rate_type, "updated_by": user_id, }, ) ) else: raise NotImplementedError( f"Conflict resolution not implemented for {dialect_name}" ) def _delete_splits_not_being_replaced( current_splits: list[Split], replacements_by_unique_key: dict[ tuple, tuple[ReplacementSplitsSchema, ReplacementSplitSchema] ], user_id: str, session: Session, ) -> list: split_ids_to_delete = [ split.split_id for split in current_splits if (split.identifier, split.split_type_id, split.collaborator_id) not in replacements_by_unique_key ] if not split_ids_to_delete: return [] is_mysql = session.connection().dialect.name == "mysql" if is_mysql: session.execute(text("SET @deleted_by = :user_id"), {"user_id": user_id}) session.query(Split).filter(Split.split_id.in_(split_ids_to_delete)).delete( synchronize_session=False ) if is_mysql: session.execute(text("SET @deleted_by = NULL")) return split_ids_to_delete