"""Worksheet payment contract advance.""" import typing from abacus_common_logic.connectors.database import db from abacus_common_logic.models.base import BaseModel from flask import abort from sqlalchemy import and_, func, literal_column, or_, select, table from payment.constants import error from payment.constants.constants import ( WORKSHEET_PAYMENT_CONTRACT_ADVANCE_ACTIONS, WORKSHEET_PAYMENT_CONTRACT_ADVANCE_STATUSES, ) class WorksheetPaymentContractAdvance(BaseModel): """Worksheet Payment Contract Advance Model.""" __tablename__ = 'worksheet_payment_contract_advance' worksheet_payment_contract_advance_id = db.Column(db.Integer, primary_key=True) contract_advance_id = db.Column(db.Integer, nullable=False) statement_period_id = db.Column(db.Integer, nullable=False) exchange_rate_statement_period_id = db.Column(db.Integer, nullable=False) payment_name = db.Column(db.String(180), nullable=False) amount = db.Column(db.Numeric(20, 2), nullable=False) currency_code = db.Column(db.String(3), nullable=False) amount_payee_currency = db.Column(db.Numeric(20, 2), nullable=False) payee_currency_code = db.Column(db.String(3), nullable=False) exchange_rate = db.Column(db.Numeric(30, 19), nullable=False) withholding_tax_amount = db.Column(db.Numeric(20, 2), nullable=False) vat_amount = db.Column(db.Numeric(20, 2), nullable=False) amount_after_withholding_and_vat = db.Column(db.Numeric(20, 2), nullable=False) withholding_tax_amount_payee_currency = db.Column(db.Numeric(20, 2), nullable=False) vat_amount_payee_currency = db.Column(db.Numeric(20, 2), nullable=False) amount_after_withholding_and_vat_payee_currency = db.Column( db.Numeric(20, 2), nullable=False ) us_source_income_rate = db.Column(db.Numeric(9, 6), nullable=True) salesforce_id = db.Column(db.String(180), nullable=True) is_internal = db.Column(db.Boolean, nullable=False, default=False) deleted_at = db.Column(db.DateTime, nullable=True) deleted_by = db.Column(db.String(255), nullable=True) filter_args = typing.TypedDict( 'filter_args', { 'salesforce_id': str, 'contract_advance_id': str, 'payment_statuses': typing.List[str], 'worksheet_ids': typing.List[str], 'active': bool, }, total=False, ) @classmethod def get_by_id(cls, obj_id): """Get object from DB by ID property. :return: object """ return ( db.session.execute( select(cls).where( cls.deleted_at.is_(None), cls.worksheet_payment_contract_advance_id == obj_id, ) ) .scalars() .first() ) @classmethod def get_by_contract_advance_id_or_error(cls, contract_advance_id, error_status=400): """Find item by contract_advance_id. Abort request if not found.""" obj = cls.get_by_contract_advance_id(contract_advance_id) if not obj: abort( code=error_status, description=error.ERROR_WORKSHEET_NOT_FOUND.format( contract_advance_id=contract_advance_id ), ) return obj @classmethod def filter_active_internal(cls): """Return select() statement with non-deleted internal items.""" return select(cls).where( cls.deleted_at.is_(None), cls.is_internal == True, # noqa ) @classmethod def default_order(cls): """Define default ordering.""" return cls.created_at.desc() @classmethod def filter_for(cls, query: filter_args) -> list: """ Create filters from dict of known parameters. Accepted dict keys: salesforce_id, contract_advance_id, payment_statuses, worksheet_ids, active """ filters = [] salesforce_id = query.get('salesforce_id') if salesforce_id: filters.append(cls.salesforce_id.ilike(salesforce_id)) contract_advance_id = query.get('contract_advance_id') if contract_advance_id: filters.append(cls.contract_advance_id == contract_advance_id) payment_statuses = query.get('payment_statuses') if payment_statuses: if None in payment_statuses: filters.append( or_( literal_column('action_status').in_(payment_statuses), literal_column('action_status').is_(None), ) ) else: filters.append(literal_column('action_status').in_(payment_statuses)) worksheet_ids = query.get('worksheet_ids') if worksheet_ids is not None: filters.append(cls.worksheet_payment_contract_advance_id.in_(worksheet_ids)) active = query.get('active') if active: filters.append(cls.deleted_at.is_(None)) return filters @classmethod def delete_by_id_or_error(cls, obj_id, error_status=404, **kwargs): """Soft delete by ID with extra status validation.""" obj = cls.get_by_id_or_error(obj_id, error_status) init_state_query = ( select(literal_column('1')) .where( and_( literal_column('parent_table_name') == cls.__tablename__, literal_column('parent_table_id') == obj_id, literal_column('action_name') == WORKSHEET_PAYMENT_CONTRACT_ADVANCE_ACTIONS.SEND_PAYMENTS, literal_column('action_status') == WORKSHEET_PAYMENT_CONTRACT_ADVANCE_STATUSES.INIT, ) ) .select_from(table('abacus_state')) ) init_status = db.session.execute(init_state_query).scalars().all() if not init_status: abort( code=400, description=error.ERROR_WORKSHEET_STATUS_REQUIRED.format( action=WORKSHEET_PAYMENT_CONTRACT_ADVANCE_ACTIONS.SEND_PAYMENTS, status=WORKSHEET_PAYMENT_CONTRACT_ADVANCE_STATUSES.INIT, ), ) obj._soft_delete() db.session.commit() @classmethod def get_filtered_query(cls, **kwargs: filter_args): """ Get worksheets by field values. Accepted dict keys: salesforce_id, contract_advance_id, payment_statuses, worksheet_ids, active """ stmt = select(cls) if 'payment_statuses' in kwargs: stmt = cls.join_states(stmt) stmt = stmt.where(*cls.filter_for(kwargs)) return stmt @classmethod def join_states(cls, stmt=None): """Join on abacus_state table.""" if stmt is None: stmt = select(cls) states_query = ( select(literal_column('action_status'), literal_column('parent_table_id')) .where( and_( literal_column('parent_table_name') == cls.__tablename__, literal_column('action_name') == WORKSHEET_PAYMENT_CONTRACT_ADVANCE_ACTIONS.SEND_PAYMENTS, ) ) .select_from(table('abacus_state')) .alias('states') ) return stmt.outerjoin( states_query, literal_column('parent_table_id') == cls.worksheet_payment_contract_advance_id, ) @classmethod def get_filtered_first(cls, **kwargs: filter_args): """Execute get_filtered_query and return first result.""" return db.session.execute(cls.get_filtered_query(**kwargs)).scalars().first() @classmethod def get_filtered_all(cls, **kwargs: filter_args) -> list: """Execute get_filtered_query and return all results.""" return db.session.execute(cls.get_filtered_query(**kwargs)).scalars().all() @classmethod def get_filtered_active_internal_records( cls, offset: int, limit: int, salesforce_id: str, contract_advance_id: int, payment_statuses: list[str], ) -> tuple[list['WorksheetPaymentContractAdvance'], int]: """Get active internal worksheets by field values.""" stmt = cls.filter_active_internal() if payment_statuses: stmt = cls.join_states(stmt) filter_args = { 'salesforce_id': salesforce_id, 'contract_advance_id': contract_advance_id, 'payment_statuses': payment_statuses, } stmt = stmt.where(*cls.filter_for(filter_args)) total_count = db.session.execute( select(func.count()).select_from(stmt.subquery()) ).scalar_one() stmt = stmt.order_by(WorksheetPaymentContractAdvance.default_order()) items = ( db.session.execute(stmt.limit(limit).offset(offset)).scalars().all() if limit > 0 else [] ) return items, total_count