"""AccountTaxInfo model.""" from datetime import date import typing from abacus_common_logic.connectors.database import db from abacus_common_logic.models.base import BaseModel from sqlalchemy import Enum from abacus_account.constants.constants import TAX_EMPLOYMENT_TYPES class AccountTaxInfo(BaseModel): """AccountTaxInfo model.""" __tablename__ = 'account_tax_info' account_tax_info_id = db.Column(db.Integer, primary_key=True) account_id = db.Column( db.Integer, db.ForeignKey('account.account_id'), nullable=False ) country_of_tax_residence = db.Column(db.String(3), nullable=True) is_sba_signed = db.Column(db.Boolean, default=False) is_vat_exempt = db.Column(db.Boolean, default=True, nullable=False) is_tax_treaty_claimed = db.Column(db.Boolean, default=False, nullable=False) tax_employment_type = db.Column( Enum(*TAX_EMPLOYMENT_TYPES, name='tax_employment_type', create_type=False), nullable=True, ) certificate_of_residence_expiration_date = db.Column(db.Date, nullable=True) is_wht_applicable = db.Column(db.Boolean, default=True, nullable=False) is_resident_of_spanish_islands = db.Column(db.Boolean, nullable=True) wht_rate_override = db.Column(db.Numeric(5, 2), nullable=True) account_tax_info_history = db.relationship( 'AccountTaxInfoHistory', backref='account_tax_info', cascade='all, delete-orphan' ) @classmethod def stream_all(cls, account_ids=None): """Stream all account tax info entries.""" query = cls.query if account_ids: query = query.filter(cls.account_id.in_(account_ids)) return query @classmethod def get_filtered_items( cls, limit: int, offset: int, account_ids: typing.Optional[typing.List[int]] = None, certificate_of_residence_expiration_date_start: date | None = None, certificate_of_residence_expiration_date_end: date | None = None ): """ Get filtered items. Arg: limit (int): pagination limit offset (int): pagination offset account_ids (list): optional accounts ids filter Returns: A tuple containing items and total count """ query = cls.query if account_ids: query = query.filter(cls.account_id.in_(account_ids)) if certificate_of_residence_expiration_date_start: query = query.filter( cls.certificate_of_residence_expiration_date.isnot(None), cls.certificate_of_residence_expiration_date >= certificate_of_residence_expiration_date_start # noqa: E501 ) if certificate_of_residence_expiration_date_end: query = query.filter( cls.certificate_of_residence_expiration_date.isnot(None), cls.certificate_of_residence_expiration_date <= certificate_of_residence_expiration_date_end # noqa: E501 ) total_count = query.count() items = query.limit(limit).offset(offset).all() if limit > 0 else [] return items, total_count