from rest_framework import status
from rest_framework.permissions import IsAdminUser
from rest_framework.response import Response
from rest_framework.views import APIView
from django.db.models import Q, Count, Avg, F, DecimalField, Sum
from django.db.models.functions import Coalesce
from decimal import Decimal

from apps.business.marketplace.models_part1 import (
    Vendor,
    VendorContract,
    ContractPricing,
)
def _serialize_contract(contract):
    """Serialize a VendorContract instance to a dict."""
    vendor = contract.vendor
    return {
        "id": str(contract.id),
        "contract_number": contract.contract_number,
        "vendor_id": str(vendor.id) if vendor else None,
        "vendor_name": (vendor.company_name or vendor.store_name) if vendor else None,
        "base_commission_percentage": float(contract.base_commission_percentage or 0),
        "start_date": contract.start_date.isoformat() if contract.start_date else None,
        "end_date": contract.end_date.isoformat() if contract.end_date else None,
        "is_active": contract.is_active,
        "created_at": contract.created_at.isoformat() if hasattr(contract, 'created_at') and contract.created_at else None,
    }


def _serialize_pricing_mapping(mapping):
    """Serialize a ContractPricing instance to a dict."""
    contract = mapping.contract
    vendor = contract.vendor if contract else None
    product = mapping.product
    buy = mapping.buy_price or Decimal("0")
    sell = mapping.sell_price or Decimal("0")
    margin_pct = float((sell - buy) / sell * 100) if sell > 0 else 0
    company_margin = float(sell - buy)
    return {
        "id": str(mapping.id),
        "contract": str(contract.id) if contract else None,
        "contract_number": contract.contract_number if contract else None,
        "vendor_id": str(vendor.id) if vendor else None,
        "vendor_name": (vendor.company_name or vendor.store_name) if vendor else None,
        "product_id": str(product.id) if product else None,
        "product_name": product.name if product else None,
        "pricing_tier": mapping.pricing_tier,
        "pincode": mapping.pincode,
        "franchise_store": str(mapping.franchise_store_id) if mapping.franchise_store_id else None,
        "buy_price": str(buy),
        "sell_price": str(sell),
        "margin_percentage": round(margin_pct, 2),
        "company_margin": round(company_margin, 2),
        "valid_from": mapping.valid_from.isoformat() if mapping.valid_from else None,
        "valid_until": mapping.valid_until.isoformat() if mapping.valid_until else None,
        "is_active": mapping.is_active,
    }


class AdminVendorContractsListView(APIView):
    """
    Get all vendor contracts across the platform with filtering and search.
    Admins can list, search, and filter contracts by vendor, status, date range, commission range.
    """
    permission_classes = [IsAdminUser]

    def get(self, request):
        try:
            page = max(1, int(request.query_params.get("page", 1) or 1))
        except (TypeError, ValueError):
            page = 1

        try:
            page_size = min(100, max(1, int(request.query_params.get("page_size", 20) or 20)))
        except (TypeError, ValueError):
            page_size = 20

        q = (request.query_params.get("q") or "").strip()
        status_filter = (request.query_params.get("status") or "").strip()

        try:
            commission_min = float(request.query_params.get("commission_min") or 0)
            commission_max = float(request.query_params.get("commission_max") or 100)
        except (TypeError, ValueError):
            commission_min = 0
            commission_max = 100

        qs = VendorContract.objects.select_related("vendor", "vendor__user").all().order_by("-created_at")

        if q:
            qs = qs.filter(
                Q(contract_number__icontains=q)
                | Q(vendor__company_name__icontains=q)
                | Q(vendor__store_name__icontains=q)
                | Q(vendor__user__mobile__icontains=q)
                | Q(vendor__user__email__icontains=q)
            )

        if status_filter:
            qs = qs.filter(is_active=(status_filter.lower() == "active"))

        if commission_min is not None or commission_max is not None:
            qs = qs.filter(
                base_commission_percentage__gte=commission_min,
                base_commission_percentage__lte=commission_max,
            )

        total = qs.count()
        start = (page - 1) * page_size
        end = start + page_size

        results = [_serialize_contract(c) for c in qs[start:end]]
        return Response(
            {
                "count": total,
                "results": results,
                "page": page,
                "page_size": page_size,
            }
        )


