from decimal import Decimal

from django.db import transaction
from rest_framework import status
from rest_framework.exceptions import ValidationError
from rest_framework.response import Response

from api.models.data.sales import Sales, SalesItems
from api.serializers.data.sales import SalesSerializer, SalesItemsSerializer, SalesListSerializer
from api.models.data.currency import BASE_CURRENCY_ID
from api.utils.currency import convert_currency
from api.views.data.base import DataRootViewSet
from api.views.mixins.order_stock import OrderStockViewMixin
from django.db.models import F, Prefetch
from api.services.fefo_allocation import (
    allocate_inventory_for_sale_item,
    revert_sale_item_allocations,
    update_sale_item_allocations,
    InsufficientInventoryError,
)
from api.services.tank_inventory import validate_tank_fill_quantity


class SalesViewSet(OrderStockViewMixin, DataRootViewSet):
    permission_module = 'sales'
    queryset = Sales.objects.select_related('customer', 'salesman').all().order_by("-sale_date")
    serializer_class = SalesSerializer
    filterset_fields = ["customer", "salesman", "currency", "status"]
    search_fields = [
        "customer__name",
        "customer__phone",
        "customer__email",
        "salesman__name",
        "invoice_number",
        "notes",
    ]

    def get_serializer_class(self):
        if self.action == 'list':
            return SalesListSerializer
        return SalesSerializer

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

        outstanding = self.request.query_params.get('outstanding')
        if outstanding in ('1', 'true', 'True'):
            qs = qs.filter(total_amount__gt=F('paid_amount')).exclude(status='cancelled')

        if self.action == 'retrieve':
            qs = qs.prefetch_related(
                'documents',
                Prefetch(
                    'sales_items',
                    queryset=SalesItems.objects.select_related(
                        'stock',
                        'stock__product',
                        'storage_tank',
                    ),
                )
            )
        return qs

    def perform_create(self, serializer):
        serializer.save()

    def perform_update(self, serializer):
        serializer.save()

    def _ensure_commission_has_salesman(self, sale, items_data=None):
        """Line commissions require a salesman on the sale header."""
        has_commission = False
        if items_data is not None:
            for item_data in items_data or []:
                try:
                    amount = Decimal(str(item_data.get('commission_amount') or 0))
                except Exception:
                    amount = Decimal('0')
                if amount > 0:
                    has_commission = True
                    break
        else:
            has_commission = sale.sales_items.filter(commission_amount__gt=0).exists()

        if has_commission and not sale.salesman_id:
            raise ValidationError({
                'salesman': 'Select a salesman when items have commission.',
            })

    def _save_stock(self, stock):
        """Save stock without triggering manual adjustment journal (sale/purchase handles GL)."""
        stock._skip_stock_journal = True
        stock.save()

    def _snapshot_unit_cost(self, item):
        """Store inventory cost in AFN (per piece for stock, per liter for tank)."""
        if Decimal(str(item.unit_cost or 0)) > 0:
            return
        if item.sale_source == SalesItems.SOURCE_TANK and item.storage_tank_id:
            _, unit_cost = validate_tank_fill_quantity(
                item.storage_tank,
                item.quantity,
                exclude_sale_item_id=item.pk,
            )
            item.unit_cost = unit_cost
            item.save(update_fields=['unit_cost'])
            return
        if item.stock_id:
            stock = item.stock
            sale = item.sales
            item.unit_cost = convert_currency(
                stock.buy_price,
                stock.currency,
                BASE_CURRENCY_ID,
                sale.sale_date,
            )
            item.save(update_fields=['unit_cost'])

    def create(self, request, *args, **kwargs):
        with transaction.atomic():
            data = request.data.copy() if hasattr(request.data, 'copy') else dict(request.data)
            sales_items_data = data.pop('sales_items', None)
            if 'invoice_number' in data:
                data.pop('invoice_number')

            serializer = self.get_serializer(data=data)
            serializer.is_valid(raise_exception=True)
            self.perform_create(serializer)
            sale = serializer.instance
            affects_stock = self._order_affects_stock(sale)

            for item_data in sales_items_data or []:
                item_data['sales'] = sale.id
                item_serializer = SalesItemsSerializer(data=item_data, context=self.get_serializer_context())
                item_serializer.is_valid(raise_exception=True)
                item = item_serializer.save()
                self._snapshot_unit_cost(item)
                if affects_stock:
                    self._update_stock_for_sales_item(item, sale)

            self._ensure_commission_has_salesman(sale, sales_items_data)

            from api.services.accounting import recalculate_sale_totals, sync_sales_journal
            recalculate_sale_totals(sale)
            sync_sales_journal(sale)
            sale.refresh_from_db()

            response_data = SalesSerializer(sale, context=self.get_serializer_context()).data
            response_data['sales_items'] = SalesItemsSerializer(
                SalesItems.objects.filter(sales=sale), many=True, context=self.get_serializer_context()
            ).data
            return Response(response_data, status=status.HTTP_201_CREATED)

    def update(self, request, *args, **kwargs):
        with transaction.atomic():
            sale = self.get_object()
            old_status = sale.status

            data = request.data.copy() if hasattr(request.data, 'copy') else dict(request.data)
            items_provided = 'sales_items' in data
            sales_items_data = data.pop('sales_items', []) if items_provided else None

            existing_items = {item.id: item for item in sale.sales_items.all()}

            serializer = self.get_serializer(sale, data=data, partial=kwargs.get('partial', False))
            serializer.is_valid(raise_exception=True)
            serializer.save()
            sale.refresh_from_db()

            affects_stock = self._order_affects_stock(sale)

            if old_status != sale.status:
                self._apply_status_stock_transition(
                    sale,
                    old_status,
                    sale.sales_items.all(),
                    apply_item=self._update_stock_for_sales_item,
                    revert_item=self._increase_stock_for_sales_item,
                )

            if items_provided and affects_stock:
                processed_item_ids = set()
                for item_data in sales_items_data:
                    item_id = item_data.get('id')

                    if item_id and item_id in existing_items:
                        item = existing_items[item_id]
                        from api.utils.packaging import sales_item_pieces
                        old_piece_amount = sales_item_pieces(item)
                        old_stock = item.stock
                        old_tank = item.storage_tank
                        old_source = item.sale_source

                        item_serializer = SalesItemsSerializer(
                            instance=item, data=item_data, partial=True, context=self.get_serializer_context()
                        )
                        item_serializer.is_valid(raise_exception=True)
                        updated_item = item_serializer.save()
                        self._snapshot_unit_cost(updated_item)
                        self._update_stock_for_sales_item_update(
                            updated_item,
                            sale,
                            old_piece_amount,
                            old_stock=old_stock,
                            old_tank=old_tank,
                            old_source=old_source,
                        )
                        processed_item_ids.add(item_id)
                    else:
                        item_data['sales'] = sale.id
                        item_serializer = SalesItemsSerializer(data=item_data, context=self.get_serializer_context())
                        item_serializer.is_valid(raise_exception=True)
                        new_item = item_serializer.save()
                        self._snapshot_unit_cost(new_item)
                        if affects_stock:
                            self._update_stock_for_sales_item(new_item, sale)
                        processed_item_ids.add(new_item.id)

                for item_id, item in existing_items.items():
                    if item_id not in processed_item_ids:
                        self._increase_stock_for_sales_item(item, sale)
                        item.delete()
            elif items_provided:
                processed_item_ids = set()
                for item_data in sales_items_data:
                    item_id = item_data.get('id')
                    if item_id and item_id in existing_items:
                        item = existing_items[item_id]
                        item_serializer = SalesItemsSerializer(
                            instance=item, data=item_data, partial=True, context=self.get_serializer_context()
                        )
                        item_serializer.is_valid(raise_exception=True)
                        updated_item = item_serializer.save()
                        self._snapshot_unit_cost(updated_item)
                        processed_item_ids.add(item_id)
                    else:
                        item_data['sales'] = sale.id
                        item_serializer = SalesItemsSerializer(data=item_data, context=self.get_serializer_context())
                        item_serializer.is_valid(raise_exception=True)
                        new_item = item_serializer.save()
                        self._snapshot_unit_cost(new_item)
                        processed_item_ids.add(new_item.id)

                for item_id, item in existing_items.items():
                    if item_id not in processed_item_ids:
                        item.delete()

            self._ensure_commission_has_salesman(
                sale,
                sales_items_data if items_provided else None,
            )

            response_data = SalesSerializer(sale, context=self.get_serializer_context()).data
            response_data['sales_items'] = SalesItemsSerializer(
                SalesItems.objects.filter(sales=sale), many=True, context=self.get_serializer_context()
            ).data
            return Response(response_data)

    def destroy(self, request, *args, **kwargs):
        with transaction.atomic():
            return super().destroy(request, *args, **kwargs)

    def perform_hard_destroy(self, instance):
        if self._order_affects_stock(instance):
            for item in instance.sales_items.all():
                self._increase_stock_for_sales_item(item, instance)
        instance.delete()

    def _update_stock_for_sales_item(self, item, sale):
        """Apply inventory impact for a sale line (packed stock FEFO or tank liters)."""
        if item.sale_source == SalesItems.SOURCE_TANK:
            validate_tank_fill_quantity(
                item.storage_tank,
                item.quantity,
                exclude_sale_item_id=item.pk,
            )
            return

        from api.utils.packaging import sales_item_pieces

        quantity_needed = sales_item_pieces(item)
        try:
            allocate_inventory_for_sale_item(item, quantity_needed)
        except InsufficientInventoryError as e:
            raise ValidationError(str(e))

    def _update_stock_for_sales_item_update(
        self,
        item,
        sale,
        old_piece_amount,
        old_stock=None,
        old_tank=None,
        old_source=None,
    ):
        from api.utils.packaging import sales_item_pieces

        old_source = old_source or SalesItems.SOURCE_STOCK
        new_source = item.sale_source
        old_pieces = Decimal(str(old_piece_amount or 0))
        new_pieces = sales_item_pieces(item)

        # Leaving packed-stock path → restore FEFO allocations.
        if old_source == SalesItems.SOURCE_STOCK and new_source != SalesItems.SOURCE_STOCK:
            revert_sale_item_allocations(item)

        if new_source == SalesItems.SOURCE_TANK:
            validate_tank_fill_quantity(
                item.storage_tank,
                item.quantity,
                exclude_sale_item_id=item.pk,
            )
            return

        old_stock = old_stock or item.stock
        new_stock = item.stock
        if old_source != SalesItems.SOURCE_STOCK:
            try:
                allocate_inventory_for_sale_item(item, new_pieces)
            except InsufficientInventoryError as e:
                raise ValidationError(str(e))
            return

        if old_stock is None or new_stock is None or old_stock.pk != new_stock.pk:
            revert_sale_item_allocations(item)
            try:
                allocate_inventory_for_sale_item(item, new_pieces)
            except InsufficientInventoryError as e:
                raise ValidationError(str(e))
            return

        update_sale_item_allocations(item, old_pieces, new_pieces)

    def _increase_stock_for_sales_item(self, item, sale):
        """Revert packed-stock allocations; tank liters free automatically when line is removed."""
        if item.sale_source == SalesItems.SOURCE_TANK:
            return
        revert_sale_item_allocations(item)


