from rest_framework import serializers
from .models import (
    Product, Location, StockLevel, Device, Transaction, TransactionItem,
    StockMovement, Size, StockLevelSize, LocationCashBalance, CashTransaction, SalesTarget, BonusRecord, Expense, ExpenseCategory,
    CreditSale
)
from django.db import transaction as db_transaction
from django.utils.dateparse import parse_datetime
from django.utils import timezone
from django.contrib.auth.models import User
import logging
from decimal import Decimal

logger = logging.getLogger(__name__)


# --- Basic Serializers ---
class ProductSerializer(serializers.ModelSerializer):
    class Meta:
        model = Product
        fields = '__all__'


class LocationSerializer(serializers.ModelSerializer):
    class Meta:
        model = Location
        fields = ['id', 'name'] 


class SizeSerializer(serializers.ModelSerializer):
    class Meta:
        model = Size
        fields = ('id', 'name')


class StockLevelSizeSerializer(serializers.ModelSerializer):
    size_id = serializers.UUIDField(write_only=True)
    size = SizeSerializer(read_only=True)

    class Meta:
        model = StockLevelSize
        fields = ('id', 'stock_level', 'size', 'size_id', 'quantity')


class StockLevelSerializer(serializers.ModelSerializer):
    product = ProductSerializer(read_only=True)
    product_id = serializers.UUIDField(write_only=True)
    location_id = serializers.UUIDField(write_only=True)
    sizes = StockLevelSizeSerializer(many=True)

    class Meta:
        model = StockLevel
        fields = ('id', 'product', 'product_id', 'location_id', 'quantity', 'updated_at', 'sizes')

    def create(self, validated_data):
        sizes_data = validated_data.pop('sizes', [])
        stock, _ = StockLevel.objects.get_or_create(
            product_id=validated_data['product_id'],
            location_id=validated_data['location_id'],
            defaults={'quantity': 0}
        )

        for size_data in sizes_data:
            StockLevelSize.objects.update_or_create(
                stock_level=stock,
                size_id=size_data['size_id'],
                defaults={'quantity': size_data['quantity']}
            )

        stock.update_total_quantity()
        return stock


class DeviceSerializer(serializers.ModelSerializer):
    location = LocationSerializer(read_only=True)
    class Meta:
        model = Device
        fields = '__all__'


class TransactionItemSerializer(serializers.ModelSerializer):
    product_id = serializers.UUIDField(write_only=True)
    size_id = serializers.UUIDField(write_only=True, required=False, allow_null=True)
    product = ProductSerializer(read_only=True)
    size = SizeSerializer(read_only=True)

    class Meta:
        model = TransactionItem
        fields = ('id', 'product', 'product_id', 'size', 'size_id', 'quantity', 'unit_price', 'line_total')


# ✅ NEW: Serializers for Bank/Cash Management
class LocationCashBalanceSerializer(serializers.ModelSerializer):
    location_name = serializers.CharField(source='location.name', read_only=True)
    payment_method_display = serializers.CharField(source='get_payment_method_display', read_only=True)
    
    class Meta:
        model = LocationCashBalance
        fields = '__all__'


class CashTransactionSerializer(serializers.ModelSerializer):
    location_name = serializers.CharField(source='location.name', read_only=True)
    payment_method_display = serializers.CharField(source='get_payment_method_display', read_only=True)
    transaction_type_display = serializers.CharField(source='get_transaction_type_display', read_only=True)
    created_by_username = serializers.CharField(source='created_by.username', read_only=True)
    
    class Meta:
        model = CashTransaction
        fields = '__all__'


# ✅ UPDATED: Transaction Serializer with Payment Method and Bank Integration