class AdminVendorContractDetailView(APIView):
    """
    Get and update individual vendor contract details.
    Admins can modify base_commission_percentage and contract dates.
    """
    permission_classes = [IsAdminUser]

    def get(self, request, contract_id):
        try:
            contract = VendorContract.objects.select_related("vendor", "vendor__user").get(id=contract_id)
        except VendorContract.DoesNotExist:
            return Response({"detail": "Contract not found."}, status=status.HTTP_404_NOT_FOUND)

        return Response(_serialize_contract(contract))

    def patch(self, request, contract_id):
        try:
            contract = VendorContract.objects.get(id=contract_id)
        except VendorContract.DoesNotExist:
            return Response({"detail": "Contract not found."}, status=status.HTTP_404_NOT_FOUND)

        base_commission = request.data.get("base_commission_percentage")
        if base_commission is not None:
            try:
                base_commission = float(base_commission)
                if not (0 <= base_commission <= 100):
                    return Response(
                        {"detail": "Commission percentage must be between 0 and 100."},
                        status=status.HTTP_400_BAD_REQUEST,
                    )
                contract.base_commission_percentage = base_commission
            except (TypeError, ValueError):
                return Response(
                    {"detail": "Invalid commission percentage value."},
                    status=status.HTTP_400_BAD_REQUEST,
                )

        if "start_date" in request.data:
            contract.start_date = request.data["start_date"]

        if "end_date" in request.data:
            contract.end_date = request.data["end_date"]

        if "is_active" in request.data:
            contract.is_active = request.data.get("is_active", True)

        contract.save()
        return Response(_serialize_contract(contract))


class AdminPricingMappingsListView(APIView):
    """
    Get all product price mappings across all vendors with filtering.
    Admins can search and filter by vendor, product, tier, and margin range.
    """
    permission_classes = [IsAdminUser]

    def get(self, request):
        try:
            page = max(1, int(request.query_params.get("page", 1) or 1))
        except (TypeError, ValueError):
            page = 1

        try:
            page_size = min(100, max(1, int(request.query_params.get("page_size", 20) or 20)))
        except (TypeError, ValueError):
            page_size = 20

        q = (request.query_params.get("q") or "").strip()
        vendor_id = (request.query_params.get("vendor_id") or "").strip()
        pricing_tier = (request.query_params.get("pricing_tier") or "").strip().upper()

        try:
            margin_min = float(request.query_params.get("margin_min") or 0)
            margin_max = float(request.query_params.get("margin_max") or 100)
        except (TypeError, ValueError):
            margin_min = 0
            margin_max = 100

        qs = ContractPricing.objects.select_related(
            "contract", "contract__vendor", "contract__vendor__user", "product", "franchise_store"
        ).all().order_by("-created_at")

        if q:
            qs = qs.filter(
                Q(product__name__icontains=q)
                | Q(contract__vendor__company_name__icontains=q)
                | Q(contract__contract_number__icontains=q)
            )

        if vendor_id:
            qs = qs.filter(contract__vendor_id=vendor_id)

        if pricing_tier:
            qs = qs.filter(pricing_tier=pricing_tier)

        if margin_min is not None or margin_max is not None:
            # Calculate margin percentage: (sell - buy) / sell * 100
            qs = qs.exclude(sell_price__lte=0)  # Avoid division by zero
            # Filter in Python after fetching (or use F expressions with annotations)
            qs_list = list(qs)
            qs_list = [
                m for m in qs_list
                if m.sell_price and m.buy_price is not None
                and (margin_min <= ((m.sell_price - m.buy_price) / m.sell_price * 100) <= margin_max)
            ]
        else:
            qs_list = list(qs)

        total = len(qs_list)
        start = (page - 1) * page_size
        end = start + page_size

        results = [_serialize_pricing_mapping(m) for m in qs_list[start:end]]
        return Response(
            {
                "count": total,
                "results": results,
                "page": page,
                "page_size": page_size,
            }
        )


