"""Per-SKU promo P&L for W22–W26 2026, anchored on ACTUAL PAST PROMO SALES.

Why a second script
-------------------
The original ``promo_calc_w22_w26.py`` calls PromoTool's ``suggest_uplift``
which can produce extreme multipliers (e.g. 66×) when a SKU's pre-promo
baseline happened to be near zero — the maths is right, the result isn't
useful for planning.

This script ignores the uplift model and uses the **observed weekly run-
rate during past promos** as the forecast level.  No multiplication
against a fragile baseline.  Each line is therefore "what this SKU does
when on promo", measured directly.

Method
------
1. Find every ERP-recorded promo event for the SKU
   (``erp_promo_calendar.csv``, grouped by ``promo_types``).
2. For each event, sum ``qty_retail + qty_webshop`` per week and average
   across the event's weeks → ``event_avg_per_week``.
3. SKU-level run-rates:
     median_run_rate = median of event averages
     max_run_rate    = peak weekly qty observed during any past promo
4. Forecast for W22–W26 (5 weeks):
     Forecast units (median) = median_run_rate × 5
     Max units (peak-week)   = max_run_rate    × 5
5. SKUs with no own promo history fall back to category siblings (same
   ``cat`` in ``sku_plan_list.csv``).  Use the siblings' median
   per-week promo qty, then scale by the ratio of this SKU's recent
   baseline to the sibling baseline so the magnitude tracks volume.
6. SKUs with neither own nor sibling promo history get the
   "capped-uplift" fallback: current baseline × ``CAP_UPLIFT`` (default
   3×) — flagged so reviewers can override manually.

P&L (RUC, revenue, margin) is computed identically to the prior script
so the two outputs are directly comparable.
"""
from __future__ import annotations

import argparse
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parent
sys.stdout.reconfigure(encoding="utf-8")  # type: ignore[attr-defined]

import pandas as pd

# --------------------------------------------------------------------------
# Constants
# --------------------------------------------------------------------------

PROMO_YEAR  = 2026
PROMO_WEEKS = [22, 23, 24, 25, 26]
N_WEEKS     = len(PROMO_WEEKS)

DATA_DIR       = ROOT / "data"
POLICY_PRIMARY = DATA_DIR / "VrstaRabatnePolitike06.xlsx"
POLICY_FALLBK  = DATA_DIR / "VrstaRabatnePolitikesummer buddy.xlsx"

# Hard ceiling for the no-history fallback.  3× is a conservative anchor
# CMs reported as "what we typically see for 20–30 % depths".
CAP_UPLIFT = 3.0

# --------------------------------------------------------------------------
# Master data
# --------------------------------------------------------------------------

def _csv(name: str) -> pd.DataFrame:
    p = DATA_DIR / name
    return pd.read_csv(p) if p.exists() else pd.DataFrame()


sales     = _csv("sales_clean.csv")
prices    = _csv("sku_prices.csv")
costs     = _csv("sku_costs.csv")
plan_df   = _csv("sku_plan_list.csv")
erp_promo = _csv("erp_promo_calendar.csv")

name_map = dict(zip(plan_df["sku"], plan_df.get("name", "")))  if "sku" in plan_df.columns else {}
cat_map  = dict(zip(plan_df["sku"], plan_df.get("cat", "")))   if "sku" in plan_df.columns else {}
tier_map = dict(zip(plan_df["sku"], plan_df.get("oznaka", ""))) if "sku" in plan_df.columns else {}
price_map = (
    dict(zip(prices["sku"], prices["normal_retail_ppp"]))
    if "sku" in prices.columns and "normal_retail_ppp" in prices.columns else {}
)
cost_map = (
    dict(zip(costs["sku"], costs["cost_price"]))
    if "sku" in costs.columns and "cost_price" in costs.columns else {}
)

# Sales is large — pre-index by SKU for speed.
if not sales.empty:
    sales["_rcm_web"] = sales["qty_retail"].fillna(0) + sales["qty_webshop"].fillna(0)
sales_by_sku = {k: g for k, g in sales.groupby("sku")} if not sales.empty else {}
erp_by_sku   = {k: g for k, g in erp_promo.groupby("sku")} if not erp_promo.empty else {}

