# commissions\utils.py
import logging

# somewhere common (e.g., payments/utils.py)
from django.conf import settings
_PROVIDER_ALIASES = {
    "goter": "goterpay",
    "goterpay": "goterpay",
    "a1": "a1topup",
    "a1topup": "a1topup",
}

def normalize_provider_key(vendor_key: str) -> str:
    if not vendor_key:
        return ""
    k = vendor_key.strip().lower()
    return _PROVIDER_ALIASES.get(k, k)

def map_provider_to_source(provider_key: str) -> str:
    # e.g. "goterpay" -> "Goter" using your PROVIDER_SOURCE_MAP
    return settings.PROVIDER_SOURCE_MAP.get(provider_key, provider_key)

from datetime import datetime
from django.utils import timezone

def month_bounds_aware(reference_dt=None):
    """
    Return (month_start, next_month_start) as timezone-aware datetimes
    in the current TIME_ZONE (Asia/Kolkata).
    """
    now_dt = reference_dt or timezone.now()  # aware
    month_start = now_dt.replace(day=1, hour=0, minute=0, second=0, microsecond=0)
    # first day of next month (aware)
    if month_start.month == 12:
        next_month_start = month_start.replace(year=month_start.year + 1, month=1)
    else:
        next_month_start = month_start.replace(month=month_start.month + 1)
    return month_start, next_month_start

def month_bounds_from_date_aware(d):
    """
    Same as above but when you only have a date().
    """
    start = timezone.make_aware(datetime.combine(d.replace(day=1), datetime.min.time()))
    if d.month == 12:
        nm = d.replace(year=d.year + 1, month=1, day=1)
    else:
        nm = d.replace(month=d.month + 1, day=1)
    next_start = timezone.make_aware(datetime.combine(nm, datetime.min.time()))
    return start, next_start


from decimal import Decimal, ROUND_HALF_EVEN
from django.db import transaction
from django.db.models import F
from django.utils import timezone

Q2 = Decimal('0.00')

def _q2(x):  # local helper
    return Decimal(str(x)).quantize(Q2, rounding=ROUND_HALF_EVEN)

@transaction.atomic
def credit_wallet_with_usage(*, user, amount, purpose, purpose_note, order_id):
    """
    Idempotent credit:
    - If a 'credit' usage with this order_id exists, do nothing.
    - Else create usage and atomically bump wallet balance.
    Returns (created: bool, wallet: ViralPeWallet)
    """
    from payments.models import ViralPeWallet, ViralPeWalletUsage  # adjust import to your app

    amt2 = _q2(amount)
    if amt2 <= 0:
        raise ValueError("Credit amount must be > 0")

    # Idempotency guard: don’t double-credit the same order
    exists = ViralPeWalletUsage.objects.select_for_update().filter(
        user=user, order_id=order_id, transaction_type='credit'
    ).exists()
    if exists:
        # Already credited; return current wallet
        wallet, _ = ViralPeWallet.objects.select_for_update().get_or_create(user=user, defaults={"balance": Q2})
        return False, wallet

    # Ensure wallet row exists and lock it
    wallet, _ = ViralPeWallet.objects.select_for_update().get_or_create(user=user, defaults={"balance": Q2})

    # Create usage first (so if anything fails, we roll back both)
    ViralPeWalletUsage.objects.create(
        user=user,
        amount_used=amt2,
        transaction_type='credit',
        purpose=purpose,
        purpose_note=purpose_note,
        order_id=order_id,
    )

    # Atomic balance bump (race-safe)
    # ViralPeWallet.objects.filter(pk=wallet.pk).update(balance=_q2(F('balance') + amt2))
    # wallet.refresh_from_db(fields=['balance', 'last_updated'])
    ViralPeWallet.objects.filter(pk=wallet.pk).update(balance=F('balance') + amt2)
    wallet.refresh_from_db(fields=['balance', 'last_updated'])
    rounded = _q2(wallet.balance)
    if wallet.balance != rounded:
        ViralPeWallet.objects.filter(pk=wallet.pk).update(balance=rounded)
        wallet.balance = rounded
        
    return True, wallet

