"""Workstation summary model.""" from sqlalchemy import asc from sqlalchemy import Column from sqlalchemy import desc from sqlalchemy import func from sqlalchemy import Integer from sqlalchemy import Numeric from sqlalchemy import String from moneyhub.constants.constants import OrderDirection from moneyhub.models.snowflake_base import BaseModel class WorkstationSummary(BaseModel): """Workstation summary model.""" __tablename__ = 'workstation_summary_dbt' account_id = Column(Integer, nullable=False, primary_key=True) account_name = Column(String(100), nullable=False) contract_id = Column(Integer, nullable=False, primary_key=True) statement_period_id = Column(Integer, nullable=False, primary_key=True) currency = Column(String(3), nullable=False) opening_balance = Column(Numeric(32, 2), nullable=False) payments = Column(Numeric(32, 2), nullable=True) net_revenue = Column(Numeric(32, 2), nullable=True) gross_revenue = Column(Numeric(32, 2), nullable=True) fee = Column(Numeric(32, 2), nullable=True) expenses = Column(Numeric(32, 2), nullable=True) adjustments = Column(Numeric(32, 2), nullable=True) closing_balance = Column(Numeric(32, 2), nullable=False) mechanicals = Column(Numeric(32, 2), nullable=True) mechanical_fees = Column(Numeric(32, 2), nullable=True) @classmethod def get_balances_by_account_id( cls, account_id: int, statemement_period_ids: list, contract_id: int | None ) -> list: """Get account balances by statement periods. Args: account_id (int): Account to get periods for. contract_id (int): Optional contract to get periods for. Returns: list: Statement periods. """ filters = [ cls.account_id == account_id, cls.statement_period_id.in_(statemement_period_ids), ] if contract_id: filters.append(cls.contract_id == contract_id) return cls.query \ .with_entities( cls.statement_period_id, cls.opening_balance, cls.closing_balance, cls.currency )\ .filter(*filters).order_by(asc(cls.statement_period_id)).all() @classmethod def get_revenue_for_account( cls, account_id: int, contract_id: int | None, statemement_period_ids: list | None, order_dir: OrderDirection = OrderDirection.ASC ) -> list: """Get revenue for a given account. Args: account_id (int): Account to get revenue for. contract_id (int): Optional contract to filter for statemement_period_ids (list): Optional list of periods to filter for order_dir (OrderDirection): Direction to order the results Returns: list: Statement periods. """ entities = [ cls.statement_period_id, func.max(cls.currency).label('currency'), func.sum(cls.gross_revenue).label('gross_revenue'), func.sum(cls.net_revenue).label('net_revenue'), func.sum(cls.fee).label('fee'), func.sum(cls.mechanicals).label('mechanicals'), func.sum(cls.mechanical_fees).label('mechanical_fees'), ] filters = [ cls.account_id == account_id, ] grouping = [ cls.account_id, cls.statement_period_id, ] if statemement_period_ids: filters.append(cls.statement_period_id.in_(statemement_period_ids)) if contract_id: filters.append(cls.contract_id == contract_id) order_direction = asc if order_dir == OrderDirection.ASC else desc return cls.query \ .with_entities(*entities)\ .filter(*filters)\ .group_by(*grouping)\ .order_by(order_direction(cls.statement_period_id))\ .all()