"""Ledger vat summary model.""" from sqlalchemy import Boolean from sqlalchemy import Column from sqlalchemy import DateTime from sqlalchemy import Enum from sqlalchemy import func from sqlalchemy import Integer from sqlalchemy import literal_column from sqlalchemy import Numeric from sqlalchemy import String from sqlalchemy import table from sqlalchemy import Text from sqlalchemy.engine import Row from moneyhub.constants.constants import VatCategory from moneyhub.models.abacus_event import AbacusEvent from moneyhub.models.mysql_base import BaseModel class LedgerVatSummary(BaseModel): """Ledger vat summary model.""" __tablename__ = 'ledger_vat_summary' ledger_vat_summary_id = Column(Integer, primary_key=True) statement_period_id = Column(Integer, nullable=False) activity_statement_period_id = Column(Integer, nullable=False) abacus_event_id = Column(Integer, nullable=False) account_id = Column(Integer, nullable=False) contract_id = Column(Integer, nullable=False) description = Column(Text, nullable=True) vat_category = Column( Enum( *VatCategory, name='vat_category', create_type=False ), nullable=False ) payee_currency_code = Column(String(3), nullable=False) vat_currency_code = Column(String(3), nullable=False) base_amount_payee_currency = Column(Numeric(20, 2), nullable=False) vat_rate = Column(Numeric(5, 2), nullable=True) vat_amount_payee_currency = Column(Numeric(20, 2), nullable=True) vat_amount_vat_currency = Column(Numeric(20, 2), nullable=True) wht_rate = Column(Numeric(20, 2), nullable=True) wht_amount_payee_currency = Column(Numeric(20, 2), nullable=True) wht_amount_vat_currency = Column(Numeric(20, 2), nullable=True) net_amount_payee_currency = Column(Numeric(20, 2), nullable=False) abacus_exempt_reason = Column(String(120), nullable=True) is_reporting_only = Column(Boolean, nullable=False) created_by = Column(String(255), nullable=False) created_at = Column(DateTime, nullable=False) last_modified_by = Column(String(255), nullable=False) last_modified = Column(DateTime, nullable=False) @classmethod def get_by_account_id( cls, account_id: int, contract_id: int | None, statement_period_id: int | None = None, activity_statement_period_id: int | None = None, vat_categories: list | None = None ) -> list: """Get a ledger vat summary by account_id and statement_period_id. Args: account_id (int): The id of an account contract_id (int): The id of a statement period statement_period_id (int): The id of the statement period vat_categories (list): vat categories to filter Returns: dict: dict of vat summary entities """ filters = [ cls.account_id == account_id ] if not statement_period_id and not activity_statement_period_id: raise Exception('Missing statement period ID') if statement_period_id: filters.append(cls.statement_period_id == statement_period_id) if activity_statement_period_id: filters.append(cls.activity_statement_period_id == activity_statement_period_id) if contract_id: filters.append(cls.contract_id == contract_id) if vat_categories: filters.append(cls.vat_category.in_(vat_categories)) return cls.query.filter(*filters).order_by( cls.ledger_vat_summary_id.desc() ).all() @classmethod def get_by_activity_statement_period(cls, statement_period_id: int) -> list: """Get entries by activity statement period. Args: statement_period_id (int): Statement period to get entries for. Returns: list: Entries made during a statement period. """ filters = [ cls.activity_statement_period_id == statement_period_id ] return cls.query \ .filter(*filters) \ .all() @classmethod def get_by_statement_period(cls, statement_period_id: int) -> list: """Get entries by statement period. Args: statement_period_id (int): Statement period to get entries for. Returns: list: Entries made during a statement period. """ filters = [ cls.statement_period_id == statement_period_id ] return cls.query \ .filter(*filters) \ .all() @classmethod def get_by_payment_entity_and_activity_statement_period( cls, payment_entity_id, statement_period_id ) -> Row: """Get entries by payment entity and activity statement period. Args: payment_entity_id (int): payment entity to get entries for. statement_period_id (int): statement period to get entries for. Returns: dict: vat entry """ filter_condition = \ [literal_column('account_payment_term.payment_entity_id') == payment_entity_id] return cls.query\ .join( table('account_payment_term'), literal_column('account_payment_term.account_id') == cls.account_id)\ .filter( *filter_condition, cls.activity_statement_period_id == statement_period_id)\ .first() @classmethod def get_by_vat_summary_file( cls, ledger_vat_summary_id: int, limit: int, offset: int ) -> tuple[list, int]: """Get the ledger VAT summary entries associated with a VAT summary file. Args: vat_summary_file_id (int): ID of the VAT summary file limit (int): how many entities to retrieve offset (int): the offset (for pagination) Returns: tuple: list of ledger VAT summary entries and the total amount of records """ query = cls.query \ .with_entities(cls, func.count().over().label('total_records')) \ .join(AbacusEvent, AbacusEvent.abacus_event_id == cls.abacus_event_id) \ .filter( AbacusEvent.target_type == 'vat_summary_file', AbacusEvent.target_id == ledger_vat_summary_id) if limit != 0: query = query.limit(limit) if offset != 0: query = query.offset(offset) rows = query.all() total_records = int(rows[0].total_records) if rows else 0 entries = [row[0] for row in rows] return (entries, total_records) @classmethod def get_visible_by_account_and_statement_periods( cls, account_id: int, contract_id: int | None, statement_period_ids: list[int], ) -> list: """Get only visible (not reporting-only) VAT summary entries. Args: account_id (int): Account to filter by contract_id (int | None): Contract to filter by statement_period_ids (int): Statement periods to filter by Returns: list: ledger VAT summary entries """ filters = [ cls.account_id == account_id, cls.statement_period_id.in_(statement_period_ids), cls.is_reporting_only == False, # noqa: E712 ] if contract_id: filters.append(cls.contract_id == contract_id) return cls.query.filter(*filters).order_by( cls.ledger_vat_summary_id.asc() ).all() @classmethod def get_activity_by_account_id( cls, account_id: int ) -> Row: """Get a ledger vat summary by account_id. Args: account_id (int): The id of an account Returns: Row: vat summary entity """ return cls.query.filter(cls.account_id == account_id).first()