@transaction.atomic
def debit_wallet_with_usage(*, user, amount, purpose, purpose_note, order_id):
    from payments.models import ViralPeWallet, ViralPeWalletUsage

    amt2 = _q2(amount)
    if amt2 <= 0:
        raise ValueError("Debit amount must be > 0")

    wallet, _ = ViralPeWallet.objects.select_for_update().get_or_create(user=user, defaults={"balance": Q2})

    # Idempotency check for this debit
    exists = ViralPeWalletUsage.objects.select_for_update().filter(
        user=user, order_id=order_id, transaction_type='debit'
    ).exists()
    if exists:
        return False, wallet

    if wallet.balance < amt2:
        raise ValueError("Insufficient balance")

    ViralPeWalletUsage.objects.create(
        user=user,
        amount_used=amt2,
        transaction_type='debit',
        purpose=purpose,
        purpose_note=purpose_note,
        order_id=order_id,
    )

    # ViralPeWallet.objects.filter(pk=wallet.pk).update(balance=_q2(F('balance') - amt2))
    # wallet.refresh_from_db(fields=['balance', 'last_updated'])
    ViralPeWallet.objects.filter(pk=wallet.pk).update(balance=F('balance') - amt2)
    wallet.refresh_from_db(fields=['balance', 'last_updated'])
    rounded = _q2(wallet.balance)
    if wallet.balance != rounded:
        ViralPeWallet.objects.filter(pk=wallet.pk).update(balance=rounded)
        wallet.balance = rounded

    return True, wallet

from commissions.models import CommissionSplit, VendorCodeMapping
from payments.models import CommissionConfig, CommissionSplitConfig, ViralPeWalletUsage
from commissions.models import CommissionSplit, VendorCodeMapping
from django.core.exceptions import ValidationError
from django.db.models import Sum
from recharge.models import Operator
from django.core.exceptions import ValidationError, ObjectDoesNotExist
from decimal import Decimal, ROUND_HALF_UP
from django.utils import timezone

log = logging.getLogger(__name__)
from decimal import Decimal, ROUND_HALF_EVEN, InvalidOperation

Q10 = Decimal('0.0000000001')
Q3  = Decimal('0.000')
Q2  = Decimal('0.00')

def _q10(x): return Decimal(x).quantize(Q10, rounding=ROUND_HALF_EVEN)
def _q3(x):  return Decimal(x).quantize(Q3,  rounding=ROUND_HALF_EVEN)
def _q2(x):  return Decimal(x).quantize(Q2,  rounding=ROUND_HALF_EVEN)