class AdminPricingMappingDetailView(APIView):
    """
    Get and update individual pricing mapping.
    Admins can adjust sell price and override margins.
    """
    permission_classes = [IsAdminUser]

    def get(self, request, pricing_id):
        try:
            mapping = ContractPricing.objects.select_related(
                "contract", "contract__vendor", "product", "franchise_store"
            ).get(id=pricing_id)
        except ContractPricing.DoesNotExist:
            return Response({"detail": "Pricing mapping not found."}, status=status.HTTP_404_NOT_FOUND)

        return Response(_serialize_pricing_mapping(mapping))

    def patch(self, request, pricing_id):
        try:
            mapping = ContractPricing.objects.get(id=pricing_id)
        except ContractPricing.DoesNotExist:
            return Response({"detail": "Pricing mapping not found."}, status=status.HTTP_404_NOT_FOUND)

        if "sell_price" in request.data:
            try:
                sell_price = Decimal(str(request.data["sell_price"]))
                if sell_price < 0:
                    return Response(
                        {"detail": "Sell price cannot be negative."},
                        status=status.HTTP_400_BAD_REQUEST,
                    )
                mapping.sell_price = sell_price
            except (TypeError, ValueError):
                return Response(
                    {"detail": "Invalid sell price value."},
                    status=status.HTTP_400_BAD_REQUEST,
                )

        if "buy_price" in request.data:
            try:
                buy_price = Decimal(str(request.data["buy_price"]))
                if buy_price < 0:
                    return Response(
                        {"detail": "Buy price cannot be negative."},
                        status=status.HTTP_400_BAD_REQUEST,
                    )
                mapping.buy_price = buy_price
            except (TypeError, ValueError):
                return Response(
                    {"detail": "Invalid buy price value."},
                    status=status.HTTP_400_BAD_REQUEST,
                )

        if "valid_from" in request.data:
            mapping.valid_from = request.data["valid_from"]

        if "valid_until" in request.data:
            mapping.valid_until = request.data["valid_until"]

        if "is_active" in request.data:
            mapping.is_active = request.data.get("is_active", True)

        mapping.save()
        return Response(_serialize_pricing_mapping(mapping))


class AdminCommercialDashboardView(APIView):
    """
    Get aggregate commercial metrics across all vendors.
    Includes: vendor count, average margin, revenue by tier, low-margin alerts.
    """
    permission_classes = [IsAdminUser]

    def get(self, request):
        # Total active contracts
        total_contracts = VendorContract.objects.filter(is_active=True).count()

        # Total active vendors with contracts
        active_vendors_with_contracts = Vendor.objects.filter(
            vendorcontract__is_active=True
        ).distinct().count()

        # Average commission percentage
        avg_commission = VendorContract.objects.filter(is_active=True).aggregate(
            avg=Avg("base_commission_percentage")
        )["avg"] or 0

        # Total active pricing mappings
        total_mappings = ContractPricing.objects.filter(is_active=True).count()

        # Mappings by tier
        mappings_by_tier = (
            ContractPricing.objects.filter(is_active=True)
            .values("pricing_tier")
            .annotate(count=Count("id"))
            .order_by("pricing_tier")
        )
        tier_breakdown = {m["pricing_tier"]: m["count"] for m in mappings_by_tier}

        # Revenue and margin by vendor
        vendor_metrics = (
            VendorContract.objects.filter(is_active=True)
            .select_related("vendor")
            .values("vendor__id", "vendor__company_name")
            .annotate(
                contract_count=Count("id"),
                avg_commission=Avg("base_commission_percentage"),
                pricing_rules_count=Count("contractpricing", filter=Q(contractpricing__is_active=True)),
                avg_margin_pct=Coalesce(
                    Avg(
                        F("contractpricing__sell_price") - F("contractpricing__buy_price"),
                        output_field=DecimalField(),
                    ),
                    Decimal("0"),
                ),
            )
            .order_by("-contract_count")[:10]
        )

        # Low margin alert (mappings with < 15% margin)
        low_margin_mappings = []
        for mapping in ContractPricing.objects.filter(is_active=True).select_related("contract", "product"):
            if mapping.sell_price and mapping.buy_price is not None:
                margin_pct = ((mapping.sell_price - mapping.buy_price) / mapping.sell_price * 100)
                if margin_pct < 15:
                    low_margin_mappings.append({
                        "id": str(mapping.id),
                        "contract_number": mapping.contract.contract_number,
                        "vendor_name": mapping.contract.vendor.company_name,
                        "product_name": mapping.product.name,
                        "buy_price": float(mapping.buy_price),
                        "sell_price": float(mapping.sell_price),
                        "margin_pct": round(margin_pct, 2),
                    })

        return Response({
            "total_contracts": total_contracts,
            "active_vendors_with_contracts": active_vendors_with_contracts,
            "avg_commission_percentage": float(avg_commission),
            "total_mappings": total_mappings,
            "mappings_by_tier": tier_breakdown,
            "top_vendor_metrics": [
                {
                    "vendor_id": str(v["vendor__id"]),
                    "vendor_name": v["vendor__company_name"],
                    "contract_count": v["contract_count"],
                    "avg_commission": float(v["avg_commission"] or 0),
                    "pricing_rules_count": v["pricing_rules_count"],
                    "avg_margin_pct": float(v["avg_margin_pct"] or 0),
                }
                for v in vendor_metrics
            ],
            "low_margin_alerts": low_margin_mappings[:20],  # Top 20 low margin rules
        })
