"""
recalc_uplift_erp.py - Recalculate promo uplift using ERP promo calendar
=========================================================================
Uses erp_promo_calendar.csv (ground-truth from Gath ERP) instead of the
statistically detected promo flags in sales_clean.csv. This gives cleaner
uplift estimates, especially for Bronze/low-volume SKUs.

Usage:  python recalc_uplift_erp.py
Input:  sales_clean.csv, erp_promo_calendar.csv, sku_plan_list.csv
Output: sku_uplift.csv, cat_uplift.csv (overwrites existing)
"""

import os, csv, sys
import pandas as pd
import numpy as np
from collections import defaultdict

from constants import (
    PROMO_UPLIFT_FALLBACK,
    UPLIFT_MIN_PROMO_WEEKS,
    UPLIFT_MIN_NORMAL_WEEKS,
    UPLIFT_MIN_VALID_SKUS,
)

# Fix Windows console encoding for Croatian characters
if sys.stdout.encoding != 'utf-8':
    sys.stdout.reconfigure(encoding='utf-8', errors='replace')

BASE_DIR = os.path.dirname(os.path.abspath(__file__))
# Check if files are in a data/ subdirectory (app layout) or in the same directory
DATA_DIR = os.path.join(BASE_DIR, 'data')
if not os.path.isdir(DATA_DIR) or not os.path.exists(os.path.join(DATA_DIR, 'sales_clean.csv')):
    DATA_DIR = BASE_DIR  # files alongside the script

def main():
    sales_path = os.path.join(DATA_DIR, 'sales_clean.csv')
    erp_path = os.path.join(DATA_DIR, 'erp_promo_calendar.csv')
    plan_path = os.path.join(DATA_DIR, 'sku_plan_list.csv')

    if not os.path.exists(erp_path):
        print("ERROR: erp_promo_calendar.csv not found. Run build_erp_promo.py first.")
        return

    # Load ERP promo lookup: (sku, year, week) -> 1
    print("Loading ERP promo calendar...")
    erp_df = pd.read_csv(erp_path)
    erp_set = set(zip(erp_df['sku'].astype(str), erp_df['year'].astype(int), erp_df['week'].astype(int)))
    print(f"  {len(erp_set):,} promo-week entries")

    # Load sales
    print("Loading sales data...")
    sales = pd.read_csv(sales_path)
    print(f"  {len(sales):,} rows")

    # Load plan SKUs for category mapping
    plan = pd.read_csv(plan_path)
    sku_cat = dict(zip(plan['sku'], plan['cat']))

    # Tag each sales row with ERP promo flag
    sales['is_erp_promo'] = sales.apply(
        lambda r: 1 if (str(r['sku']), int(r['year']), int(r['week'])) in erp_set else 0, axis=1)

    # Calculate per-SKU uplift
    print("\nCalculating SKU-level uplift...")
    sku_stats = []
    for sku, grp in sales.groupby('sku'):
        promo_mask = grp['is_erp_promo'] == 1
        normal_mask = grp['is_erp_promo'] == 0

        promo_weeks = int(promo_mask.sum())
        normal_weeks = int(normal_mask.sum())

        avg_normal = float(grp.loc[normal_mask, 'qty_total'].mean()) if normal_weeks > 0 else 0
        avg_promo = float(grp.loc[promo_mask, 'qty_total'].mean()) if promo_weeks > 0 else 0

        if avg_normal > 0 and promo_weeks >= UPLIFT_MIN_PROMO_WEEKS:
            uplift = avg_promo / avg_normal
        else:
            uplift = 1.0

        # Wholesale spike detection (keep from original logic)
        ws_spike_weeks = int(grp['is_wholesale_spike'].sum()) if 'is_wholesale_spike' in grp.columns else 0
        ws_normal = grp.loc[grp.get('is_wholesale_spike', 0) == 0, 'qty_wholesale'] if 'is_wholesale_spike' in grp.columns else grp['qty_wholesale']
        ws_promo = grp.loc[grp.get('is_wholesale_spike', 0) == 1, 'qty_wholesale'] if 'is_wholesale_spike' in grp.columns else pd.Series()
        ws_uplift = float(ws_promo.mean() / ws_normal.mean()) if len(ws_normal) > 0 and ws_normal.mean() > 0 and len(ws_promo) > 0 else 1.0

        cat = sku_cat.get(sku, '')

        sku_stats.append({
            'sku': sku,
            'grupacija': cat,
            'promo_weeks': promo_weeks,
            'normal_weeks': normal_weeks,
            'avg_normal': round(avg_normal, 2),
            'avg_promo': round(avg_promo, 2),
            'promo_uplift': round(uplift, 4),
            'ws_spike_weeks': ws_spike_weeks,
            'ws_uplift': round(ws_uplift, 4),
        })

    sku_df = pd.DataFrame(sku_stats)

    # Filter for meaningful uplift (same criteria as engine)
    valid_uplift = sku_df[(sku_df['promo_uplift'] > 1) & (sku_df['promo_weeks'] >= UPLIFT_MIN_PROMO_WEEKS) & (sku_df['normal_weeks'] >= UPLIFT_MIN_NORMAL_WEEKS)]
    print(f"  Total SKUs: {len(sku_df)}")
    print(f"  SKUs with valid uplift (>1.0, >={UPLIFT_MIN_PROMO_WEEKS} promo weeks, >={UPLIFT_MIN_NORMAL_WEEKS} normal weeks): {len(valid_uplift)}")
    print(f"  Median uplift: {valid_uplift['promo_uplift'].median():.2f}")
    print(f"  Mean uplift: {valid_uplift['promo_uplift'].mean():.2f}")

    # Write sku_uplift.csv
    out_sku = os.path.join(DATA_DIR, 'sku_uplift.csv')
    sku_df.to_csv(out_sku, index=False)
    print(f"\n  Written: {out_sku}")

    # Calculate category-level uplift
    print("\nCalculating category-level uplift...")
    cat_stats = []
    for cat, grp in sku_df[sku_df['grupacija'] != ''].groupby('grupacija'):
        valid = grp[(grp['promo_uplift'] > 1) & (grp['promo_weeks'] >= UPLIFT_MIN_PROMO_WEEKS) & (grp['normal_weeks'] >= UPLIFT_MIN_NORMAL_WEEKS)]
        if len(valid) >= UPLIFT_MIN_VALID_SKUS:
            cat_uplift = float(valid['promo_uplift'].median())
        else:
            cat_uplift = PROMO_UPLIFT_FALLBACK
        cat_stats.append({
            'grupacija': cat,
            'cat_promo_uplift': round(cat_uplift, 4),
        })

    cat_df = pd.DataFrame(cat_stats)
    out_cat = os.path.join(DATA_DIR, 'cat_uplift.csv')
    cat_df.to_csv(out_cat, index=False)
    print(f"  Written: {out_cat}")
    print(f"\nCategory uplifts:")
    for _, r in cat_df.iterrows():
        print(f"  {r['grupacija']}: {r['cat_promo_uplift']:.2f}")

    print("\nDone.")


if __name__ == '__main__':
    main()
