from decimal import Decimal, InvalidOperation

from rest_framework import serializers

from .models import Employee, PayrollRates, Payslip, PayrollRun

RUN_FIELDS = (
    'id', 'period_year', 'period_month', 'status', 'employee_count',
    'total_gross', 'total_deductions', 'total_net', 'created_at',
)


class EmployeeSerializer(serializers.ModelSerializer):
    gross_pay = serializers.DecimalField(max_digits=14, decimal_places=2, read_only=True)

    class Meta:
        model = Employee
        fields = '__all__'
        read_only_fields = ('id', 'company')


class PayslipSerializer(serializers.ModelSerializer):
    # Payslips are only ever generated by a payroll run, never written directly.
    class Meta:
        model = Payslip
        fields = '__all__'


class PayrollRunSerializer(serializers.ModelSerializer):
    class Meta:
        model = PayrollRun
        fields = RUN_FIELDS
        read_only_fields = RUN_FIELDS


class PayrollRunDetailSerializer(serializers.ModelSerializer):
    payslips = PayslipSerializer(many=True, read_only=True)

    class Meta:
        model = PayrollRun
        fields = RUN_FIELDS + ('payslips', 'rates_snapshot')
        read_only_fields = RUN_FIELDS


class PayrollRatesSerializer(serializers.ModelSerializer):
    class Meta:
        model = PayrollRates
        fields = (
            'paye_bands', 'personal_relief',
            'nssf_rate', 'nssf_tier_1_limit', 'nssf_tier_2_limit',
            'shif_rate', 'shif_minimum', 'housing_levy_rate',
            'nssf_deductible_for_paye', 'shif_deductible_for_paye',
            'housing_levy_deductible_for_paye',
            'updated_at',
        )
        read_only_fields = ('updated_at',)

    def validate_paye_bands(self, bands):
        if not isinstance(bands, list) or not bands:
            raise serializers.ValidationError('Add at least one PAYE band.')

        previous_limit = None
        for index, band in enumerate(bands):
            if not isinstance(band, dict) or 'rate' not in band:
                raise serializers.ValidationError(f'Band {index + 1} is missing a rate.')
            try:
                rate = Decimal(str(band['rate']))
            except (InvalidOperation, TypeError):
                raise serializers.ValidationError(f'Band {index + 1} has an invalid rate.')
            if not 0 <= rate <= 1:
                raise serializers.ValidationError(
                    f'Band {index + 1}: rate must be between 0% and 100%.'
                )

            upper = band.get('up_to')
            is_last = index == len(bands) - 1
            if is_last:
                if upper not in (None, ''):
                    raise serializers.ValidationError('The final band must have no upper limit.')
                continue
            if upper in (None, ''):
                raise serializers.ValidationError(
                    f'Band {index + 1} needs an upper limit — only the final band is open-ended.'
                )
            try:
                limit = Decimal(str(upper))
            except (InvalidOperation, TypeError):
                raise serializers.ValidationError(f'Band {index + 1} has an invalid upper limit.')
            if limit <= 0:
                raise serializers.ValidationError(f'Band {index + 1}: upper limit must be positive.')
            if previous_limit is not None and limit <= previous_limit:
                raise serializers.ValidationError('PAYE band limits must increase from one band to the next.')
            previous_limit = limit

        return bands

    def validate(self, attrs):
        tier_1 = attrs.get('nssf_tier_1_limit', getattr(self.instance, 'nssf_tier_1_limit', None))
        tier_2 = attrs.get('nssf_tier_2_limit', getattr(self.instance, 'nssf_tier_2_limit', None))
        if tier_1 is not None and tier_2 is not None and tier_2 < tier_1:
            raise serializers.ValidationError(
                {'nssf_tier_2_limit': 'The upper NSSF limit must be at or above the lower limit.'}
            )
        return attrs
