from django.db import models, transaction
from django.db.models import Q
from rest_framework.decorators import action
from rest_framework.response import Response

from api.models.data.stock import Stock
from api.serializers.data.stock import StockSerializer
from api.views.data.base import DataRootViewSet


def get_expiring_stock_queryset(days_ahead=30, include_expired=True):
    """Legacy Stock rows that have on-hand lots needing expiry attention."""
    from django.db.models import Exists, OuterRef

    from api.models.data.inventory_batch import InventoryBatch
    from api.services.batch_queries import expiring_attention_q

    matching_batches = InventoryBatch.objects.filter(
        product_id=OuterRef('product_id'),
        condition=OuterRef('condition'),
    ).filter(expiring_attention_q(days_ahead, include_expired=include_expired))

    return Stock.objects.filter(Exists(matching_batches))


def get_stock_queryset(qs, user, field='warehouse_id', query_params=None):
    return qs


def _search_batches(qs, search):
    if not search:
        return qs
    from django.db.models import Q

    return qs.filter(
        Q(batch_number__icontains=search)
        | Q(product__name__icontains=search)
        | Q(product__barcode__icontains=search)
        | Q(supplier__name__icontains=search)
    )


class StockViewSet(DataRootViewSet):
    permission_module = 'stock'
    queryset = Stock.objects.select_related().all().order_by("-id")
    serializer_class = StockSerializer
    search_fields = ["product__name", "product__barcode", "notes"]
    
    def get_queryset(self):
        queryset = super().get_queryset()
        
        stock_status = self.request.query_params.get('stock_status')
        if stock_status == 'out_of_stock':
            queryset = queryset.filter(piece_amount=0)
        elif stock_status == 'low_stock':
            queryset = queryset.filter(piece_amount__gt=0, piece_amount__lte=100)
        elif stock_status == 'in_stock':
            queryset = queryset.filter(piece_amount__gt=100)
        elif stock_status == 'expiring':
            days = int(self.request.query_params.get('days_ahead', 30))
            include_expired = str(
                self.request.query_params.get('include_expired', '1')
            ).lower() not in ('0', 'false', 'no')
            queryset = get_expiring_stock_queryset(
                days_ahead=days, include_expired=include_expired
            )
        
        return queryset

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

    @action(detail=False, methods=['get'])
    def expiring(self, request):
        """Near-expiry (and optionally already-expired) on-hand lots."""
        from api.models.data.inventory_batch import InventoryBatch
        from api.serializers.data.inventory_batch import InventoryBatchListSerializer
        from api.services.batch_queries import expiring_attention_q, near_expiry_q

        days = int(request.query_params.get('days_ahead', 30))
        include_expired = str(
            request.query_params.get('include_expired', '1')
        ).lower() not in ('0', 'false', 'no')
        status_filter = (request.query_params.get('expiry_status') or '').strip().lower()

        if status_filter == 'expired':
            from api.services.batch_queries import expired_on_hand_q
            q = expired_on_hand_q()
        elif status_filter == 'expiring':
            q = near_expiry_q(days)
        else:
            q = expiring_attention_q(days, include_expired=include_expired)

        qs = (
            InventoryBatch.objects.filter(q)
            .select_related('product', 'supplier')
            .order_by('expire_date', 'created_at')
        )
        qs = _search_batches(qs, request.query_params.get('search'))

        page = self.paginate_queryset(qs)
        serializer = InventoryBatchListSerializer(
            page if page is not None else qs,
            many=True,
            context={'request': request},
        )
        if page is not None:
            return self.get_paginated_response(serializer.data)
        return Response(serializer.data)

    @action(detail=False, methods=['get'])
    def expired(self, request):
        """Already-expired on-hand lots (active or expired status) with remaining qty."""
        from api.models.data.inventory_batch import InventoryBatch
        from api.serializers.data.inventory_batch import InventoryBatchListSerializer
        from api.services.batch_queries import expired_on_hand_q

        qs = (
            InventoryBatch.objects.filter(expired_on_hand_q())
            .select_related('product', 'supplier')
            .order_by('expire_date', 'created_at')
        )
        qs = _search_batches(qs, request.query_params.get('search'))

        # Flip lingering active+past-date rows to expired without hiding them
        InventoryBatch.objects.filter(
            pk__in=qs.filter(status=InventoryBatch.STATUS_ACTIVE).values('pk')
        ).update(status=InventoryBatch.STATUS_EXPIRED)

        page = self.paginate_queryset(qs)
        serializer = InventoryBatchListSerializer(
            page if page is not None else qs,
            many=True,
            context={'request': request},
        )
        if page is not None:
            return self.get_paginated_response(serializer.data)
        return Response(serializer.data)

    @action(detail=False, methods=['post'])
    def scan_expiry_notifications(self, request):
        """Scan for expiring / low stock and create notifications (delegates to service)."""
        from api.services.notifications import scan_stock_alerts

        days = int(request.data.get('days_ahead', 30) or 30)
        threshold = int(request.data.get('low_stock_threshold', 100) or 100)
        result = scan_stock_alerts(
            days_ahead=days,
            low_stock_threshold=threshold,
            include_low_stock=True,
            mark_expired=True,
        )
        return Response(
            {
                'created_or_updated': result['created'],
                'refreshed': result['refreshed'],
                'resolved': result['resolved'],
            }
        )

    @action(detail=True, methods=['post'])
    def move_condition(self, request, pk=None):
        """
        Move quantity from this stock lot into another condition lot
        (typically good → damaged).

        Body:
          - condition: target condition (required), e.g. "damaged"
          - dana_amount: pieces to move (required)
          - buy_price: optional per-piece price (defaults to source buy_price)
          - notes: optional
        """
        from api.services.stock_condition_move import move_stock_to_condition

        source = self.get_object()
        target_condition = (request.data.get('condition') or '').strip()
        dana_amount = request.data.get('dana_amount')
        buy_price = request.data.get('buy_price', None)
        notes = (request.data.get('notes') or '').strip()

        target = move_stock_to_condition(
            source=source,
            target_condition=target_condition,
            piece_amount=request.data.get('piece_amount', dana_amount),
            dana_amount=dana_amount,
            buy_price=buy_price,
            notes=notes,
            batch_id=request.data.get('batch_id') or None,
        )
        return Response(self.get_serializer(target).data)

    @action(detail=True, methods=['post'])
    def return_to_main(self, request, pk=None):
        """

        Body:
          - piece_amount / dana_amount: pieces to return
          - batch_id: optional specific batch
          - reason: expired | damaged | other (optional; inferred from lot)
          - notes: optional
        """
        from api.models.data.inventory_batch import InventoryBatch
        from api.serializers.data.stock_transfer import StockTransferSerializer
        from api.services.stock_return_to_main import return_stock_to_main

        source = self.get_object()
        batch_id = request.data.get('batch_id') or None
        batch = None
        if batch_id:
            batch = InventoryBatch.objects.filter(pk=batch_id).first()
            if not batch:
                return Response(
                    {'batch_id': f'Batch {batch_id} not found.'},
                    status=400,
                )

        transfer, target = return_stock_to_main(
            source=source,
            batch=batch,
            piece_amount=request.data.get(
                'piece_amount', request.data.get('dana_amount')
            ),
            notes=(request.data.get('notes') or '').strip(),
            reason=(request.data.get('reason') or '').strip() or None,
            user=request.user,
        )
        return Response(
            {
                'transfer': StockTransferSerializer(transfer).data,
                'target_stock': self.get_serializer(target).data,
            }
        )

    @action(detail=True, methods=['post'])
    def write_off(self, request, pk=None):
        """
        Write off / dispose expired or damaged stock (typical at main warehouse).

        Posts Dr Stock Adjustments · Cr Inventory. Reduces on-hand qty.

        Body:
          - piece_amount / dana_amount: pieces to write off
          - batch_id: optional specific batch
          - adjustment_type: expired | damaged | lost (optional; inferred)
          - notes: optional
        """
        from api.models.data.inventory_batch import InventoryBatch
        from api.serializers.data.inventory_batch import StockAdjustmentSerializer
        from api.services.stock_write_off import write_off_stock

        source = self.get_object()
        batch_id = request.data.get('batch_id') or None
        batch = None
        if batch_id:
            batch = InventoryBatch.objects.filter(pk=batch_id).first()
            if not batch:
                return Response(
                    {'batch_id': f'Batch {batch_id} not found.'},
                    status=400,
                )

        result = write_off_stock(
            source=source,
            batch=batch,
            piece_amount=request.data.get(
                'piece_amount', request.data.get('dana_amount')
            ),
            adjustment_type=(request.data.get('adjustment_type') or '').strip() or None,
            notes=(request.data.get('notes') or '').strip(),
            user=request.user,
        )
        source.refresh_from_db()
        return Response(
            {
                'adjustment_type': result['adjustment_type'],
                'piece_amount': str(result['piece_amount']),
                'adjustments': StockAdjustmentSerializer(
                    result['adjustments'], many=True
                ).data,
                'stock': self.get_serializer(source).data,
            }
        )
