"""Schema for payment_group_payment_account serialization.""" import json from abacus_common_data.currency import get_currency_object_from_code from abacus_common_logic.marshalling.base import PaginationSchema from abacus_common_logic.marshalling.custom_fields import ma from marshmallow import ( fields, pre_load, validate, validates, validates_schema, ValidationError, ) from payment.constants.constants import ( DEFAULT_PAGE_LIMIT, DEFAULT_SORT_BY, DEFAULT_SORT_ORDER, PAYMENT_ACCOUNT_SORTABLE_COLUMNS, ) from payment.constants.error import ( ERROR_DUPLICATE_ACCOUNT_ID_BULK_FILTER, ERROR_INVALID_SORT_BY, ERROR_INVALID_SORT_ORDER, ) class ContractsPayableSchema(ma.Schema): """Contract payable balance information. To be saved in a JSON DB column. NOTE: current_balance has to be a string because a decimal is not a valid JSON type. """ def dump(self, data, **kwargs): """Load data from string to json.""" new_data = data if isinstance(data, str): new_data = json.loads(data) return super().dump(new_data, **kwargs) contract_id = ma.IntegerId(required=True) current_balance = ma.String(required=True) currency_code = ma.NonemptyString(required=True) class PaymentGroupPaymentAccountSchema(ma.Schema): """Payment Group Payment Account detail schema.""" account_id = ma.NonNegativeInteger(required=True) contracts_payable = ma.Nested(ContractsPayableSchema(many=True), required=True) currency_code = ma.NonemptyString(required=True) currency_name = ma.NonemptyString() current_balance = ma.Decimal(as_string=True, required=True) tax_withholding = ma.Decimal(as_string=True, allow_none=True) vat_amount = ma.Decimal(as_string=True, allow_none=True) balance_after_tax = ma.Decimal(as_string=True, required=True) last_payment = ma.Decimal(as_string=True, required=True) note = ma.String() payment_group_payment_account_id = ma.NonNegativeInteger(required=True) payment_group_payment_id = ma.NonNegativeInteger(required=True) payoneer_program_id = ma.NonNegativeInteger(required=True) prior_payment_group_payment_id = ma.NonNegativeInteger() last_statement_period_id = ma.NonNegativeInteger(allow_none=True) last_statement_period_name = ma.String() current_statement_period_id = ma.NonNegativeInteger(required=True) current_statement_period_name = ma.NonemptyString() class PaymentGroupPaymentAccountListSchema(PaymentGroupPaymentAccountSchema): """Payment Group Payment Account list schema.""" account_name = ma.NonemptyString() account_payee_id = ma.NonNegativeInteger() currency_name = fields.Method('currency_name') payment_difference = ma.Decimal(as_string=True) percent_difference = ma.Decimal(as_string=True, places=2) payoneer_payee_id = ma.NonNegativeInteger() payment_status = ma.NonemptyString() payment_error_code = ma.NonemptyString() @staticmethod def currency_name(obj): # noqa """Get currency name.""" if obj.get('currency_code'): return get_currency_object_from_code(obj['currency_code'])['currency_name'] class PaymentGroupPaymentAccountListItemsSchema(ma.Schema): """Schema for PaymentGroupPaymentAccountList response.""" items = ma.Nested(PaymentGroupPaymentAccountListSchema, many=True) total_count = ma.NonNegativeInteger(required=True) class PaymentGroupPaymentAccountPutSchema(ma.Schema): """Payment Group Payment Account put schema.""" note = ma.NonemptyString(required=True) class PaymentGroupPaymentAccountListResponseSchema(ma.Schema): """Schema for PaymentGroupPaymentAccountList response list.""" items = ma.Nested(PaymentGroupPaymentAccountSchema, many=True) total_count = ma.NonNegativeInteger(required=True) class PaymentGroupPaymentAccountFilterKeySchema(ma.Schema): """Filter key for filtering.""" account_ids = ma.List(ma.NonNegativeInteger(), required=False) contract_ids = ma.List(ma.NonNegativeInteger(), required=False) payment_statuses = ma.List(ma.NonemptyString(), required=False) @pre_load def split_comma_separated(self, data: dict, **kwargs): data = data.to_dict() if hasattr(data, 'to_dict') else data for field, value in data.items(): if not isinstance(value, str): continue data[field] = [p.strip() for p in value.split(',') if p.strip()] return data class PaymentGroupPaymentAccountFilterParamsPostSchema(ma.Schema): """Post request params for entries filtering.""" filters = ma.Nested( PaymentGroupPaymentAccountFilterKeySchema, required=False, load_default={}, ) class LastPaymentBulkFilterKeySchema(ma.Schema): """Filter by account IDs only.""" account_ids = ma.List( ma.NonNegativeInteger(), required=True, validate=[validate.Length(min=1)] ) @validates_schema def validate_account_ids(self, data, **kwargs): """Validate duplicates for account ids.""" account_ids = data.get('account_ids') if len(set(account_ids)) != len(account_ids): raise ValidationError(ERROR_DUPLICATE_ACCOUNT_ID_BULK_FILTER) class LastPaymentBulkFilterParamsPostSchema(ma.Schema): """Post request params for entries filtering.""" filters = ma.Nested( LastPaymentBulkFilterKeySchema, required=True, ) class PaymentGroupPaymentAccountPaginationSchema(PaginationSchema): limit = ma.Integer( load_default=DEFAULT_PAGE_LIMIT, required=False, validate=[validate.Range(min=1, max=5000)], ) sort_by = ma.String(load_default=DEFAULT_SORT_BY) sort_order = ma.String(load_default=DEFAULT_SORT_ORDER) search_term = ma.String(required=False, allow_none=True) @validates('sort_by') def validate_sort_by(self, value): if value.lower() not in PAYMENT_ACCOUNT_SORTABLE_COLUMNS: raise ValidationError(ERROR_INVALID_SORT_BY.format(sort_by=value)) @validates('sort_order') def validate_sort_order(self, value): if value.lower() not in ['asc', 'desc']: raise ValidationError(ERROR_INVALID_SORT_ORDER)