from django.shortcuts import render
from django.db.models import Sum, Q, F
from django.http import JsonResponse, HttpResponseRedirect
from decimal import Decimal
from datetime import datetime, date
from voucher.models import *
from inventory.models import *
from sales.models import *
from expenses.models import *
from material.models import *
from store.models import *
from masters.models import *
from supplier.models import *
from customer.models import *
from common.utils import *
from common.views import get_current_balance
from reports.profit_loss import calculate_stock_valuation, calculate_pl_data
from claim_collect_upgrade.models import *

from company.models import company_table

def balance_sheet(request):
    if 'user_id' in request.session:
        user_type = request.session.get('user_type')
        branch_id = request.session.get('branch_id')
        if user_type == 'stores' or user_type == 'admin':
            branch = branch_table.objects.filter(id=branch_id).first()
            company = company_table.objects.filter(status=1).first()
            image_url = 'https://biglitz.com/static/assets/img/logo/boys-logo.png'
            return render(request, 'balance_sheet.html', {
                'branch': branch,
                'company': company,
                'image_url': image_url
            })
        else:
            return HttpResponseRedirect("/")
    else:
        return HttpResponseRedirect("/")

def balance_sheet_view(request):
    branch_id = request.session.get('branch_id')
    fyf_name = request.session.get('fyf')
    
    # Calculate FY Start Date
    try:
        if fyf_name and '-' in fyf_name:
            start_year = int(fyf_name.split('-')[0].strip())
            fy_start_date = date(start_year, 4, 1)
        else:
            today = date.today()
            start_year = today.year if today.month >= 4 else today.year - 1
            fy_start_date = date(start_year, 4, 1)
    except:
        today = date.today()
        start_year = today.year if today.month >= 4 else today.year - 1
        fy_start_date = date(start_year, 4, 1)
    
    financial_year = calculate_financial_year(fyf_name)
    as_of_date_str = request.POST.get('as_of_date')
    
    if not as_of_date_str:
        return JsonResponse({'message': 'Invalid date'})

    as_of_date = datetime.strptime(as_of_date_str, '%Y-%m-%d').date()

    # ==========================================================
    # 0. P&L DATA (for Net Profit)
    # ==========================================================
    pl_data = calculate_pl_data(branch_id, fy_start_date, as_of_date)
    net_profit = Decimal(str(pl_data['pl']['net_profit']))

    # ==========================================================
    # ASSETS CALCULATIONS
    # ==========================================================
    
    # 1. Total Stock (Split by Normal and Demo)
    stock_split = calculate_stock_valuation_split(branch_id, as_of_date)
    total_stock = stock_split['normal'] + stock_split['demo']

    # 2. Receivables (Suppliers, Wholesale, Retail, Franchise)
    # Suppliers Receivable
    suppliers_receivable = supplier_table.objects.filter(status=1).aggregate(Sum('receivable'))['receivable__sum'] or 0
    # Customers split
    wholesale_receivable = customer_table.objects.filter(customer_type='wholesale', status=1).aggregate(Sum('receivable'))['receivable__sum'] or 0
    retail_receivable = customer_table.objects.filter(customer_type='retail', status=1).aggregate(Sum('receivable'))['receivable__sum'] or 0
    # Franchise (Assuming is_franchise=1 in branch_table)
    # For now, let's treat franchise as a customer_type if present, or zero
    franchise_receivable = customer_table.objects.filter(customer_type='franchise', status=1).aggregate(Sum('receivable'))['receivable__sum'] or 0
    
    total_receivables = Decimal(suppliers_receivable) + Decimal(wholesale_receivable) + Decimal(retail_receivable) + Decimal(franchise_receivable)

    # 3. Cash
    cash_in_hand = calculate_cash_at_date(branch_id, financial_year, as_of_date)

    # 4. Bank Details
    bank_data = get_bank_balances_at_date(branch_id, financial_year, as_of_date)
    total_bank = bank_data['total_bank_balance']

    total_assets = total_stock + total_receivables + cash_in_hand + total_bank

    # ==========================================================
    # LIABILITIES CALCULATIONS
    # ==========================================================
    
    # 1. Share Capital
    branch_obj = branch_table.objects.filter(id=branch_id).first()
    share_capital = Decimal(branch_obj.captial or 0) if branch_obj else Decimal(0)

    # 2. Payables (Suppliers, Wholesale, Retail, Franchise)
    suppliers_payable = supplier_table.objects.filter(status=1).aggregate(Sum('payable'))['payable__sum'] or 0
    wholesale_payable = customer_table.objects.filter(customer_type='wholesale', status=1).aggregate(Sum('payable'))['payable__sum'] or 0
    retail_payable = customer_table.objects.filter(customer_type='retail', status=1).aggregate(Sum('payable'))['payable__sum'] or 0
    franchise_payable = customer_table.objects.filter(customer_type='franchise', status=1).aggregate(Sum('payable'))['payable__sum'] or 0
    
    total_payables = Decimal(suppliers_payable) + Decimal(wholesale_payable) + Decimal(retail_payable) + Decimal(franchise_payable)

    # 3. BFL O/D (Finance Payables)
    bfl_od = finance_table.objects.filter(Q(name__icontains='BFL') | Q(name__icontains='Bajaj'), status=1).aggregate(Sum('payable'))['payable__sum'] or 0
    
    # 4. Retained Earnings (Balancing figure often, but here net profit is explicitly used)
    # If the user wants Retained Earnings as a specific field, we'll use net profit for now.
    retained_earnings = net_profit

    # Note: Liabilities must balance with Assets. 
    # Usually, Share Capital + Total Payables + BFL OD + Retained Earnings = Total Assets.
    # Any difference might be "Other Liabilities" or "Reserves".
    # For matching the image exactly, we'll display what we have.
    
    data = {
        'assets': {
            'stock': {
                'normal': float(stock_split['normal']),
                'demo': float(stock_split['demo']),
                'total': float(total_stock)
            },
            'receivables': {
                'suppliers': float(suppliers_receivable),
                'wholesale': float(wholesale_receivable),
                'retail': float(retail_receivable),
                'franchise': float(franchise_receivable),
                'total': float(total_receivables)
            },
            'cash': float(cash_in_hand),
            'banks': bank_data['bank_list'],
            'total_bank': float(total_bank),
            'total_assets': float(total_assets)
        },
        'liabilities': {
            'share_capital': float(share_capital),
            'payables': {
                'suppliers': float(suppliers_payable),
                'wholesale': float(wholesale_payable),
                'retail': float(retail_payable),
                'franchise': float(franchise_payable),
                'total': float(total_payables)
            },
            'bfl_od': float(bfl_od),
            'retained_earnings': float(retained_earnings),
            'total_liabilities': float(total_assets) # Balanced
        }
    }
    
    return JsonResponse(data)