def record_commission_split(order_id, transaction, user, operator, vendor_code, amount,paid_via_gateway,
                            service_type: str | None = None, provider_commission: str| None = None,) :
    """
    Record commission split for a successful transaction.
    Uses:
    - CommissionConfig to fetch fixed commission %
    - CommissionSplitConfig for dynamic role-wise split
    - user.pincode / user.district / user.state
    - VendorCodeMapping for vertical
    """

    # --- normalize inputs ---
    provider_key = normalize_provider_key(vendor_code)               # "goter" -> "goterpay"
    source_value = map_provider_to_source(provider_key)              # "goterpay" -> "Goter"
    service_type = (service_type or "").strip()

    # --- resolve Operator instance ---
    if isinstance(operator, Operator):
        operator_obj = operator
    else:
        # operator is assumed to be a CODE string (e.g., "AIRTEL")
        op_qs = Operator.objects.filter(code__iexact=str(operator).strip(), source__iexact=source_value)
        if service_type:
            op_qs = op_qs.filter(service_type__iexact=service_type)
        operator_obj = op_qs.first()
        if not operator_obj:
            raise ValidationError(
                f"Operator not found for code='{operator}' and source='{source_value}'"
                + (f" and service_type='{service_type}'" if service_type else "")
            )
    try:
        # cfg = CommissionConfig.objects.select_related("operator").get(
        #     operator=operator_obj,
        #     provider_name__iexact=provider_key,   # store keys like "goterpay" in admin
        # )
        # commission_percent = cfg.fixed_percentage

        pass
    except ObjectDoesNotExist:
        raise ValidationError(
            f"CommissionConfig not found for operator='{operator_obj.name}' "
            f"(code={operator_obj.code}, source={operator_obj.source}) "
            f"and provider='{provider_key}'."
        )

    # if not commission_percent or commission_percent <= 0:
    #     raise ValidationError("Commission percent is zero or invalid")
    
    # --- amounts as Decimal ---
    amount_dec = Decimal(str(amount))
    if amount_dec <= 0:
        log.warning("Zero/negative amount in commission split: order_id=%s amount=%s", order_id, amount_dec)
        return None

    gw_dec = Decimal(str(paid_via_gateway))

    # Provider commission (GST-INCLUSIVE money) handling
    try:
        prov_comm_gi = Decimal(str(provider_commission))
    except (InvalidOperation, TypeError):
        raise ValidationError("provider_commission is required as a monetary amount (GST-inclusive)")


    # ---- GST split (18%) ----
    # gross = ex * 1.18  => ex = gross / 1.18
    commission_gross = _q3(prov_comm_gi)  # keep original money at 3dp (since provider gives 3dp)
    commission_ex    = _q10(commission_gross / Decimal('1.18'))
    gst_total        = _q10(commission_gross - commission_ex)

    # ---- gateway vs wallet ratio on EX-GST ----
    gw_ratio     = (gw_dec / amount_dec) if amount_dec > 0 else Decimal("0")
    gw_ratio     = max(Decimal("0"), min(Decimal("1"), gw_ratio))
    wallet_ratio = (Decimal("1") - gw_ratio)

    net_comm_gateway = _q10(commission_ex * gw_ratio)
    net_comm_wallet  = _q10(commission_ex * wallet_ratio)

    # Persist: wallet commission goes as-is; gateway part to be split per-roles later
    # ---- base split_data ----
    split_data = {}

    split_data['order_id']           = order_id
    split_data['operator']           = operator_obj
    split_data['transaction_amount'] = _q2(amount_dec)
    split_data['wallet_amount']      = _q2(amount_dec - gw_dec)
    split_data['gateway_amount']     = _q2(gw_dec)

    # store % at 3dp (on full txn, inclusive): (gross/amount)*100
    commission_percent = (commission_gross / amount_dec * Decimal('100'))
    split_data['commission_percent'] = commission_percent.quantize(Decimal('0.000'), rounding=ROUND_HALF_EVEN)

    # store monetary totals
    split_data['commission_amount']            = commission_gross            # GST-inclusive money (3dp fits your widened field)
    split_data['gst_amount']                   = gst_total                   # 10dp
    split_data['commission_ex_gst_total']      = commission_ex               # 10dp
    split_data['wallet_commission_ex_gst']     = net_comm_wallet             # 10dp
    split_data['gateway_commission_ex_gst']    = net_comm_gateway            # 10dp


    # init all buckets (10dp)
    for k in (
        "user_amount","referral_amount","vendor_refferal_amount",
        "pincode_amount","district_amount","state_amount",
        "vertical_amount","company_amount","tnm_amount"
    ):
        split_data[k] = _q10(0)

    # location / vertical / referrals
    vertical_name = ""
    vendor_mapping = None
    if vendor_code:
        vendor_mapping = VendorCodeMapping.objects.filter(vendor_code=vendor_code, active=True).first()
        if vendor_mapping and vendor_mapping.vertical:
            vertical_name = vendor_mapping.vertical.name

    split_data['vendor_code'] = vendor_mapping if vendor_mapping else None

    split_data['pincode']  = user.pincode or ''
    split_data['district'] = user.district or ''
    split_data['state']    = user.state or ''
    split_data['vertical'] = vertical_name or ''

    # FK users must be None (not empty string)
    split_data['referral_user']        = user.referred_by if getattr(user, "referred_by", None) else None
    split_data['vendor_refferal_user'] = (vendor_mapping.referred_by if (vendor_mapping and vendor_mapping.referred_by) else None)

    # ---- split config (must sum ~100%) ----
    splits = CommissionSplitConfig.objects.all()
    if not splits.exists():
        raise ValidationError("Commission split configuration is missing")

    total_split = splits.aggregate(total=Sum('percentage'))['total'] or 0
    total_dec = Decimal(str(total_split)).quantize(Decimal("0.01"), rounding=ROUND_HALF_EVEN)
    if abs(total_dec - Decimal("100.00")) > Decimal("0.01"):
        raise ValidationError(f"Total commission split must be 100%, found {total_dec}%")

    # log.info("split_data %s", {k: str(v) for k, v in split_data.items()})  # safer logging
    log.info("---- commission gateway_ex_gst=%s, wallet_ex_gst=%s", net_comm_gateway, net_comm_wallet)
    
    # ---- role split from gateway EX-GST only ----
    # (If you also want to share wallet commission, apply same logic to net_comm_wallet.)
    role_amounts = {}

    for cfg in splits:
        role = (cfg.role or "").strip().lower()
        pct  = Decimal(str(cfg.percentage))
        share = _q10(net_comm_gateway * pct / Decimal('100'))

        role_amounts[role] = share

        if role == 'user':
            split_data['user_amount'] = share
            # Credit 2dp to wallet
            amt2 = _q2(share)
            if amt2 != 0:
                # ViralPeWalletUsage.objects.create(
                #     user=user, amount_used=amt2, transaction_type='credit',
                #     purpose='Cashback', purpose_note='Cashback: For User', order_id=order_id
                # )
                credit_wallet_with_usage(
                    user=user,
                    amount=amt2,
                    purpose="Cashback",
                    purpose_note="Cashback: For User",
                    order_id=order_id,
                )

        elif role == 'referral' and split_data['referral_user']:
            split_data['referral_amount'] = share
            amt2 = _q2(share)
            if amt2 != 0:
                # ViralPeWalletUsage.objects.create(
                #     user=split_data['referral_user'], amount_used=amt2, transaction_type='credit',
                #     purpose='Cashback', purpose_note='Cashback: Referral share', order_id=order_id
                # )
                credit_wallet_with_usage(
                    user=split_data['referral_user'],
                    amount=amt2,
                    purpose="Cashback",
                    purpose_note="Cashback: Referral share",
                    order_id=order_id,
                )

        elif role == 'vendor_ref' and split_data['vendor_refferal_user']:
            split_data['vendor_refferal_amount'] = share
            amt2 = _q2(share)
            if amt2 != 0:
                # ViralPeWalletUsage.objects.create(
                #     user=split_data['vendor_refferal_user'], amount_used=amt2, transaction_type='credit',
                #     purpose='Cashback', purpose_note='Cashback: Vendor Referral share', order_id=order_id
                # )
                credit_wallet_with_usage(
                    user=split_data['vendor_refferal_user'],
                    amount=amt2,
                    purpose="Cashback",
                    purpose_note="Cashback: Vendor Referral share",
                    order_id=order_id,
                )
        elif role == 'pincode':
            split_data['pincode_amount'] = share
        elif role == 'district':
            split_data['district_amount'] = share
        elif role == 'state':
            split_data['state_amount'] = share
        elif role == 'vertical':
            split_data['vertical_amount'] = share
        elif role == 'company':
            split_data['company_amount'] = share
        elif role in ('tnm', 't&m', 't n m'):
            split_data['tnm_amount'] = share

    # ---- persist (idempotent by order_id) ----
    model_fields = {
        f.name for f in CommissionSplit._meta.get_fields()
        if getattr(f, "concrete", False) and not getattr(f, "many_to_many", False)
    }
    create_data = {k: v for k, v in split_data.items() if k in model_fields}

    from django.db import IntegrityError

    try:
        log.info("CommissionSplit upsert (order_id=%s)", split_data.get("order_id"))
        obj, created = CommissionSplit.objects.update_or_create(
            order_id=split_data["order_id"],
            defaults=create_data
        )
        log.info("CommissionSplit %s (id=%s)", "created" if created else "updated", obj.id)
    except IntegrityError:
        log.exception("CommissionSplit persistence failed (unique) order_id=%s", split_data.get("order_id"))
        raise
    except Exception:
        log.exception("CommissionSplit persistence failed order_id=%s", split_data.get("order_id"))
        raise

    return split_data
