"""Split Data Persister. Handles doing CUD operations on the split data table in Snowflake. """ from datetime import datetime from typing import Any from ddtrace.trace import tracer from sqlalchemy import insert, or_, select, tuple_ from sqlalchemy.orm.session import Session from collaborator.connectors import snowflake from collaborator.constants import error from collaborator.constants.split_type import SplitTypeId from collaborator.models.snowflake.split import Split from collaborator.utils.error import OwsError _SNOWFLAKE_IN_CLAUSE_LIMIT = 32_768 class SplitPersister: """Handles high level operations for splits in Snowflake.""" @classmethod @snowflake.db_session def get_splits_for_template( cls, tuids: list[str], session: Session ) -> list[tuple[str, int, float, str]]: """Return (identifier, collaborator_id, split_rate, rate_type) for the given TUIDs. Excludes rows with null collaborator_id and fivetran-deleted rows. Chunks the TUID list to stay within Snowflake's IN-clause limit. """ if not tuids: return [] int_tuids = [str(int(t)) for t in tuids if t.isdigit()] results: list[Any] = [] for i in range(0, len(int_tuids), _SNOWFLAKE_IN_CLAUSE_LIMIT): chunk = int_tuids[i : i + _SNOWFLAKE_IN_CLAUSE_LIMIT] rows = session.execute( select( Split.identifier, Split.collaborator_id, Split.split_rate, Split.rate_type, ) .where( Split.identifier.in_(chunk), Split.split_type_id == SplitTypeId.TRACK, or_( Split.fivetran_deleted.is_(None), Split.fivetran_deleted != True, # noqa: E712 ), ) .distinct() ).all() results.extend(rows) return [ (identifier, collaborator_id, split_rate, rate_type) for identifier, collaborator_id, split_rate, rate_type in results ] @classmethod @tracer.wrap(name="SplitPersister.replace_splits") @snowflake.db_writer_session def replace_splits(cls, split_data: list[dict], deleted_split_ids: list, session): """Replace existing splits with given split data. Args: split_data (SplitsRequest): list of dictionaries [{split_id, split_rate}...] session (sqlalchemy.orm.session.Session): database session. Returns: Response(response): updated status """ _insert_or_update_splits(split_data, session) _delete_splits(deleted_split_ids, session) session.commit() @tracer.wrap(name="SplitPersister._get_splits") def _get_splits(split_data: list, session: Session) -> list: tuid_collabid_map = [ (split["identifier"], split["collaborator_id"]) for split in split_data ] result = ( session.query(Split) .filter(tuple_(Split.identifier, Split.collaborator_id).in_(tuid_collabid_map)) .all() ) return [split.to_dict() for split in result if split is not None] @tracer.wrap(name="SplitPersister._insert_or_update_splits") def _insert_or_update_splits(split_data: list[dict], session: Session): existing_splits = _get_splits(split_data, session) splits_to_insert = [] splits_to_update = [] for split in split_data: split_exists = any( split["collaborator_id"] == existing_split["collaborator_id"] and split["identifier"] == existing_split["identifier"] for existing_split in existing_splits ) if split_exists: splits_to_update.append(split) else: splits_to_insert.append(split) _insert_splits(splits_to_insert, session) _update_splits(splits_to_update, session) @tracer.wrap(name="SplitPersister._insert_splits") def _insert_splits(splits_to_insert: list, session: Session): if len(splits_to_insert) == 0: return insert_split_stmt = insert(Split).values(splits_to_insert) session.execute(insert_split_stmt) @tracer.wrap(name="SplitPersister._update_splits") def _update_splits(splits_to_update: list, session: Session): if len(splits_to_update) == 0: return for item in splits_to_update: row = session.query(Split).filter(Split.split_id == item["id"]).first() if not row: raise OwsError.not_found( code=error.ERROR_CODE_SPLIT_NOT_FOUND, message=error.ERROR_MESSAGE_SPLIT_NOT_FOUND, ) row.split_rate = item["split_rate"] row.collaborator_id = item["collaborator_id"] row.rate_type = item["rate_type"] if "updated_date" in item and item["updated_date"]: row.updated_date = datetime.fromisoformat(item["updated_date"]) session.add(row) session.flush() @tracer.wrap(name="SplitPersister._delete_splits") def _delete_splits(split_ids_to_delete: list, session: Session) -> None: delete_query = session.query(Split).filter(Split.split_id.in_(split_ids_to_delete)) delete_query.delete(synchronize_session=False)