class SalesItemsViewSet(DataRootViewSet):
    permission_module = 'sales'
    queryset = SalesItems.objects.all().order_by("-id")
    serializer_class = SalesItemsSerializer
    filterset_fields = ["sales", "stock"]

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

    def _stock_view(self):
        return SalesViewSet()

    def create(self, request, *args, **kwargs):
        with transaction.atomic():
            response = super().create(request, *args, **kwargs)
            item = SalesItems.objects.get(id=response.data['id'])
            sale = item.sales
            view = self._stock_view()
            view._snapshot_unit_cost(item)
            if view._order_affects_stock(sale):
                view._update_stock_for_sales_item(item, sale)
            return response

    def update(self, request, *args, **kwargs):
        with transaction.atomic():
            item = self.get_object()
            sale = item.sales
            from api.utils.packaging import sales_item_pieces
            old_piece_amount = sales_item_pieces(item)
            old_stock = item.stock
            old_tank = item.storage_tank
            old_source = item.sale_source
            view = self._stock_view()

            response = super().update(request, *args, **kwargs)
            item = self.get_object()
            view._snapshot_unit_cost(item)
            if view._order_affects_stock(sale):
                view._update_stock_for_sales_item_update(
                    item,
                    sale,
                    old_piece_amount,
                    old_stock=old_stock,
                    old_tank=old_tank,
                    old_source=old_source,
                )
            return response

    def destroy(self, request, *args, **kwargs):
        with transaction.atomic():
            item = self.get_object()
            sale = item.sales
            view = self._stock_view()
            response = super().destroy(request, *args, **kwargs)
            if view._order_affects_stock(sale):
                view._increase_stock_for_sales_item(item, sale)
            return response