class TransactionSerializer(serializers.ModelSerializer):
    id        = serializers.UUIDField()
    device_id = serializers.UUIDField(write_only=True, required=False, allow_null=True)
    items     = TransactionItemSerializer(many=True)
 
    # ── Standard payment fields ────────────────────────────────────────────
    payment_method = serializers.ChoiceField(
        choices=Transaction.PAYMENT_METHODS,
        default='cash'
    )
    mpesa_reference = serializers.CharField(
        required=False, allow_blank=True, allow_null=True
    )
    payment_phone = serializers.CharField(
        required=False, allow_blank=True, allow_null=True
    )
 
    # ── Credit-sale fields (write-only; ignored for non-credit payments) ──
    credit_customer_name  = serializers.CharField(
        required=False, allow_blank=True, allow_null=True, write_only=True
    )
    credit_customer_phone = serializers.CharField(
        required=False, allow_blank=True, allow_null=True, write_only=True
    )
    credit_due_date = serializers.DateField(
        required=False, allow_null=True, write_only=True
    )
 
    # ── Read-only credit info (returned in response) ───────────────────────
    credit_info = serializers.SerializerMethodField(read_only=True)
 
    class Meta:
        model = Transaction
        fields = (
            'id',
            'device_id',
            'created_at',
            'total_amount',
            'payment_method',
            'mpesa_reference',
            'payment_phone',
            'credit_customer_name',
            'credit_customer_phone',
            'credit_due_date',
            'credit_info',
            'items',
            'raw',
        )
 
    def get_credit_info(self, obj):
        """Expose credit details in the response (read-only)."""
        try:
            c = obj.credit_sale
            return {
                'customer_name':  c.customer_name,
                'customer_phone': c.customer_phone,
                'due_date':       c.due_date.isoformat(),
                'is_paid':        c.is_paid,
                'is_overdue':     c.is_overdue,
            }
        except Exception:
            return None
 
    # ── Validation ─────────────────────────────────────────────────────────
    def validate(self, data):
        if isinstance(data.get('created_at'), str):
            data['created_at'] = parse_datetime(data['created_at'])
 
        payment_method = data.get('payment_method', 'cash')
 
        # M-Pesa reference required
        if payment_method in ['mpesa_paybill', 'mpesa_till']:
            if not data.get('mpesa_reference'):
                raise serializers.ValidationError(
                    "M-Pesa reference is required for M-Pesa payments"
                )
 
        # Credit fields required
        if payment_method == 'on_credit':
            if not (data.get('credit_customer_name') or '').strip():
                raise serializers.ValidationError(
                    "Customer name is required for credit sales"
                )
            if not (data.get('credit_customer_phone') or '').strip():
                raise serializers.ValidationError(
                    "Customer phone is required for credit sales"
                )
            if not data.get('credit_due_date'):
                raise serializers.ValidationError(
                    "Due date is required for credit sales"
                )
 
        return data
 
    # ── Cash balance helper ────────────────────────────────────────────────
    def _update_cash_balance(self, location, payment_method, amount, transaction_obj, user):
        """
        Updates LocationCashBalance and logs a CashTransaction.
        SKIPPED for 'on_credit' — money has not been received yet.
        Balance is updated later when admin calls mark_credit_paid.
        """
        if payment_method == 'on_credit':
            return
 
        amount = Decimal(str(amount))
 
        cash_balance, _ = LocationCashBalance.objects.get_or_create(
            location=location,
            payment_method=payment_method,
            defaults={
                'balance':          Decimal('0.00'),
                'total_sales':      Decimal('0.00'),
                'total_purchases':  Decimal('0.00'),
            }
        )
        cash_balance.balance     += amount
        cash_balance.total_sales += amount
        cash_balance.save()
 
        CashTransaction.objects.create(
            location=location,
            payment_method=payment_method,
            transaction_type='sale',
            amount=amount,
            sale_transaction=transaction_obj,
            balance_after=cash_balance.balance,
            created_by=user,
            description=(
                f"Sale transaction "
                f"{transaction_obj.server_receipt_number or transaction_obj.id}"
            )
        )
 
    # ── Create ─────────────────────────────────────────────────────────────
    def create(self, validated_data):
        items     = validated_data.pop('items')
        device_id = validated_data.pop('device_id', None)
        tx_id     = validated_data.get('id')
 
        payment_method = validated_data.get('payment_method', 'cash')
        mpesa_reference = validated_data.pop('mpesa_reference', None)
        payment_phone   = validated_data.pop('payment_phone', None)
 
        # Pop credit fields (not on Transaction model)
        credit_customer_name  = validated_data.pop('credit_customer_name', None)
        credit_customer_phone = validated_data.pop('credit_customer_phone', None)
        credit_due_date       = validated_data.pop('credit_due_date', None)
 
        with db_transaction.atomic():
 
            # ── Idempotency check ─────────────────────────────────────────
            existing = Transaction.objects.select_for_update().filter(id=tx_id).first()
            if existing:
                return existing
 
            device   = Device.objects.filter(id=device_id).first() if device_id else None
            location = self._get_location(device)
            if not location:
                raise serializers.ValidationError(
                    "Cannot determine location for this device"
                )
 
            created_at = validated_data.get('created_at') or timezone.now()
 
            # ── Total verification ────────────────────────────────────────
            frontend_total    = float(validated_data.get('total_amount', 0))
            calculated_total  = sum(
                int(item['quantity']) * float(item['unit_price']) for item in items
            )
            if abs(frontend_total - calculated_total) > 0.01:
                raise serializers.ValidationError(
                    f"Transaction total mismatch. "
                    f"Expected {calculated_total}, got {frontend_total}"
                )
 
            # ── Lock stock rows ───────────────────────────────────────────
            product_ids  = [item['product_id'] for item in items]
            stock_locks  = StockLevel.objects.select_for_update().filter(
                product_id__in=product_ids,
                location=location
            )
            stock_dict   = {sl.product_id: sl for sl in stock_locks}
 
            # ── Pre-validate every item before touching any stock ─────────
            validation_data = []
            for item in items:
                product = Product.objects.filter(id=item['product_id']).first()
                if not product:
                    raise serializers.ValidationError(
                        f"Product {item['product_id']} not found"
                    )
 
                qty   = int(item['quantity'])
                stock = stock_dict.get(product.id)
                if not stock:
                    raise serializers.ValidationError(
                        f"No stock for {product.name} at {location.name}"
                    )
 
                size_instance = None
                size_id = item.get('size_id')
                if size_id:
                    size_instance = StockLevelSize.objects.select_for_update().filter(
                        id=size_id,
                        stock_level=stock
                    ).first()
                    if not size_instance:
                        raise serializers.ValidationError(
                            f"Stock size not found for {product.name} at {location.name}"
                        )
                    if size_instance.quantity < qty:
                        raise serializers.ValidationError(
                            f"Not enough stock for {product.name} "
                            f"(Size {size_instance.size.name}). "
                            f"Need: {qty}, Have: {size_instance.quantity}"
                        )
 
                validation_data.append({
                    'product':       product,
                    'stock':         stock,
                    'size_instance': size_instance,
                    'quantity':      qty,
                    'unit_price':    float(item['unit_price']),
                })
 
            # ── Create Transaction ────────────────────────────────────────
            tx = Transaction.objects.create(
                id=tx_id,
                device=device,
                created_at=created_at,
                total_amount=calculated_total,
                payment_method=payment_method,
                mpesa_reference=mpesa_reference,
                payment_phone=payment_phone,
                raw=validated_data.get('raw'),
                synced=False,
            )
 
            # ── Deduct stock and create items / movements ─────────────────
            for item_data in validation_data:
                product       = item_data['product']
                size_instance = item_data['size_instance']
                qty           = item_data['quantity']
                unit_price    = item_data['unit_price']
 
                if size_instance:
                    size_instance.quantity -= qty
                    size_instance.save()
 
                TransactionItem.objects.create(
                    transaction=tx,
                    product=product,
                    size=size_instance.size if size_instance else None,
                    quantity=qty,
                    unit_price=unit_price,
                    line_total=qty * unit_price,
                )
 
                StockMovement.objects.create(
                    product=product,
                    delta=-qty,
                    reason='sale',
                    transaction=tx,
                    location=location,
                )
 
            # ── Assign receipt number ─────────────────────────────────────
            last = (
                Transaction.objects
                .select_for_update()
                .order_by("-server_receipt_number")
                .first()
            )
            tx.server_receipt_number = (
                (last.server_receipt_number or 1000) + 1 if last else 1000
            )
            tx.synced = True
            tx.save()
 
            # ── Create CreditSale record if applicable ────────────────────
            if payment_method == 'on_credit':
                from .models import CreditSale
                CreditSale.objects.create(
                    transaction=tx,
                    customer_name=credit_customer_name.strip(),
                    customer_phone=credit_customer_phone.strip(),
                    due_date=credit_due_date,
                )
 
            # ── Update cash balance (skipped for on_credit) ───────────────
            user = (
                self.context['request'].user
                if self.context.get('request') else None
            )
            self._update_cash_balance(
                location, payment_method, calculated_total, tx, user
            )
 
            return tx
 
    def _get_location(self, device):
        if not device:
            return None
        if device.location:
            return device.location
        if device.assigned_to:
            user = User.objects.filter(username=device.assigned_to).first()
            if user and hasattr(user, 'userprofile') and user.userprofile.location:
                return user.userprofile.location
        return None
 



