"""Collaborator reports Model. This is a model class for owscollaborator_reports table in rds. """ from datetime import datetime, timezone from typing import List, Optional from sqlalchemy import and_, asc, desc, func, select from sqlalchemy.dialects.mysql import insert from sqlalchemy.orm import aliased from sqlalchemy.orm.session import Session from collaborator.connectors import mysql from collaborator.constants import error from collaborator.constants import reports as constants from collaborator.constants.reports import ReportStatus from collaborator.models.rds.collaborator import Collaborator from collaborator.models.rds.report import Report from collaborator.models.rds.report_contract import ReportContract from collaborator.models.rds.report_run import ReportRun from collaborator.models.rds.transaction import Transaction from collaborator.utils import logging from collaborator.utils.error import OwsError from collaborator.utils.helpers import sanitize_text from collaborator.utils.typing import Account, Sort, User _DEFAULT_SORT = Sort("requested_datetime", "DESC") class ReportPersister(object): """Report persister class.""" @classmethod def create_reports( cls, collaborators: list, period_name: str, report_run_id: int, report_run_name: str, currency: str, request_datetime: datetime, ): """Create reports in rds. Args: collaborators (list): List of collaborators period_name (str): Name of the period report_run_id (int): ID of the report run report_run_name (str): Name of the report run currency (str): Currency of the report request_datetime (datetime): Datetime when the report was requested Returns: list: data for the created reports in rds """ reports: List[Report] = [] for collaborator in collaborators: filename = sanitize_text( "{report_run}_{collaborator}_{period}_{date}".format( report_run=report_run_name, collaborator=collaborator["name"], period=period_name, date=request_datetime.strftime("%Y%m%d"), ) ) exclude_transaction_types = ( None if collaborator["performance_rights"] else ",".join(constants.TRANSACTION_TYPES_PERFORMANCE_RIGHTS) ) reports.append( Report( report_run_id=report_run_id, amount=None, currency=currency, collaborator_id=collaborator["id"], filename=filename, exclude_transaction_types=exclude_transaction_types, requested_datetime=request_datetime, generated_datetime=None, status=ReportStatus.REQUESTED, collaborator_dp_enabled=bool(collaborator["dp_enabled_date"]), ) ) return reports @classmethod @mysql.db_session def get_report_ids_with_overlapping_period( cls, report_ids: list, session: Session ) -> dict: """Get report ids with overlapping period range. Args: period_ids (list): List of report IDs Returns: dict: map of report IDs to arrays of report IDs with overlapping periods. """ overlaps: dict = {int(report_id): [] for report_id in report_ids} ThisReport = aliased(Report) OtherReport = aliased(Report) ThisReportRun = aliased(ReportRun) OtherReportRun = aliased(ReportRun) # Match up every report with every other report that both a) has the same # collaborator ID and b) has a transaction_id (meaning it is already # added to royalties) so that we can compare their period ID ranges report_to_report_lines = ( session.query( ThisReport.report_id.label("this_report_id"), OtherReport.report_id.label("comp_report_id"), ThisReportRun.period_ids.label("this_period_ids"), OtherReportRun.period_ids.label("comp_period_ids"), ) .join( ThisReportRun, ThisReport.report_run_id == ThisReportRun.report_run_id, ) .join( OtherReport, and_( OtherReport.collaborator_id == ThisReport.collaborator_id, OtherReport.report_id != ThisReport.report_id, OtherReport.transaction_id != None, # noqa: E711 ), ) .join( OtherReportRun, OtherReport.report_run_id == OtherReportRun.report_run_id, ) .filter( ThisReport.report_id.in_(report_ids), ThisReport.transaction_id == None, # noqa: E711 ) .all() ) # Compare every report-to-report matching's period ID ranges for overlaps. # If it has overlaps, add those report IDs to the overlaps map for row in report_to_report_lines: this_report_id, comp_report_id = ( int(row.this_report_id), int(row.comp_report_id), ) this_period_ids = set(row.this_period_ids.split(",")) comp_period_ids = set(row.comp_period_ids.split(",")) if any(("" in this_period_ids, "" in comp_period_ids)): continue if len(this_period_ids & comp_period_ids): overlaps[this_report_id].append(comp_report_id) return overlaps @classmethod @mysql.db_session def get_items( cls, session: Session, account: Optional[Account] = None, sort: Optional[Sort] = None, collaborator_id: Optional[int] = None, has_transaction: Optional[bool] = None, status: Optional[str] = None, report_run_uuid: Optional[str] = None, report_id: Optional[str] = None, collaborator_dp_enabled: Optional[bool] = None, ) -> list: """Get report items from owscollaborator-reports table in rds. Args: account (Account): account to use for query partition key session (Session): DB session sort (Sort?): direction to sort the results in filters (dict): additional filters to use in query Returns: list: collaborator reports """ query = session.query(Report).join( Collaborator, Report.collaborator_id == Collaborator.collaborator_id ) query_filters = [ (Report.deleted_datetime == None), # noqa: E711 ] if account: query_filters.append(Collaborator.vendor_id == account.id) if collaborator_id: query_filters.append(Collaborator.collaborator_id == collaborator_id) if report_id: query_filters.append(Report.report_id == report_id) if has_transaction is not None: query = query.join( Transaction, and_( Transaction.deleted_date == None, # noqa: E711 Transaction.report_id == Report.report_id, ), isouter=True, ) if has_transaction: query_filters.append(Transaction.transaction_id != None) # noqa: E711 else: query_filters.append(Transaction.transaction_id == None) # noqa: E711 if status: query_filters.append(Report.status == status) if collaborator_dp_enabled is not None: query_filters.append( Report.collaborator_dp_enabled == collaborator_dp_enabled ) if report_run_uuid: query = query.join( ReportRun, Report.report_run_id == ReportRun.report_run_id, ) query_filters.append(ReportRun.uuid == report_run_uuid) sort = sort or _DEFAULT_SORT order_attr = getattr(Report, sort.key) order_direction = asc if sort.direction == "ASC" else desc results = ( query.filter(*query_filters).order_by(order_direction(order_attr)).all() ) reports = [item.to_dict() for item in results] return reports @classmethod @mysql.db_session def get_item(cls, account: Account, report_id: int, session: Session) -> dict: """Get a single report item. Args: account (Account): optional param for external clients. report_id(int): report unique identifier session (Session): DB session Returns: dict: a collaborator report """ if account.type != "vendor": raise Exception(f"Unsupported account type: {account.type}") result = ( session.query(Report) .join(Collaborator, Report.collaborator_id == Collaborator.collaborator_id) .filter(Report.report_id == report_id, Collaborator.vendor_id == account.id) .first() ) if not result: raise OwsError.not_found( code=error.ERROR_CODE_REPORT_NOT_FOUND, message=error.ERROR_MESSAGE_REPORT_NOT_FOUND, ) return result.to_dict() @classmethod @mysql.db_session def get_item_with_report_id(cls, report_id: int, session: Session) -> dict: """Get a single report item. Args: report_id(int): report unique identifier session (Session): DB session Returns: dict: a collaborator report """ result = session.query(Report).filter_by(report_id=report_id).first() if not result: raise OwsError.not_found( code=error.ERROR_CODE_REPORT_NOT_FOUND, message=error.ERROR_MESSAGE_REPORT_NOT_FOUND, ) return result.to_dict() @classmethod @mysql.db_session def get_reports_with_report_ids(cls, report_ids: list, session: Session) -> list: """Get a list of reports by given IDs. Args: report_ids (list): report unique identifiers session (Session): DB session Returns: list: of report information """ results = ( session.query(Report) .filter(Report.report_id.in_(report_ids)) .filter(Report.deleted_datetime == None) # noqa: E711 .all() ) if not results: raise OwsError.not_found( code=error.ERROR_CODE_REPORT_NOT_FOUND, message=error.ERROR_MESSAGE_REPORT_NOT_FOUND, ) return [item.to_dict() for item in results] @classmethod @mysql.db_session def bulk_delete_items( cls, account: Account, report_ids: list, user: User, session: Session ) -> list: """Soft-delete report items. Args: account (Account): optional param for external clients. report_ids (list): list with report ids user (User): user tuple session (Session): Database session Returns: list with the deleted reports """ report_query = ( session.query(Report) .join(Collaborator, Collaborator.collaborator_id == Report.collaborator_id) .filter( and_( Report.report_id.in_(report_ids), Collaborator.vendor_id == account.id, Report.transaction_id.is_(None), ) ) ) result = [item.to_dict() for item in report_query.all()] if len(result) != len(report_ids): raise OwsError( code=error.ERROR_CODE_INVALID_REPORT_IDS, message=error.ERROR_MESSAGE_INVALID_REPORT_IDS, ) if report_query: session.query(Report).filter(Report.report_id.in_(report_ids)).update( { "deleted_by": user.id, "deleted_datetime": datetime.now(tz=timezone.utc), }, synchronize_session=False, ) session.commit() for report in result: logging.log_event( logging.LOG_EVENT_DELETE, "report", report["id"], report, None, user ) return result @classmethod @mysql.db_session def update_reports_with_transactions( cls, transactions, user: User, session: Session ) -> list: """Add the transaction_id that corresponds to a report. Args: transactions(list): list with created transactions account (Account): optional param for external clients. user (User): optional param for orchard user session (Session): Database session Returns: list: with the updated reports """ report_ids = [t["report_id"] for t in transactions] query = ( session.query(Report) .with_for_update() .filter(Report.report_id.in_(report_ids)) ) previous_reports = { report.report_id: report.to_dict() for report in query.all() } if session.connection().dialect.name == "mysql": values = [ { "id": txn["report_id"], "transaction_id": txn["transaction_id"], } for txn in transactions if previous_reports.get(txn["report_id"]) ] stmt = insert(Report).values(values) session.execute( stmt.on_duplicate_key_update( { "transaction_id": stmt.inserted.transaction_id, } ) ) else: session.bulk_update_mappings(Report, transactions) # type: ignore[arg-type] session.flush() updated_reports = [report.to_dict() for report in query.all()] logging.bulk_log_events( logging.LOG_EVENT_UPDATE, "report", user=user, data=[ { "id": report["id"], "original": previous_reports.get(report["id"]), "updated": report, } for report in updated_reports ], ) return updated_reports @classmethod @mysql.db_session def update_report_status(cls, report_id: int, status: str, session: Session): """Update the status of a report. Args: report_id (int): report unique identifier status (str): report status session (Session): Database session """ report = ( session.query(Report) .filter_by(report_id=report_id) .filter(Report.deleted_datetime == None) # noqa: E711 .first() ) if not report: raise OwsError.not_found( code=error.ERROR_CODE_REPORT_NOT_FOUND, message=error.ERROR_MESSAGE_REPORT_NOT_FOUND, ) report_previous_state = report.to_dict() report.status = ReportStatus(status) session.commit() logging.log_event( logging.LOG_EVENT_UPDATE, "report", report_id, report_previous_state, report.to_dict(), None, ) @classmethod @mysql.db_session def remove_transaction_from_report(cls, transaction_id: int, session: Session): """Update the status of a report. Args: transaction_id (int): report unique identifier status (str): report status session (Session): Database session """ report = session.query(Report).filter_by(transaction_id=transaction_id).first() report_previous_state = report.to_dict() if report else None if report: report.transaction_id = None session.commit() logging.log_event( logging.LOG_EVENT_UPDATE, "report", report.report_id, report_previous_state, report.to_dict(), None, ) @classmethod @mysql.db_session def remove_transactions_from_reports( cls, transaction_ids: List[int], session: Session ): """Remove transactions from reports. Args: transaction_ids (List[int]): List of transaction ids. session (Session) """ reports_base_query = session.query(Report).filter( Report.transaction_id.in_(transaction_ids) ) reports = reports_base_query.all() original_report_dicts = { report.report_id: report.to_dict() for report in reports } reports_base_query.update({"transaction_id": None}) session.commit() updated_report_dicts = { report.report_id: report.to_dict() for report in reports } logging.bulk_log_events( logging.LOG_EVENT_UPDATE, "report", data=[ { "id": report_id, "original": original_report_dicts.get(report_id), "updated": updated_report_dicts.get(report_id), } for report_id in original_report_dicts.keys() ], user=None, ) @classmethod @mysql.db_session def get_latest_reports_by_collaborator_ids( cls, collaborator_ids: List[int], session: Session ): """Get latest reports by collaborator IDs. Args: collaborator_ids (List[int]): List of collaborator IDs session (Session): Database session Returns: List: List of reports """ Report1 = aliased(Report) Report2 = aliased(Report) query = ( select( Report1.report_id, Report1.collaborator_id, Collaborator.vendor_id, ) .join( Report2, and_( Report1.collaborator_id == Report2.collaborator_id, Report1.requested_datetime < Report2.requested_datetime, Report2.deleted_datetime == None, # noqa ), isouter=True, ) .join(Collaborator, Report1.collaborator_id == Collaborator.collaborator_id) .where( and_( Report1.collaborator_id.in_(collaborator_ids), Report1.deleted_datetime == None, # noqa Report2.report_id == None, # noqa ) ) ) return session.execute(query).all() @classmethod @mysql.db_session def get_report_contract_subtotals( cls, report_ids: List[int], session: Session, ): """Get non-zero report contract subtotals for given report IDs. Zero-amount entries are excluded. Args: report_ids (List[int]): List of report IDs to get subtotals for contracts. session (Session): DB session Returns: List[Dict[str, Any]]: Report contract subtotals """ return ( session.query( ReportContract.report_id, ReportContract.contract_id, ReportContract.currency, ReportContract.amount, Collaborator.vendor_id, ) .join(Report, Report.report_id == ReportContract.report_id) .join(Collaborator, Collaborator.collaborator_id == Report.collaborator_id) .filter( ReportContract.report_id.in_(report_ids), ReportContract.amount != 0, ) .all() ) @classmethod @mysql.db_session def get_report_contract_subtotal_aggregations( cls, report_run_id: int, collaborator_dp_enabled: bool | None, session: Session, ): """Get report contract subtotal aggregations for a given report run ID. Zero-amount entries are excluded. """ query = ( select( func.sum(ReportContract.amount).label("currency_agnostic_total_amount"), func.count().label("total_count"), ) .select_from(ReportContract) .join(Report, Report.report_id == ReportContract.report_id) .where( Report.report_run_id == report_run_id, ReportContract.amount != 0, ) ) if collaborator_dp_enabled is not None: query = query.where( Report.collaborator_dp_enabled == collaborator_dp_enabled ) return session.execute(query).one()