import logging from typing import Any, List from abacus_common_logic.connectors.database import db from payee.constants.error import ERROR_ACCOUNT_PAYEE_ID_NOT_FOUND from payee.models.account_payee import AccountPayee from payee.models.tax_withholding_override import TaxWithholdingOverride from payee.utils.exception import ( AccountPayeeNotFoundException, TaxWithholdingOverrideException, ) logger = logging.getLogger('tax_withholding_override') def get_tax_withholding_override_by_account_payee_ids( account_payee_ids: List[int], ) -> dict: """Fetch tax withholding overrides by account payee IDs.""" overrides = TaxWithholdingOverride.get_by_account_payee_ids(account_payee_ids) overrides_map = {override.account_payee_id: override for override in overrides} return { 'items': [ {'data': overrides_map.get(account_payee_id)} for account_payee_id in account_payee_ids ] } def get_tax_withholding_override_by_account_payee_id( account_payee_id: int, ) -> dict: """Fetch tax withholding override by account payee ID.""" override = TaxWithholdingOverride.get_by_account_payee_id(account_payee_id) return override def create_or_update_tax_withholding_override( account_payee_id: int, **created_params: Any ) -> dict: """Create tax withholding override for account payee.""" account_payee = AccountPayee.get_payee_by_id(account_payee_id) if account_payee is None: raise AccountPayeeNotFoundException( ERROR_ACCOUNT_PAYEE_ID_NOT_FOUND.format(account_payee_id=account_payee_id) ) try: tax_withhold_override = TaxWithholdingOverride.get_by_account_payee_id( account_payee_id ) if tax_withhold_override is None: tax_withhold_override = TaxWithholdingOverride.build( account_payee_id=account_payee_id, rate_override=created_params.get('rate_override', None), certificate_expiration_date=created_params.get( 'certificate_expiration_date', None ), message=created_params.get('message', None), ) else: tax_withhold_override.update_attributes(**created_params) db.session.commit() return tax_withhold_override except Exception as e: db.session.rollback() logger.error(f'Tax withholding override error: {str(e)}') raise TaxWithholdingOverrideException(str(e))