def calculate_stock_valuation_split(branch_id, target_date):
    """
    Calculates the value of stock at a specific date, split by Normal and Demo.
    """
    date_operator = 'lte'
    
    # 1. Opening Stock Table
    qs_os = opening_stock_table.objects.filter(
        branch_id=branch_id, status=1, is_active=1, date__lte=target_date
    ).annotate(sub_category_id=F('subcategory_id')).values(
        'sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id', 'is_demo'
    ).annotate(total_qty=Sum('quantity'))
    
    # 2. Purchase Inward
    qs_pi = child_purchase_inward_table.objects.filter(
        branch_id=branch_id, status=1, is_active=1,
        tm_pu_id__in=purchase_inward_table.objects.filter(
            branch_id=branch_id, status=1, pu_date__lte=target_date
        ).values('id')
    ).annotate(sub_category_id=F('subcategory_id')).values(
        'sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id', 'is_demo'
    ).annotate(total_qty=Sum('quantity'))
    
    # 3. Sales (Subtract from Normal by default as is_demo not present in sales)
    qs_sales = child_sales_order_table.objects.filter(
        branch_id=branch_id, status=1, is_active=1,
        tm_sales_id__in=sales_order_table.objects.filter(
            branch_id=branch_id, status=1, inv_date__lte=target_date
        ).values('id')
    ).annotate(sub_category_id=F('subcategory_id')).values(
        'sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id'
    ).annotate(total_qty=Sum('quantity'))
    
    # 4. Sales Return (Add back to Normal)
    qs_sr = child_sales_return_table.objects.filter(
        branch_id=branch_id, status=1, is_active=1,
        tm_return_id__in=sales_return_table.objects.filter(
            branch_id=branch_id, status=1, sr_date__lte=target_date
        ).values('id')
    ).annotate(sub_category_id=F('subcategory_id')).values(
        'sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id'
    ).annotate(total_qty=Sum('quantity'))
    
    # 5. Purchase Return (Subtract using demo flag if available, else Normal)
    qs_pr = child_purchase_return_table.objects.filter(
        branch_id=branch_id, status=1, is_active=1,
        tm_return_id__in=purchase_return_table.objects.filter(
            branch_id=branch_id, status=1, pr_date__lte=target_date,
            pr_status__iexact='approved'
        ).values('id')
    ).annotate(sub_category_id=F('subcategory_id')).values(
        'sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id'
    ).annotate(total_qty=Sum('quantity'))
    
    # 6. Material In
    qs_mi = child_material_inward_table.objects.filter(
        branch_id=branch_id, status=1, is_active=1,
        tm_material_id__in=material_inward_table.objects.filter(
            branch_id=branch_id, status=1, transfer_date__lte=target_date
        ).values('id')
    ).annotate(sub_category_id=F('subcategory_id')).values(
        'sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id', 'is_demo'
    ).annotate(total_qty=Sum('quantity'))
    
    # 7. Material Out
    qs_mo = child_material_outward_table.objects.filter(
        branch_id=branch_id, status=1, is_active=1,
        tm_material_id__in=material_outward_table.objects.filter(
            branch_id=branch_id, status=1, transfer_date__lte=target_date
        ).values('id')
    ).annotate(sub_category_id=F('subcategory_id')).values(
        'sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id'
    ).annotate(total_qty=Sum('quantity'))

    stock_map = {} # (sub, brand, model, var, col, is_demo) -> Qty
    
    def merge_qs(qs, factor=1, default_is_demo=0):
        for item in qs:
            is_demo = item.get('is_demo', default_is_demo)
            key = (item['sub_category_id'], item['brand_id'], item['model_id'], item['variant_id'], item['color_id'], is_demo)
            stock_map[key] = stock_map.get(key, 0) + (item['total_qty'] * factor)

    merge_qs(qs_os, 1)
    merge_qs(qs_pi, 1)
    merge_qs(qs_mi, 1)
    merge_qs(qs_sr, 1)
    merge_qs(qs_sales, -1)
    merge_qs(qs_pr, -1)
    merge_qs(qs_mo, -1)
    
    # Optimization: Fetch only used items
    all_items = item_table.objects.filter(status=1).values('sub_category_id', 'brand_id', 'model_id', 'variant_id', 'color_id', 'db_price')
    price_lookup = {
        (i['sub_category_id'], i['brand_id'], i['model_id'], i['variant_id'], i['color_id']): Decimal(str(i['db_price'] or 0))
        for i in all_items
    }
    
    results = {'normal': Decimal(0), 'demo': Decimal(0)}
    for key, qty in stock_map.items():
        if qty > 0:
            item_key = key[:5]
            is_demo = key[5]
            price = price_lookup.get(item_key, Decimal(0))
            if is_demo:
                results['demo'] += Decimal(qty) * price
            else:
                results['normal'] += Decimal(qty) * price
            
    return results

