"""Contract marshmallow schemas.""" from datetime import datetime from typing import TypedDict from abacus_common_logic.marshalling.custom_fields import ma from marshmallow import fields, validates_schema, ValidationError from abacus_contract.constants import constants, error from abacus_contract.schemas.contract_lifecycle import ContractLifecyclePostSchema from abacus_contract.schemas.contract_lifecycle_schedule import \ ContractLifecycleSchedulePostSchema class BaseContractSchema(ma.Schema): """Shared contract fields.""" contract_type = ma.Enum(options=constants.CONTRACT_TYPES, required=True) execution_date = ma.FormattedDate(allow_none=True) reference_signing_entity_id = ma.NonNegativeInteger(required=False) term_end = ma.FormattedDate(required=False, allow_none=True) term_start = ma.FormattedDate(required=False, allow_none=True) summary_note = ma.String(required=False, allow_none=True) general_note = ma.String(required=False, allow_none=True) class ContractDetailSchema(BaseContractSchema): """Contract GET response.""" contract_id = ma.IntegerId(required=True) contract_name = ma.NonemptyString(required=True) oa_contract_id = ma.Pluck( 'LegacyContractSchema', 'oa_contract_id', attribute='legacy_contract', allow_none=True ) account_id = ma.NonNegativeInteger(required=True) sap_created_at = ma.FormattedDateTime(allow_none=True) initial_start_date = ma.FormattedDate(allow_none=True) is_excluded_from_accounting_run = ma.Boolean() class ContractPostSchema(BaseContractSchema): """Contract POST body.""" contract_id = ma.IntegerId(allow_none=True, required=False) account_id = ma.IntegerId() contract_name = ma.NonemptyString(required=True) oa_contract_id = ma.IntegerId() is_excluded_from_accounting_run = ma.Boolean(allow_none=True) class ContractPutSchema(ma.Schema): """Contract PUT body.""" contract_name = ma.NonemptyString() reference_signing_entity_id = ma.NonNegativeInteger(required=False) sap_created_at = ma.FormattedDateTime(allow_none=True) summary_note = ma.String(required=False, allow_none=True) general_note = ma.String(required=False, allow_none=True) term_end = ma.FormattedDate(allow_none=True) term_start = ma.FormattedDate(allow_none=True) is_excluded_from_accounting_run = ma.Boolean() execution_date = ma.FormattedDate(allow_none=True) initial_start_date = ma.FormattedDate(allow_none=True) class ContractTerminationSchema(ma.Schema): """Contract termination body.""" termination_effective = ma.FormattedDate(required=True) termination_notice_received = ma.FormattedDate(required=False, load_default=None) class ContractVatInfoSchema(ma.Schema): """Contract Vat Info list response.""" account_id = ma.IntegerId() contract_id = ma.IntegerId() country_of_tax_residence = ma.String(allow_none=True) account_is_sba_signed = ma.Boolean(allow_none=True) client_tax_rate = ma.Decimal(as_string=True, allow_none=True) supplier_tax_rate = ma.Decimal(as_string=True, allow_none=True) class ContractSapFormattedSchema(ma.Schema): """Contract info for SAP API response.""" account_id = ma.NonemptyString(data_key='AccountId', required=True) contract_id = ma.NonemptyString(data_key='ContractId', required=True) contract_name = ma.TruncatedString( data_key='ContractName', required=True, metadata={'truncate': 120}) contract_type = ma.Enum( options=constants.CONTRACT_TYPES, required=True, data_key='ContractType' ) term_end = ma.FormattedDateTime(data_key='DateTo') term_start = ma.FormattedDateTime(data_key='DateFrm') Bukrs = ma.String(allow_none=False) Prctr = ma.String(allow_none=False) BusUnit = ma.String(allow_none=True, default=None) Zzfield1 = ma.String(allow_none=True, default=None) Zzfield2 = ma.String(allow_none=True, default=None) Zzfield3 = ma.String(allow_none=True, default=None) class ContractAndLifecyclePostSchema(ma.Schema): """Schema for contract, contract_lifecycle and contract_lifecycle_schedule POST request.""" # noqa: E501 contract = fields.Nested(ContractPostSchema, required=True) contract_lifecycle = fields.Nested(ContractLifecyclePostSchema, required=True) contract_lifecycle_schedules = fields.List( fields.Nested(ContractLifecycleSchedulePostSchema, required=True), required=True ) @validates_schema def validate_contract_lifecycle_schedules(self, data, **kwargs): """Check whether contract_lifecycle_schedules field is empty.""" contract_lifecycle_schedules = data.get('contract_lifecycle_schedules') if len(contract_lifecycle_schedules) == 0: raise ValidationError( error.ERROR_CONTRACT_LIFECYCLE_SCHEDULES_LIST_EMPTY ) class ContractDetail(TypedDict): """TypedDict that mirrors the serialized output from ContractDetailSchema.""" contract_id: int contract_name: str oa_contract_id: int | None account_id: int sap_created_at: datetime | None initial_start_date: datetime | None is_excluded_from_accounting_run: bool