from decimal import Decimal

from django.db import transaction
from django.db.models import Sum
from django.utils import timezone
from django.utils.dateparse import parse_datetime
from rest_framework.decorators import action
from rest_framework.response import Response

from api.models.data.loan import Loan, LoanPayment
from api.serializers.data.loan import LoanPaymentSerializer, LoanSerializer
from api.views.data.base import DataRootViewSet


class LoanPaymentViewSet(DataRootViewSet):
    permission_module = 'loans'
    serializer_class = LoanPaymentSerializer
    filterset_fields = ['loan']

    def get_queryset(self):
        qs = super().get_queryset()
        return qs


class LoanViewSet(DataRootViewSet):
    permission_module = 'loans'
    queryset = Loan.objects.select_related().prefetch_related('payments').all().order_by('-loan_date')
    serializer_class = LoanSerializer
    search_fields = ['customer__name', 'customer__phone', 'customer__email', 'vendor__name', 'bill_number', 'notes']

    @action(detail=False, methods=['get'])
    def summary(self, request):
        queryset = self.get_queryset()

        loan_in_total = queryset.filter(loan_type='loan_in').aggregate(total=Sum('amount'))['total'] or 0
        loan_out_total = queryset.filter(loan_type='loan_out').aggregate(total=Sum('amount'))['total'] or 0

        unpaid_loan_in = Decimal('0')
        unpaid_loan_out = Decimal('0')
        for loan in queryset.filter(is_paid=False):
            due = loan.balance_due
            if loan.loan_type == 'loan_in':
                unpaid_loan_in += due
            else:
                unpaid_loan_out += due

        summary = {
            'total_loan_in': loan_in_total,
            'total_loan_out': loan_out_total,
            'unpaid_loan_in': float(unpaid_loan_in),
            'unpaid_loan_out': float(unpaid_loan_out),
            'net_balance': loan_in_total - loan_out_total,
            'total_loans': queryset.count(),
            'unpaid_loans': queryset.filter(is_paid=False).count(),
        }

        return Response(summary)

    @action(detail=True, methods=['post'])
    @transaction.atomic
    def add_payment(self, request, pk=None):
        loan = self.get_object()
        amount = Decimal(str(request.data.get('amount', 0)))
        if amount <= 0:
            return Response({'amount': 'Payment amount must be greater than zero.'}, status=400)

        balance = loan.balance_due
        if amount > balance:
            return Response(
                {'amount': f'Payment cannot exceed remaining balance ({balance}).'},
                status=400,
            )

        payment_date = request.data.get('payment_date') or timezone.now()
        if isinstance(payment_date, str):
            payment_date = parse_datetime(payment_date) or timezone.now()

        method = request.data.get('payment_method') or loan.payment_method or 'cash'
        if method not in ('cash', 'bank', 'sarafi'):
            return Response({'payment_method': 'Must be cash, bank, or sarafi.'}, status=400)

        payment = LoanPayment.objects.create(
            loan=loan,
            amount=amount,
            payment_date=payment_date,
            payment_method=method,
            reference_number=request.data.get('reference_number', ''),
            notes=request.data.get('notes', ''),
        )
        return Response(LoanPaymentSerializer(payment).data)

    @action(detail=True, methods=['post'])
    @transaction.atomic
    def mark_paid(self, request, pk=None):
        loan = self.get_object()
        remaining = loan.balance_due
        if remaining > 0:
            LoanPayment.objects.create(
                loan=loan,
                amount=remaining,
                payment_method=request.data.get('payment_method') or loan.payment_method or 'cash',
                notes='Marked as fully paid',
            )
        else:
            loan.is_paid = True
            loan.save(update_fields=['is_paid', 'updated_at'])
        return Response(LoanSerializer(loan).data)

    @action(detail=True, methods=['post'])
    @transaction.atomic
    def mark_unpaid(self, request, pk=None):
        loan = self.get_object()
        loan.payments.all().delete()
        loan.amount_paid = Decimal('0')
        loan.is_paid = False
        loan.save(update_fields=['amount_paid', 'is_paid', 'updated_at'])
        return Response({'message': 'Loan marked as unpaid'})
