"""Schema validation for vendor contract.""" from abacus_common_logic.marshalling.custom_fields import Range, ma from marshmallow import ValidationError, validates_schema from abacus_legacy_sync.constants.constants import ( VENDOR_CONTRACT_CONTRACT_TYPES, VENDOR_CONTRACT_PAY_AFTER_INTERVALS, VENDOR_CONTRACT_PAYMENT_INTERVALS, VENDOR_CONTRACT_YES_NO_ENUM, ) class VendorContractSchema(ma.Schema): """Shared vendor contract fields.""" vendor_id = ma.NonNegativeInteger(required=True) cont_start = ma.FormattedDate(allow_none=True) cont_end = ma.FormattedDate(allow_none=True) contract_type = ma.Enum(options=VENDOR_CONTRACT_CONTRACT_TYPES, required=False) release_term = ma.NonNegativeInteger(required=False) opt_out = ma.Enum(options=VENDOR_CONTRACT_YES_NO_ENUM, required=False) is_automatic_rollover = ma.Enum(options=VENDOR_CONTRACT_YES_NO_ENUM, required=False) payment_interval = ma.Enum( options=VENDOR_CONTRACT_PAYMENT_INTERVALS, required=False, allow_none=True ) pay_after = ma.Enum( options=VENDOR_CONTRACT_PAY_AFTER_INTERVALS, required=False, allow_none=True ) digital_split = ma.Float(validate=Range(min=0.0, max=1.0), required=False) class VendorContractPostSchema(VendorContractSchema): """Schema for POST request.""" vendor_id = ma.Int(validate=Range(min=0, max=10**11), required=True) currency_code = ma.NonemptyString(required=True) country_exclusion = ma.List(ma.NonemptyString()) store_exclusion = ma.List(ma.NonemptyString()) distribution_type_id = ma.NonNegativeInteger(required=False) @validates_schema def validate_distribution_type_id(self, data, **kwargs): """Validate distribution_type_id field.""" store_exclusion = data.get('store_exclusion', None) if ( store_exclusion is not None and len(store_exclusion) > 0 and 'distribution_type_id' not in data ): raise ValidationError( 'distribution_type_id is required when store_exclusion is present.' ) class VendorContractDetailSchema(VendorContractSchema): """Schema for Detail response.""" vendor_contract_id = ma.IntegerId(required=True) currency_id = ma.IntegerId(required=True) territory_carve_out = ma.NonemptyString(required=False) class VendorContractPutSchema(ma.Schema): """Schema for PUT request.""" country_exclusion = ma.List(ma.NonemptyString()) store_exclusion = ma.List(ma.NonemptyString()) distribution_type_id = ma.NonNegativeInteger(required=False) @validates_schema def validate_distribution_type_id(self, data, **kwargs): """Validate distribution_type_id field.""" store_exclusion = data.get('store_exclusion', None) if ( store_exclusion is not None and len(store_exclusion) > 0 and 'distribution_type_id' not in data ): raise ValidationError( 'distribution_type_id is required when store_exclusion is present.' )