"""Schema for tax form body filter parameters.""" from abc import ABC from abacus_common_logic.marshalling.custom_fields import ma from marshmallow import validates_schema, ValidationError from marshmallow_oneofschema import OneOfSchema from payee.constants.constants import TAX_FORM_TYPES class FilterParamsPostSchema(ma.Schema): """Post request params for entries filtering.""" expiration_date_start = ma.Date(required=False) expiration_date_end = ma.Date(required=False) is_active = ma.Boolean(required=False) account_payee_ids = ma.List(ma.NonNegativeInteger(), required=False) @validates_schema def validate_dates(self, data, **kwargs): if data.get('expiration_date_start') and data.get('expiration_date_end'): if data['expiration_date_start'] > data['expiration_date_end']: raise ValidationError( 'expiration_date_start must be before expiration_date_end', ) class TaxFormInfoOutputSchema(ma.Schema): items = ma.List(ma.Nested('TaxFormInfoSchema'), required=True) total_count = ma.NonNegativeInteger(required=True) class TaxFormInfoSchema(ma.Schema): """Schema for account payee tax form info output.""" account_payee_tax_form_info_id = ma.NonNegativeInteger( required=True, dump_only=True ) tax_form_type = ma.Enum(options=TAX_FORM_TYPES, required=True) account_payee_id = ma.IntegerId(required=True) signed_date = ma.Date(required=False, allow_none=True) expiration_date = ma.Date(required=False, allow_none=True) last_modified = ma.DateTime(required=False) @validates_schema def validate_dates(self, data, **kwargs): if data.get('expiration_date') and data.get('signed_date'): if data['expiration_date'] < data['signed_date']: raise ValidationError( 'expiration_date must be after signed_date', ) class BaseUSTaxFormSchema(TaxFormInfoSchema, ABC): """Base US tax form schema.""" _tax_form_type: str account_payee_tax_form_info_id = ma.NonNegativeInteger( required=True, dump_only=True ) tax_id_country = ma.String(required=True) tin_type = ma.String(required=True, allow_none=True) tin = ma.String(required=True, allow_none=True) tax_name = ma.String(required=True, allow_none=False) def _tax_form_type(self) -> str: pass class USTaxFormW9Schema(BaseUSTaxFormSchema): """W-9 tax form schema.""" _tax_form_type = TAX_FORM_TYPES.W9 tax_classification = ma.String(required=True, allow_none=False) class USTaxFormW8BENSchema(BaseUSTaxFormSchema): """W-8BEN tax form schema.""" _tax_form_type = TAX_FORM_TYPES.W8BEN tax_treaty_claim = ma.Boolean(required=True, allow_none=False) signed_date = ma.Date(required=True) expiration_date = ma.Date(required=False) class USTaxFormW8BENESchema(BaseUSTaxFormSchema): """W-8BEN-E tax form schema.""" _tax_form_type = TAX_FORM_TYPES.W8BENE tax_treaty_claim = ma.Boolean(required=True) type_of_entity = ma.String(required=True, allow_none=False) lob = ma.String(required=False, allow_none=True) signed_date = ma.Date(required=True) expiration_date = ma.Date(required=False) @validates_schema def validate_lob(self, data, **kwargs): if data.get('tax_treaty_claim') and not data.get('lob'): raise ValidationError('lob is required for tax_treaty_claim', 'lob') class USTaxFormW8IMYSchema(BaseUSTaxFormSchema): """W-8IMY tax form schema.""" _tax_form_type = TAX_FORM_TYPES.W8IMY tax_treaty_claim = ma.Boolean(required=True, allow_none=False) type_of_entity = ma.String(required=True, allow_none=False) signed_date = ma.Date(required=True) expiration_date = ma.Date(required=False) class USTaxFormW8ECISchema(BaseUSTaxFormSchema): """W-8ECI tax form schema.""" _tax_form_type = TAX_FORM_TYPES.W8ECI tax_treaty_claim = ma.Boolean(required=True, allow_none=False) type_of_entity = ma.String(required=True, allow_none=False) signed_date = ma.Date(required=True) expiration_date = ma.Date(required=False) class TaxFormInputSchema(OneOfSchema): """Union input schema for tax forms.""" type_field = 'tax_form_type' type_field_remove = False type_schemas = { cls._tax_form_type: cls( exclude=('account_payee_id', 'last_modified', 'expiration_date') ) for cls in BaseUSTaxFormSchema.__subclasses__() } class TaxFormOutputSchema(OneOfSchema): """Union output schema for tax forms.""" type_field = 'tax_form_type' type_schemas = { cls._tax_form_type: cls for cls in BaseUSTaxFormSchema.__subclasses__() } def get_obj_type(self, obj): try: return obj.tax_form_type except Exception: raise Exception('Unknown object type: %s' % repr(obj)) class TaxFormInfoDetailsOutputSchema(ma.Schema): items = ma.List(ma.Nested(TaxFormOutputSchema), required=True) total_count = ma.NonNegativeInteger(required=True)