# Pre-compute siblings by category for fallback.
cat_to_skus: dict[str, list[str]] = {}
for s, c in cat_map.items():
    if c:
        cat_to_skus.setdefault(c, []).append(s)


# --------------------------------------------------------------------------
# Policy loader (same as the first script — duplicated rather than imported
# to keep this file standalone and easy to copy).
# --------------------------------------------------------------------------

def _read_policy(path: Path) -> pd.DataFrame:
    if not path.exists():
        return pd.DataFrame()
    df = pd.read_excel(path, sheet_name="RabatnaPolitika")
    df = df[df["Artikal"].notna() & (df["Artikal"].astype(str).str.strip() != "")]
    df["Artikal"] = df["Artikal"].astype(str).str.strip()
    df["_source"] = path.name
    return df


def load_policy_map() -> dict[str, dict]:
    pieces = [p for p in (_read_policy(POLICY_PRIMARY), _read_policy(POLICY_FALLBK)) if not p.empty]
    if not pieces:
        return {}
    combined = pd.concat(pieces, ignore_index=True)
    out: dict[str, dict] = {}
    for _, r in combined.iterrows():
        sku = r["Artikal"]
        if sku in out:
            continue
        def _f(col):
            v = r.get(col)
            try:
                return float(v) if v is not None and pd.notna(v) else None
            except (TypeError, ValueError):
                return None
        out[sku] = {
            "discount_pct":    _f("Rabat %"),
            "promo_price":     _f("Cijena"),
            "min_qty":         _f("Količina >"),
            "policy_name":     r.get("Rabatna politika"),
            "policy_source":   r.get("_source"),
            "policy_sku_name": r.get("Naziv artikla"),
        }
    return out


# --------------------------------------------------------------------------
# Past-promo run-rate extraction
# --------------------------------------------------------------------------

def past_promo_events(sku: str) -> list[dict]:
    """List of past promo events for ``sku`` with per-event run-rate.

    Each entry:
      {
        "promo_types": "AKCIJE 07/25",
        "year": 2025, "start_week": 27, "end_week": 31,
        "n_weeks": 5, "total_qty": 161, "avg_per_week": 32.2,
        "peak_week_qty": 91,
      }
    """
    if sku not in erp_by_sku or sku not in sales_by_sku:
        return []
    ep = erp_by_sku[sku]
    sc = sales_by_sku[sku]
    out: list[dict] = []
    for ptype, grp in ep.groupby("promo_types"):
        yws = set(zip(grp["year"].astype(int), grp["week"].astype(int)))
        # Filter sales to weeks belonging to this event.
        in_ev = sc[sc.apply(lambda r, _yws=yws:
                             (int(r["year"]), int(r["week"])) in _yws, axis=1)]
        if in_ev.empty:
            continue
        wks = sorted(zip(in_ev["year"].astype(int), in_ev["week"].astype(int)))
        out.append({
            "promo_types":    ptype,
            "year":           wks[0][0],
            "start_week":     wks[0][1],
            "end_week":       wks[-1][1],
            "n_weeks":        len(in_ev),
            "total_qty":      float(in_ev["_rcm_web"].sum()),
            "avg_per_week":   float(in_ev["_rcm_web"].mean()),
            "peak_week_qty":  float(in_ev["_rcm_web"].max()),
        })
    # Sort most-recent first by start (year, week).
    out.sort(key=lambda e: (e["year"], e["start_week"]), reverse=True)
    return out


