from typing import List from abacus_common_logic.connectors.database import db from sqlalchemy import asc, desc, func, literal_column, or_, select, table from payment.models import WorksheetPayableBalanceAfterTax def get_filtered_active_records( event_id: int, statement_period_id: int, contract_ids: List[int] | None, limit: int, offset: int, search_term: str | None = None, sort_by: str | None = None, sort_order: str | None = None, ) -> tuple[list[WorksheetPayableBalanceAfterTax], int]: """Get filtered active non deleted records.""" query = WorksheetPayableBalanceAfterTax.filter_active() if statement_period_id: query = query.where( WorksheetPayableBalanceAfterTax.statement_period_id == statement_period_id ) if event_id: query = query.where(WorksheetPayableBalanceAfterTax.abacus_event_id == event_id) if contract_ids: query = query.where( WorksheetPayableBalanceAfterTax.contract_id.in_(contract_ids) ) order_func = desc if sort_order and sort_order.lower() == 'desc' else asc if search_term or sort_by == 'account_name': query = query.join( table('account').alias('account'), literal_column('account.account_id') == WorksheetPayableBalanceAfterTax.account_id, ) if search_term: query = query.where( or_( literal_column('account.account_name').ilike(f'%{search_term}%'), WorksheetPayableBalanceAfterTax.account_id.ilike( # for consistency with ows-abacus-account search f'%{search_term}%' ), ) ) if sort_by == 'account_name': query = query.order_by(order_func(literal_column('account.account_name'))) elif sort_by == 'contract_name': query = query.join( table('contract').alias('contract'), literal_column('contract.contract_id') == WorksheetPayableBalanceAfterTax.contract_id, ).order_by(order_func(literal_column('contract.contract_name'))) if sort_by in WorksheetPayableBalanceAfterTax.__table__.columns.keys(): query = query.order_by(order_func(literal_column(sort_by))) total_count = db.session.execute( select(func.count()).select_from(query.subquery()) ).scalar_one() if limit: query = query.limit(limit) if offset: query = query.offset(offset) return db.session.execute(query).scalars().all(), total_count