"""Run controller model.""" from abacus_common_logic.connectors.database import db from abacus_common_logic.models import BaseModel, NormalizedDateTime from sqlalchemy import func from royalties.constants import constants from royalties.models.run_controller_contract import RunControllerContract class RunController(BaseModel): """Run controller model.""" __tablename__ = 'run_controller' run_controller_id = db.Column(db.Integer, primary_key=True) contract_type = db.Column( db.Enum(*constants.CONTRACT_TYPES, name='contract_type', create_type=False), nullable=False, ) run_controller_name = db.Column(db.String(255), nullable=False) deleted_at = db.Column(NormalizedDateTime(), nullable=True) deleted_by = db.Column(db.String(180), nullable=True) run_controller_contracts = db.relationship( 'RunControllerContract', backref='run_controller', cascade='all, delete-orphan' ) @classmethod def base_list_query(cls, contract_type=None, active_only=False) -> db.session.query: """Create base query to get list of run_controllers. Args: contract_type (str): contract_type to filter list by active_only (bool): flag to filter list for only active run_controllers Returns: a db query """ query = ( db.session.query( RunController.contract_type, RunController.run_controller_name, RunController.run_controller_id, ) .add_columns( func.count(RunControllerContract.contract_id).label('contract_count') ) .outerjoin( RunControllerContract, RunController.run_controller_id == RunControllerContract.run_controller_id, ) .group_by(RunController.run_controller_id) .order_by(RunController.run_controller_id) ) if contract_type: query = query.filter(cls.contract_type == contract_type) if active_only: query = query.filter(cls.deleted_at.is_(None), cls.deleted_by.is_(None)) return query @classmethod def default_order(cls): """Override to customize default ordering.""" return func.lower(cls.run_controller_name) @classmethod def find_by_name(cls, run_controller_name): """Get run_controller by run_controller_name value. :return: object """ return cls.query.filter_by(run_controller_name=run_controller_name).first() @classmethod def get_all_associated_to_contracts(cls, contract_type): """Get all run controllers that are associated to contracts.""" return ( cls.query.join( RunControllerContract, RunController.run_controller_id == RunControllerContract.run_controller_id, ) .filter(cls.contract_type == contract_type) .group_by(RunController.run_controller_id) .all() ) @classmethod def count_contracts(cls, run_controller_id): """Get count of contracts associated with run_controller. Args: run_controller_id (int): ID of the run controller Returns: int: Count of associated contracts """ return ( db.session.query(func.count(RunControllerContract.contract_id)) .filter(RunControllerContract.run_controller_id == run_controller_id) .scalar() )