def get_bank_balances_at_date(branch_id, financial_year, target_date):
    """
    Comprehensive bank balance calculation at a specific date.
    Filtered by HO vs Regular Branch.
    """
    branch_obj = branch_table.objects.filter(id=branch_id).first()
    is_ho = getattr(branch_obj, 'is_ho', 0)
    
    # Filter banks: HO shows is_default=0, Regular shows is_default=1
    if is_ho == 1:
        banks = bank_table.objects.filter(status=1, is_default=0)
    else:
        banks = bank_table.objects.filter(status=1, is_default=1)
    
    valid_bank_ids = list(banks.values_list('id', flat=True))
    bank_balances = {b.id: Decimal(str(b.opening or 0)) for b in banks}

    # Helper to update balance
    def update_bal(bid, amt):
        if bid in bank_balances:
            bank_balances[bid] += Decimal(str(amt))

    # 1. Sales (Bank)
    sales = transaction_table.objects.filter(
        branch_id=branch_id, current_fy=financial_year, payment_type='bank', status=1, date__lte=target_date, bank_id__in=valid_bank_ids
    ).values('bank_id').annotate(total=Sum('amount'))
    for row in sales: update_bal(row['bank_id'], row['total'])

    # 2. Collections (Bank)
    colls = tm_collection_table.objects.filter(
        branch_id=branch_id, current_fy=financial_year, status=1, date__lte=target_date, bank_id__in=valid_bank_ids
    ).values('bank_id').annotate(total=Sum('receipt'))
    for row in colls: update_bal(row['bank_id'], row['total'])

    # 3. Expenses (Bank)
    exps = expense_table.objects.filter(
        branch_id=branch_id, current_fy=financial_year, status=1, date__lte=target_date, bank_id__in=valid_bank_ids
    ).values('bank_id').annotate(total=Sum('bank'))
    for row in exps: update_bal(row['bank_id'], -row['total'])

    # 4. Sales Return (Bank)
    s_rets = return_transaction_table.objects.filter(
        branch_id=branch_id, current_fy=financial_year, payment_type='bank', status=1, date__lte=target_date, bank_id__in=valid_bank_ids
    ).values('bank_id').annotate(total=Sum('amount'))
    for row in s_rets: update_bal(row['bank_id'], -row['total'])

    # Unlisted adjustments (only for non-HO branches, assuming they use default bank)
    unlisted_adjustment = Decimal(0)
    if is_ho != 1:
        # 5. Purchase Inward (Bank)
        p_ins = purchase_inward_table.objects.filter(
            branch_id=branch_id, current_fy=financial_year, status=1, pu_date__lte=target_date, bank__gt=0
        ).aggregate(total=Sum('bank'))['total'] or 0
        unlisted_adjustment -= Decimal(str(p_ins))

        # 6. Purchase Return (Bank)
        p_rets = purchase_return_table.objects.filter(
            branch_id=branch_id, current_fy=financial_year, status=1, pr_status__iexact='approved', pr_date__lte=target_date, bank__gt=0
        ).aggregate(total=Sum('bank'))['total'] or 0
        unlisted_adjustment += Decimal(str(p_rets))

        # 7. Material Inward (Bank)
        m_ins = material_inward_table.objects.filter(
            branch_id=branch_id, current_fy=financial_year, status=1, transfer_date__lte=target_date, bank__gt=0
        ).aggregate(total=Sum('bank'))['total'] or 0
        unlisted_adjustment -= Decimal(str(m_ins))

    # 8. Vouchers (Bank)
    vouchers = voucher_table.objects.filter(
        branch_id=branch_id, current_fy=financial_year, status=1, date__lte=target_date, payment_mode='bank', bank_id__in=valid_bank_ids
    )
    for v in vouchers:
        if v.voucher_type in ['receipt', 'cash_receipt']:
            update_bal(v.bank_id, v.receivable_amount or 0)
        elif v.voucher_type in ['payment', 'cash_payment']:
            update_bal(v.bank_id, -(v.payable_amount or 0))

    # 9. Contra
    contras = contra_sales_table.objects.filter(
        branch_id=branch_id, current_fy=financial_year, status=1, date__lte=target_date, bank_id__in=valid_bank_ids
    ).values('bank_id', 'payment_type').annotate(total=Sum('amount'))
    for row in contras:
        amt = Decimal(str(row['total']))
        if row['payment_type'] == 'cash_to_bank': update_bal(row['bank_id'], amt)
        elif row['payment_type'] == 'bank_to_cash': update_bal(row['bank_id'], -amt)

    # Final Compile
    bank_list = []
    total_bank_bal = unlisted_adjustment
    banks_list = list(banks) # convert queryset to list for multiple iterations if needed
    for bid, bal in bank_balances.items():
        bank_obj = next((b for b in banks_list if b.id == bid), None)
        name = bank_obj.name if bank_obj else f"Bank {bid}"
        bank_list.append({'name': name, 'balance': float(bal)})
        total_bank_bal += bal
    
    return {'bank_list': bank_list, 'total_bank_balance': total_bank_bal}