# --- Other Serializers ---
class GlobalStockSerializer(serializers.Serializer):
    product_id = serializers.UUIDField()
    product_name = serializers.CharField()
    total_qty = serializers.IntegerField()


class ProductWithStockSerializer(serializers.ModelSerializer):
    stock = serializers.SerializerMethodField()
    image = serializers.SerializerMethodField()
    
    class Meta:
        model = Product
        fields = ('id', 'name', 'unit_price', 'image', 'stock')
    
    def get_image(self, obj):
        """Return full URL for product image"""
        if obj.image:
            request = self.context.get('request')
            if request:
                return request.build_absolute_uri(obj.image.url)
            return obj.image.url
        return None
    
    def get_stock(self, obj):
        """
        Returns stock information grouped by location with sizes.
        Each size entry includes the StockLevelSize id, size name, and quantity.
        """
        levels = StockLevel.objects.filter(product=obj).prefetch_related(
            'location',
            'sizes',      # This prefetches StockLevelSize objects
            'sizes__size' # This prefetches the related Size objects
        )
        
        data = []
        for lvl in levels:
            # lvl.sizes.all() returns StockLevelSize objects
            sizes_data = []
            for stock_level_size in lvl.sizes.all():
                sizes_data.append({
                    "id": str(stock_level_size.id),        # StockLevelSize ID (UUID)
                    "size": stock_level_size.size.name,    # Size name (e.g., "42")
                    "quantity": stock_level_size.quantity  # Quantity for this size
                })
            
            data.append({
                "location": lvl.location.name,
                "total_quantity": lvl.quantity,
                "sizes": sizes_data
            })
        
        return data

