from abacus_common_logic.marshalling.custom_fields import ma from marshmallow import fields, pre_load, validate, validates_schema, ValidationError from payee.constants import error class TaxWithholdingOverrideDataloaderSchema(ma.Schema): """Schema that validates a list of account payee IDs directly.""" account_payee_ids = ma.List(ma.IntegerId(required=True, strict=True), required=True) @pre_load() def pre_load(self, data, *args, **kwargs): """Process top level list.""" return {'account_payee_ids': data} @validates_schema def check_duplicate(self, data, **kwargs): input_list = data.get('account_payee_ids', []) if len(set(input_list)) != len(input_list): raise ValidationError( error.ERROR_TAX_WITHHOLDING_OVERRIDE_DATALOADER_SCHEMA_DUPLICATE ) class TaxWithholdingOverrideCreateSchema(ma.Schema): rate_override = ma.Decimal( as_string=True, required=False, validate=validate.Range(min=0, max=100) ) certificate_expiration_date = ma.Date(required=False, allow_none=True) message = ma.String( required=False, validate=validate.Length(min=1), allow_none=True ) @validates_schema def check_empty_input(self, data, **kwargs): if not data: raise ValidationError( error.ERROR_TAX_WITHHOLDING_OVERRIDE_CREATE_SCHEMA_EMPTY_INPUT ) class TaxWithholdingOverrideSchema(ma.Schema): tax_withholding_override_id = ma.IntegerId(required=True) account_payee_id = ma.IntegerId(required=True) rate_override = ma.Decimal(as_string=True) certificate_expiration_date = ma.Date() message = ma.String() created_by = ma.String() created_at = ma.DateTime() last_modified_by = ma.String() last_modified = ma.DateTime() class DataloaderEntrySchema(ma.Schema): data = ma.Nested(TaxWithholdingOverrideSchema, required=False, allow_none=True) class TaxWithholdingOverrideDataloaderOutputSchema(ma.Schema): """Tax withholding override dataloader output schema.""" items = ma.Nested(DataloaderEntrySchema, many=True)