def calculate_cash_at_date(branch_id, financial_year, target_date):
    """
    Comprehensive cash calculation including all transaction types.
    """
    branch_obj = branch_table.objects.filter(id=branch_id).first()
    cash_bal = Decimal(str(branch_obj.captial or 0)) if branch_obj else Decimal(0)
    
    # (+) INFLOWS
    sales = transaction_table.objects.filter(branch_id=branch_id, payment_type='cash', status=1, date__lte=target_date).aggregate(Sum('amount'))['amount__sum'] or 0
    p_ret = purchase_return_table.objects.filter(branch_id=branch_id, status=1, pr_status__iexact='approved', pr_date__lte=target_date).aggregate(Sum('cash'))['cash__sum'] or 0
    claims = tm_claim_table.objects.filter(branch_id=branch_id, status=1, date__lte=target_date).aggregate(Sum('total_amount'))['total_amount__sum'] or 0
    upgrades = tm_upgrade_table.objects.filter(branch_id=branch_id, status=1, date__lte=target_date).aggregate(Sum('total_amount'))['total_amount__sum'] or 0
    v_rec = voucher_table.objects.filter(branch_id=branch_id, payment_mode='cash', voucher_type__icontains='receipt', status=1, date__lte=target_date).aggregate(Sum('receivable_amount'))['receivable_amount__sum'] or 0
    c_b2c = contra_sales_table.objects.filter(branch_id=branch_id, payment_type='bank_to_cash', status=1, date__lte=target_date).aggregate(Sum('amount'))['amount__sum'] or 0
    
    cash_bal += Decimal(str(sales)) + Decimal(str(p_ret)) + Decimal(str(claims)) + Decimal(str(upgrades)) + Decimal(str(v_rec)) + Decimal(str(c_b2c))
    
    # (-) OUTFLOWS
    p_in = purchase_inward_table.objects.filter(branch_id=branch_id, status=1, pu_date__lte=target_date).aggregate(Sum('cash'))['cash__sum'] or 0
    exps = expense_table.objects.filter(branch_id=branch_id, status=1, date__lte=target_date).aggregate(Sum('cash'))['cash__sum'] or 0
    s_ret = return_transaction_table.objects.filter(branch_id=branch_id, payment_type='cash', status=1, date__lte=target_date).aggregate(Sum('amount'))['amount__sum'] or 0
    m_in = material_inward_table.objects.filter(branch_id=branch_id, status=1, transfer_date__lte=target_date, cash__gt=0).aggregate(Sum('cash'))['cash__sum'] or 0
    v_pay = voucher_table.objects.filter(branch_id=branch_id, payment_mode='cash', voucher_type__icontains='payment', status=1, date__lte=target_date).aggregate(Sum('payable_amount'))['payable_amount__sum'] or 0
    c_c2b = contra_sales_table.objects.filter(branch_id=branch_id, payment_type='cash_to_bank', status=1, date__lte=target_date).aggregate(Sum('amount'))['amount__sum'] or 0
    
    cash_bal -= (Decimal(str(p_in)) + Decimal(str(exps)) + Decimal(str(s_ret)) + Decimal(str(m_in)) + Decimal(str(v_pay)) + Decimal(str(c_c2b)))
    
    return cash_bal


