from django.utils import timezone
from django.db import models
from django.contrib.contenttypes.models import ContentType
from account.models import User
from api.models.data.base import BaseModel

from api.models.data.currency import CURRENCY_CHOICES
from api.models.data.vendors import Vendor  
from api.models.data.products import Product
from api.models.data.sales import Sales, SalesItems  # Import the sales models
import uuid

class Return(BaseModel):
    """Sales return model"""
    sales = models.ForeignKey(Sales, on_delete=models.CASCADE, related_name='returns')

    return_number = models.CharField(max_length=50, unique=True, blank=True)
    return_date = models.DateTimeField(default=timezone.now)
    reason = models.CharField(max_length=200, null=True, blank=True)
    currency = models.IntegerField(choices=CURRENCY_CHOICES, default=1)
    total_amount = models.DecimalField(max_digits=15, decimal_places=2, default=0)
    total_amount_base = models.DecimalField(max_digits=15, decimal_places=2, default=0)
    refund_amount = models.DecimalField(max_digits=15, decimal_places=2, default=0)
    refund_amount_base = models.DecimalField(max_digits=15, decimal_places=2, default=0)
    exchange_rate = models.DecimalField(max_digits=18, decimal_places=6, default=1)
    notes = models.CharField(blank=True, null=True, max_length=255)
    
    def save(self, *args, **kwargs):
        if not self.return_number:
            self.return_number = f"RET-{uuid.uuid4().hex[:8].upper()}"
        super().save(*args, **kwargs)
    
    def __str__(self):
        return f"Return {self.return_number} for Sale {self.sales.invoice_number}"
    
    def calculate_total_amount(self):
        """Calculate total return amount from return items"""
        total = sum(item.return_amount for item in self.return_items.all())
        return total

class ReturnItems(BaseModel):
    """Sales return line items"""
    return_order = models.ForeignKey(Return, on_delete=models.CASCADE, related_name='return_items')
    sales_item = models.ForeignKey(SalesItems, on_delete=models.CASCADE)
    liter_amount = models.DecimalField(max_digits=12, decimal_places=3, default=0)
    return_price = models.DecimalField(max_digits=12, decimal_places=2, default=0)
    return_amount = models.DecimalField(max_digits=15, decimal_places=2, default=0)
    
    def __str__(self):
        product_name = self.sales_item.stock.product.name if self.sales_item.stock and self.sales_item.stock.product else 'Unknown Product'
        return f"Return {self.liter_amount}L of {product_name} for {self.return_order.return_number}"
    
    def save(self, *args, **kwargs):
        """Return price is per carton; liter_amount is pieces."""
        from api.utils.packaging import PIECES_PER_CARTON
        from decimal import Decimal, ROUND_HALF_UP

        pieces = Decimal(str(self.liter_amount or 0))
        price = Decimal(str(self.return_price or 0))
        if pieces > 0 and price > 0:
            self.return_amount = (
                (pieces / PIECES_PER_CARTON) * price
            ).quantize(Decimal('0.01'), rounding=ROUND_HALF_UP)
        else:
            self.return_amount = Decimal('0.00')
        super().save(*args, **kwargs)
    
    def clean(self):
        """Validate that return liter_amount doesn't exceed original sale liter_amount"""
        if self.liter_amount > self.sales_item.liter_amount:
            raise ValueError("Return liter amount cannot exceed original sale liter amount")

def get_return_models():
    """
    Returns a dictionary containing the Return and ReturnItems models.
    
    Returns:
        dict: A dictionary with keys 'Return' and 'ReturnItems' containing the respective model classes.
    """
    return {
        'Return': Return,
        'ReturnItems': ReturnItems
    }