# Add these to your existing serializers.py

from .models import SalesTarget, BonusRecord, Expense, ExpenseCategory

class SalesTargetSerializer(serializers.ModelSerializer):
    location_name = serializers.CharField(source='location.name', read_only=True)
    current_sales = serializers.SerializerMethodField()
    bonus_amount = serializers.SerializerMethodField()
    progress_percentage = serializers.SerializerMethodField()
    target_met = serializers.SerializerMethodField()
    
    class Meta:
        model = SalesTarget
        fields = '__all__'
    
    def get_current_sales(self, obj):
        return float(obj.get_current_sales())
    
    def get_bonus_amount(self, obj):
        return float(obj.calculate_bonus())
    
    def get_progress_percentage(self, obj):
        return float(obj.get_progress_percentage())
    
    def get_target_met(self, obj):
        return obj.get_current_sales() >= obj.target_amount


class BonusRecordSerializer(serializers.ModelSerializer):
    location_name = serializers.CharField(source='location.name', read_only=True)
    
    class Meta:
        model = BonusRecord
        fields = '__all__'


class ExpenseCategorySerializer(serializers.ModelSerializer):
    class Meta:
        model = ExpenseCategory
        fields = ['id', 'name', 'category_type', 'is_active']


class MobileExpenseSerializer(serializers.ModelSerializer):
    """Simplified expense serializer for mobile app (business expenses only)"""
    category_id = serializers.UUIDField(write_only=True)
    category_name = serializers.CharField(source='category.name', read_only=True)
    
    class Meta:
        model = Expense
        fields = [
            'id', 'date', 'amount', 'category_id', 'category_name',
            'description', 'reference_number', 'created_at'
        ]
        read_only_fields = ['id', 'created_at']
    
    def create(self, validated_data):
        category_id = validated_data.pop('category_id')
        
        # Get user from request context
        user = self.context['request'].user
        
        # Get user's location
        location = None
        if hasattr(user, 'userprofile') and user.userprofile.location:
            location = user.userprofile.location
        else:
            # Default to MainStore
            location = Location.objects.filter(name__icontains='main').first()
        
        if not location:
            raise serializers.ValidationError("Cannot determine location for expense")
        
        # Create expense (always business type, always bank transfer)
        expense = Expense.objects.create(
            category_id=category_id,
            expense_type='business',  # Always business
            paid_from_location=location,
            payment_method='bank_transfer',  # Always bank
            created_by=user,
            **validated_data
        )
        
        return expense
    
class CreditSaleSerializer(serializers.ModelSerializer):
    transaction_id     = serializers.UUIDField(source='transaction.id', read_only=True)
    receipt_number     = serializers.IntegerField(source='transaction.server_receipt_number', read_only=True)
    amount             = serializers.DecimalField(
        source='transaction.total_amount', max_digits=12, decimal_places=2, read_only=True
    )
    location           = serializers.SerializerMethodField()
    is_overdue         = serializers.BooleanField(read_only=True)
    days_overdue       = serializers.IntegerField(read_only=True)
    days_until_due     = serializers.IntegerField(read_only=True)
 
    class Meta:
        from .models import CreditSale
        model  = CreditSale
        fields = (
            'id',
            'transaction_id',
            'receipt_number',
            'amount',
            'customer_name',
            'customer_phone',
            'due_date',
            'is_paid',
            'paid_at',
            'paid_payment_method',
            'is_overdue',
            'days_overdue',
            'days_until_due',
            'notes',
            'location',
            'created_at',
        )
 
    def get_location(self, obj):
        tx = obj.transaction
        if tx.device and tx.device.location:
            return tx.device.location.name
        return 'Unknown'