def median(xs: list[float]) -> float:
    if not xs:
        return 0.0
    s = sorted(xs)
    n = len(s)
    return s[n // 2] if n % 2 else (s[n // 2 - 1] + s[n // 2]) / 2


def recent_baseline(sku: str, lookback: int = 13, n_clean: int = 4) -> float:
    """Same logic as PromoTool's base_run_rate — average of last `n_clean`
    non-promo, non-zero weeks within the past `lookback` weeks."""
    if sku not in sales_by_sku:
        return 0.0
    sc = sales_by_sku[sku].sort_values(["year", "week"], ascending=False).head(lookback)
    promo_yws = set()
    if sku in erp_by_sku:
        ep = erp_by_sku[sku]
        promo_yws = set(zip(ep["year"].astype(int), ep["week"].astype(int)))
    chosen = []
    for _, r in sc.iterrows():
        y, w = int(r["year"]), int(r["week"])
        if (y, w) in promo_yws:
            continue
        q = float(r["_rcm_web"])
        if q <= 0:
            continue
        chosen.append(q)
        if len(chosen) >= n_clean:
            break
    return (sum(chosen) / len(chosen)) if chosen else 0.0


# --------------------------------------------------------------------------
# Per-SKU calc
# --------------------------------------------------------------------------

def forecast_from_actuals(sku: str) -> dict:
    """Return forecast run-rate (per-week) for the SKU plus diagnostics."""
    events = past_promo_events(sku)
    if events:
        med_rr = median([e["avg_per_week"] for e in events])
        max_rr = max((e["peak_week_qty"] for e in events), default=0.0)
        return {
            "source":          "own",
            "n_events":        len(events),
            "median_run_rate": med_rr,
            "max_week_qty":    max_rr,
            "events_summary":  "; ".join(
                f"{e['promo_types']} CW{e['start_week']}-CW{e['end_week']}/{e['year']} "
                f"avg {e['avg_per_week']:.0f} u/wk (peak {e['peak_week_qty']:.0f})"
                for e in events[:4]
            ),
        }

    # Category-sibling fallback — siblings' median promo run-rate scaled
    # by ratio of this SKU's baseline to siblings' baseline (so a small
    # SKU doesn't inherit a big sibling's volume).
    cat = cat_map.get(sku, "")
    siblings = [s for s in cat_to_skus.get(cat, []) if s != sku]
    sib_rates: list[float] = []
    sib_baselines: list[float] = []
    for s in siblings[:50]:
        evs = past_promo_events(s)
        if not evs:
            continue
        sib_rates.append(median([e["avg_per_week"] for e in evs]))
        sib_baselines.append(recent_baseline(s))
    if sib_rates:
        sib_med_rr  = median(sib_rates)
        sib_med_bas = median([b for b in sib_baselines if b > 0]) or 1.0
        own_bas     = recent_baseline(sku)
        scale = (own_bas / sib_med_bas) if sib_med_bas > 0 else 1.0
        return {
            "source":          f"sibling-cat ({cat}, n={len(sib_rates)})",
            "n_events":        0,
            "median_run_rate": sib_med_rr * scale,
            "max_week_qty":    sib_med_rr * scale * 1.5,   # rough peak heuristic
            "events_summary":  f"no own promo — sibling median {sib_med_rr:.0f} u/wk × scale {scale:.2f}",
        }

    # Last-resort: capped uplift on current baseline.
    own_bas = recent_baseline(sku)
    return {
        "source":          f"capped-{CAP_UPLIFT:g}x-baseline",
        "n_events":        0,
        "median_run_rate": own_bas * CAP_UPLIFT,
        "max_week_qty":    own_bas * CAP_UPLIFT * 1.5,
        "events_summary":  f"no own/sibling promo history — baseline {own_bas:.1f} × {CAP_UPLIFT:g}",
    }


def compute_sku(sku: str, policy: dict | None) -> dict:
    name  = name_map.get(sku, "") or (policy.get("policy_sku_name") if policy else "") or ""
    cat   = cat_map.get(sku, "")
    tier  = tier_map.get(sku, "")
    price = float(price_map.get(sku, 0) or 0)
    cost  = float(cost_map.get(sku, 0) or 0)
    missing_master: list[str] = []
    if sku not in name_map:  missing_master.append("catalog")
    if price <= 0:           missing_master.append("price")
    if cost  <= 0:           missing_master.append("cost")

    row: dict = {
        "SKU": sku,
        "Naziv": name,
        "Kategorija": cat,
        "Tier": tier,
        "Period": f"W{PROMO_WEEKS[0]:02d}-W{PROMO_WEEKS[-1]:02d} {PROMO_YEAR}",
        "Trajanje (tjedni)": N_WEEKS,
        "Normal retail price (EUR)": round(price, 2) if price else None,
        "Cost (EUR)": round(cost, 2) if cost else None,
    }
    if policy is None:
        row["Note"] = "SKU not in either rabatna politika file"
        row["Master data missing"] = ",".join(missing_master)
        return row

    discount = policy.get("discount_pct")
    cijena_pol = policy.get("promo_price")
    if cijena_pol and cijena_pol > 0:
        promo_price = float(cijena_pol)
        eff_disc = ((price - promo_price) / price * 100) if price > 0 else (discount or 0)
    else:
        eff_disc = float(discount or 0)
        promo_price = price * (1 - eff_disc / 100) if price > 0 else 0.0

    row.update({
        "Rabatna politika":     policy.get("policy_name"),
        "Policy izvor":         policy.get("policy_source"),
        "Rabat % (policy)":     discount,
        "Cijena (policy EUR)":  cijena_pol,
        "Effective discount %": round(eff_disc, 2),
        "Promo price (EUR)":    round(promo_price, 4),
        "Min qty (Količina >)": policy.get("min_qty"),
    })

    base_avg = recent_baseline(sku)
    fc = forecast_from_actuals(sku)

    promo_qty_median = fc["median_run_rate"] * N_WEEKS
    promo_qty_max    = fc["max_week_qty"]    * N_WEEKS
    base_qty         = base_avg * N_WEEKS
    incremental      = promo_qty_median - base_qty

    promo_ruc_unit = max(0.0, promo_price - cost)
    promo_revenue  = promo_qty_median * promo_price
    promo_revenue_max = promo_qty_max * promo_price
    promo_ruc      = promo_qty_median * promo_ruc_unit
    promo_ruc_max  = promo_qty_max    * promo_ruc_unit
    base_revenue   = base_qty * price
    base_ruc       = base_qty * max(0.0, price - cost)
    rev_delta      = promo_revenue - base_revenue
    ruc_delta      = promo_ruc - base_ruc

    # Implied uplift (just for reference — not used to drive the forecast).
    implied_uplift = (fc["median_run_rate"] / base_avg) if base_avg > 0 else None
    margin_pct = ((promo_price - cost) / promo_price * 100) if promo_price > 0 else None

    row.update({
        "Forecast source":              fc["source"],
        "Past promo events n":          fc["n_events"],
        "Past promo events summary":    fc["events_summary"],
        "Baseline / tjedan (RCM+WEB)":  round(base_avg, 1),
        "Promo run-rate median (u/wk)": round(fc["median_run_rate"], 1),
        "Promo peak week qty (u)":      round(fc["max_week_qty"], 1),
        "Implied uplift (×)":           round(implied_uplift, 2) if implied_uplift is not None else None,

        "Baseline units (no promo)":    int(round(base_qty)),
        "Forecast units (median)":      int(round(promo_qty_median)),
        "Max units (peak-week × 5)":    int(round(promo_qty_max)),
        "Incremental units":            int(round(incremental)),

        "Promo revenue (EUR)":          round(promo_revenue, 2),
        "Max revenue (EUR)":            round(promo_revenue_max, 2),
        "Revenue delta vs no-promo":    round(rev_delta, 2),
        "Promo unit RUC (EUR)":         round(promo_ruc_unit, 4),
        "Promo RUC total (EUR)":        round(promo_ruc, 2),
        "Max RUC (EUR)":                round(promo_ruc_max, 2),
        "RUC delta vs no-promo":        round(ruc_delta, 2),
        "Promo margin %":               round(margin_pct, 2) if margin_pct is not None else None,
        "Master data missing":          ",".join(missing_master),
    })
    notes = []
    if base_avg <= 0:
        notes.append("Baseline=0 (no clean non-promo weeks)")
    if fc["source"].startswith("capped"):
        notes.append("Fallback: capped uplift × baseline")
    elif fc["source"].startswith("sibling"):
        notes.append("Fallback: category siblings' promo rate")
    if missing_master:
        notes.append("Master data missing: " + ",".join(missing_master))
    row["Note"] = " · ".join(notes) if notes else ""
    return row


# --------------------------------------------------------------------------
# CLI
# --------------------------------------------------------------------------

def _read_sku_list(path: Path) -> list[str]:
    raw = path.read_text(encoding="utf-8-sig").splitlines()
    out: list[str] = []
    seen: set[str] = set()
    for line in raw:
        s = line.strip().lstrip("﻿").strip()
        if not s or s.startswith("#") or s.lower() == "sku":
            continue
        sku = s.split(",", 1)[0].strip().strip('"').strip("'")
        if sku and sku not in seen:
            seen.add(sku)
            out.append(sku)
    return out


def main() -> None:
    ap = argparse.ArgumentParser(description="Promo calc W22–W26 — actuals-based.")
    ap.add_argument("--skus", type=Path, required=True,
                    help="Text/CSV file with one SKU per line.")
    ap.add_argument("--out",  type=Path, default=None,
                    help="Output xlsx path. Defaults to "
                         "data/promo_w22_w26_actual_<stem>.xlsx.")
    args = ap.parse_args()

    skus = _read_sku_list(args.skus)
    out_path = args.out or DATA_DIR / f"promo_w22_w26_actual_{args.skus.stem}.xlsx"

    print(f"Loading policies…")
    pmap = load_policy_map()
    print(f"  {len(pmap)} SKUs in combined policy map")
    print(f"\nComputing actuals-based promo metrics for {len(skus)} SKUs, "
          f"period W{PROMO_WEEKS[0]}–W{PROMO_WEEKS[-1]} {PROMO_YEAR}…")

    rows = []
    missing = []
    for sku in skus:
        pol = pmap.get(sku)
        rows.append(compute_sku(sku, pol))
        if pol is None:
            missing.append(sku)

    df = pd.DataFrame(rows)
    # Stable column order — diagnostics last.
    cols = [
        "SKU", "Naziv", "Kategorija", "Tier", "Period", "Trajanje (tjedni)",
        "Rabatna politika", "Policy izvor",
        "Rabat % (policy)", "Cijena (policy EUR)", "Effective discount %",
        "Normal retail price (EUR)", "Promo price (EUR)", "Cost (EUR)",
        "Min qty (Količina >)",
        "Forecast source", "Past promo events n",
        "Baseline / tjedan (RCM+WEB)",
        "Promo run-rate median (u/wk)", "Promo peak week qty (u)",
        "Implied uplift (×)",
        "Baseline units (no promo)",
        "Forecast units (median)", "Max units (peak-week × 5)",
        "Incremental units",
        "Promo revenue (EUR)", "Max revenue (EUR)", "Revenue delta vs no-promo",
        "Promo unit RUC (EUR)", "Promo RUC total (EUR)",
        "Max RUC (EUR)", "RUC delta vs no-promo",
        "Promo margin %",
        "Past promo events summary",
        "Master data missing", "Note",
    ]
    df = df.reindex(columns=cols)

    out_path.parent.mkdir(parents=True, exist_ok=True)
    with pd.ExcelWriter(out_path, engine="openpyxl") as wr:
        df.to_excel(wr, sheet_name="Promo_W22_W26_actual", index=False)
        ws = wr.sheets["Promo_W22_W26_actual"]
        for i, col in enumerate(df.columns, 1):
            ml = max([len(str(col))] +
                     [len(str(v)) for v in df[col].astype(str).head(200)])
            ws.column_dimensions[ws.cell(row=1, column=i).column_letter].width = min(ml + 2, 50)

    print(f"\n✓ wrote {out_path} — {len(df)} rows")
    if missing:
        print(f"\n[warn] {len(missing)} SKU(s) not found in either policy file:")
        for s in missing:
            print(f"  - {s}  ({name_map.get(s, '?')})")

    # Quick portfolio summary
    print("\n=== Portfolio totals ===")
    print(f"  Forecast units total:  {df['Forecast units (median)'].sum():,.0f}")
    print(f"  Max units total:       {df['Max units (peak-week × 5)'].sum():,.0f}")
    print(f"  Promo revenue total:   €{df['Promo revenue (EUR)'].sum():,.0f}")
    print(f"  Promo RUC total:       €{df['Promo RUC total (EUR)'].sum():,.0f}")
    src_counts = df["Forecast source"].value_counts().to_dict()
    print("\n  Forecast source breakdown:")
    for s, n in src_counts.items():
        print(f"    {s:<40s}  {n}")


if __name__ == "__main__":
    main()
