"""StatementPeriodAdjustmentBatchCriteria model.""" import json from typing import Optional from abacus_common_logic.connectors.database import db from abacus_common_logic.models import BaseModel from sqlalchemy import func, or_ from sqlalchemy.orm import Mapped from royalties.models.statement_period_adjustment_file import ( StatementPeriodAdjustmentFile, ) class StatementPeriodAdjustmentBatchCriteria(BaseModel): """StatementPeriodAdjustmentBatchCriteria model.""" __tablename__ = 'statement_period_adjustment_batch_criteria' statement_period_adjustment_batch_criteria_id: Mapped[int] = db.Column( db.Integer, primary_key=True, autoincrement=True, comment='Primary key.' ) batch_criteria: Mapped[Optional[dict]] = db.Column( db.JSON, nullable=False, comment='The set of attributes used to create adjustment batch. eg: {"payment_schedules": ["30_days_after_quarter_end"], "reference_payment_entities": [1,2,3]}.', ) statement_period_adjustment_file_id: Mapped[int] = db.Column( db.Integer, db.ForeignKey( 'statement_period_adjustment_file.statement_period_adjustment_file_id' ), nullable=False, comment='The statement_period_adjustment_file created using this batch.', ) statement_period_adjustment_file = db.relationship('StatementPeriodAdjustmentFile') @classmethod def get_by_file_id( cls, statement_period_adjustment_file_id: int ) -> Optional['StatementPeriodAdjustmentBatchCriteria']: """Get the statement_period_adjustment_batch_criteria by file id. Args: statement_period_adjustment_file_id (int): ID of the statement_period_adjustment_file Returns: adjustment batch criteria details """ return ( cls.query.join( StatementPeriodAdjustmentFile, cls.statement_period_adjustment_file_id == StatementPeriodAdjustmentFile.statement_period_adjustment_file_id, ) .filter( cls.statement_period_adjustment_file_id == statement_period_adjustment_file_id, StatementPeriodAdjustmentFile.deleted_at.is_(None), StatementPeriodAdjustmentFile.deleted_by.is_(None), ) .first() ) @classmethod def get_by_batch_criteria_and_period_id( cls, batch_criteria: dict, statement_period_id: int ) -> Optional['StatementPeriodAdjustmentBatchCriteria']: """Get the statement_period_adjustment_batch_criteria by file id. Args: batch_criteria (dict): dict of attributes used to create adjustment batch statement_period_id (int): ID of the statement_period Returns: adjustment batch criteria details """ payment_schedule_filter = [ func.json_contains( cls.batch_criteria, json.dumps(val), '$.payment_schedules' ) == 1 for val in batch_criteria['payment_schedules'] ] reference_entities_filter = [ func.json_contains( cls.batch_criteria, json.dumps(val), '$.reference_payment_entities' ) == 1 for val in batch_criteria['reference_payment_entities'] ] return ( cls.query.join( StatementPeriodAdjustmentFile, cls.statement_period_adjustment_file_id == StatementPeriodAdjustmentFile.statement_period_adjustment_file_id, ) .filter( or_(*payment_schedule_filter), or_(*reference_entities_filter), StatementPeriodAdjustmentFile.statement_period_id == statement_period_id, StatementPeriodAdjustmentFile.deleted_at.is_(None), StatementPeriodAdjustmentFile.deleted_by.is_(None), ) .all() )