"""Contract Term Model.""" from abacus_common_logic.connectors.database import db from abacus_common_logic.models.base import BaseModel from sqlalchemy import Enum, func, or_ from abacus_contract.constants import constants from abacus_contract.models.account_contract import AccountContract from abacus_contract.models.contract import Contract from abacus_contract.models.contract_term_condition import ContractTermCondition class ContractTerm(BaseModel): """Contract Term Model.""" __tablename__ = 'contract_term' contract_term_id = db.Column(db.Integer, primary_key=True) contract_id = db.Column( db.Integer, db.ForeignKey('contract.contract_id'), nullable=False ) contract_term_name = db.Column(db.String(180), nullable=True) term_type = db.Column( Enum( *constants.CONTRACT_TERM_TYPES, name='term_type', create_type=False), nullable=False ) attachments = db.Column( db.JSON, nullable=True, default=[] ) attachments_relations = db.Column( db.JSON, nullable=True ) is_base_term = db.Column(db.Boolean, nullable=False, default=True) deleted_at = db.Column(db.DateTime, nullable=True) deleted_by = db.Column(db.String(180), nullable=True) conditions = db.relationship( 'ContractTermCondition', backref='contract_term', cascade='all, delete-orphan' ) contract_term_schedule = db.relationship( 'ContractTermSchedule', backref='contract_term', cascade='all, delete-orphan' ) @classmethod def get_contract_base_term(cls, contract_id): """Return contract's base term.""" return cls.query.filter( cls.contract_id == contract_id, cls.is_base_term.is_(True), cls.deleted_at.is_(None), cls.deleted_by.is_(None) ).first() @classmethod def stream_all(cls, contract_ids=None): """Stream all contract terms.""" db_query = db.session.query(Contract, ContractTerm, ContractTermCondition)\ .join(Contract.contract_terms) \ .join(ContractTerm.conditions, isouter=True) \ .filter(ContractTerm.deleted_at.is_(None)) \ .filter(ContractTermCondition.deleted_at.is_(None)) if contract_ids: db_query = db_query.filter(Contract.contract_id.in_(contract_ids)) return db_query.yield_per(1000) @classmethod def get_contract_terms_by_account_and_term_type( cls, account_id, attachments, term_type ): """Get contract terms for a specific account and term_type.""" contract_attachments = list() for attachment in attachments: contract_attachments.append( func.json_contains(ContractTerm.attachments, f'"{attachment}"') ) return cls.query\ .join(AccountContract, cls.contract_id == AccountContract.contract_id) \ .filter(AccountContract.account_id == account_id) \ .filter(cls.term_type == term_type) \ .filter(cls.deleted_at.is_(None)) \ .filter(or_(*contract_attachments)) \ .all()