"""Build a per-SKU promo P&L for W22–W26 2026.

Inputs
------
* SKU list (`SKUS` constant, taken from the user's screenshot).
* Rabatne politike workbooks:
    1. ``data/VrstaRabatnePolitike06.xlsx`` — primary source
    2. ``data/VrstaRabatnePolitikesummer buddy.xlsx`` — fallback for SKUs
       missing from the primary policy
* PromoTool's own helpers (``promo_data.base_run_rate`` and
  ``promo_data.suggest_uplift``) so the median uplift, p90 upside and
  baseline run-rate match exactly what the Promo Tool would show
  interactively.

Output
------
``data/promo_w22_w26_per_sku.xlsx`` — one row per SKU with columns
matching the metrics surfaced by the Promo Tool ("Total promo qty",
"Realistic upside (p90)", promo revenue + RUC, etc.).
"""
from __future__ import annotations

import argparse
import os
import sys
from pathlib import Path

# PromoTool's helpers import `streamlit` which is fine in this script — we
# just call functions; no UI is rendered.  Make PromoTool importable as a
# top-level package by tweaking sys.path before the import.
ROOT = Path(__file__).resolve().parent
PROMO_TOOL_DIR = ROOT / "PromoTool"
sys.path.insert(0, str(PROMO_TOOL_DIR))

# Force UTF-8 stdout so Croatian characters print cleanly on Windows.
sys.stdout.reconfigure(encoding="utf-8")  # type: ignore[attr-defined]

import pandas as pd

# promo_data is module-level so streamlit caches don't matter here.
from promo_data import (                 # noqa: E402  (intentional ordering)
    base_run_rate,
    suggest_uplift,
)

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

PROMO_YEAR     = 2026
PROMO_WEEKS    = [22, 23, 24, 25, 26]    # W22–W26 inclusive  (5 weeks)
N_WEEKS        = len(PROMO_WEEKS)

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

# Exact SKU list from the screenshot, in the order shown.
SKUS = [
    "POL09890", "POL09891", "POL09739", "POL09892", "POL09736",
    "POL09829", "POL09734", "POL09738", "POL09735", "POL09893",
    "POL09894", "POL12887",
    "ZOE12362", "ZOE12365", "ZOE12366", "ZOE12361", "ZOE12405",
    "CEL11618", "CEL11619", "CEL11616", "CEL11615",
    "POL12768", "POL09848",
    "POL12860", "POL12870", "POL12871", "POL12872",
    "POL09857", "POL09792",
    "ZOE12408", "ZOE12382", "ZOE12377", "ZOE12397",   # ZOE12397 = "ZOE Boss Reds 250g Raspberry"
    "POL09929", "POL09928", "POL09927",
    "ZOE12395",
    "POL04462", "POL12741", "POL04473", "POL04475", "POL04474",
    "ZOE12453", "ZOE12392", "ZOE12391",
    "POL12921", "POL12920", "POL09895",
    "POL12853", "POL12854",
    "ZOE12386", "ZOE12387", "ZOE12400",
    "POL12866", "POL12899",
    "ZOE12411",
    "POL04375", "POL04376", "POL09714", "POL04373", "POL09711",
]

# --------------------------------------------------------------------------
# Policy loader — combine the two workbooks, primary wins on SKU collisions.
# --------------------------------------------------------------------------

def _read_policy(path: Path) -> pd.DataFrame:
    """Return a single dataframe of policy *items* (header rows dropped)."""
    if not path.exists():
        print(f"[warn] policy file missing: {path}")
        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]:
    """SKU → {discount_pct, promo_price, min_qty, vp_margin_pct, vp_margin_eur,
                mp_margin_pct, mp_margin_eur, valid_from, valid_to, source}.
    Primary file wins on collisions (06/26 promo takes precedence over
    summer-buddy).
    """
    pieces = [_read_policy(POLICY_PRIMARY), _read_policy(POLICY_FALLBK)]
    combined = pd.concat([p for p in pieces if not p.empty], ignore_index=True)
    if combined.empty:
        return {}

    out: dict[str, dict] = {}
    for _, r in combined.iterrows():
        sku = r["Artikal"]
        # Primary appended first → keep the first occurrence we see.
        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 >"),
            "vp_margin_pct":    _f("VP Marža %"),
            "vp_margin_eur":    _f("VP Marža"),
            "mp_margin_pct":    _f("MP Marža %"),
            "mp_margin_eur":    _f("MP Marža"),
            "valid_from":       r.get("Trajanje od"),
            "valid_to":         r.get("Trajanje do"),
            "policy_name":      r.get("Rabatna politika"),
            "policy_source":    r.get("_source"),
            "policy_sku_name":  r.get("Naziv artikla"),
        }
    return out


# --------------------------------------------------------------------------
# Reference data — sales / prices / costs / plan list / ERP promo calendar.
# Same files PromoTool reads under DATA_DIR (../data/ from the tool).
# --------------------------------------------------------------------------

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


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

# Lookup maps — same logic as promo_data.load_all().
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 {}

# Retail PPP — matches PromoTool: discount applied off normal retail price.
price_map: dict[str, float] = {}
if "sku" in prices.columns and "normal_retail_ppp" in prices.columns:
    price_map = dict(zip(prices["sku"], prices["normal_retail_ppp"]))
cost_map: dict[str, float] = {}
if "sku" in costs.columns and "cost_price" in costs.columns:
    cost_map = dict(zip(costs["sku"], costs["cost_price"]))


# --------------------------------------------------------------------------
# Per-SKU calculation
# --------------------------------------------------------------------------

def compute_sku(sku: str, policy: dict | None) -> dict:
    """Build the per-SKU row mirroring Promo Tool's headline numbers."""
    name  = name_map.get(sku, "")
    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)

    # Fallback name from the policy when the SKU isn't in sku_plan_list.csv
    # (new brand introductions aren't always onboarded into the master yet).
    if not name and policy is not None:
        name = str(policy.get("policy_sku_name") or "")
    # Master-data coverage flags so reviewers know which numbers come from a
    # complete master record vs. a partial one.
    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")

    # Result skeleton — any branch can stuff fields and `return result`.
    result: 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:
        result["Note"] = "SKU not in either rabatna politika file"
        return result

    discount = policy.get("discount_pct")
    cijena_policy = policy.get("promo_price")
    # Effective promo price:
    #  * Use Cijena from the policy if it's set (some campaigns specify
    #    the absolute promo price directly).
    #  * Otherwise apply the Rabat % to the normal retail price.
    if cijena_policy and cijena_policy > 0:
        promo_price = float(cijena_policy)
        # Back out effective discount for reporting.
        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

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

    # ---- Baseline run-rate (same fn the Promo Tool uses) ----
    rr = base_run_rate(sales, erp_promo, sku)
    base_avg  = float(rr.get("avg") or 0)
    n_clean_w = int(rr.get("weeks") or 0)

    # ---- Predicted uplift (median + p90 upside) ----
    upl = suggest_uplift(sales, erp_promo, sku, eff_disc)
    uplift_med = float(upl.get("uplift") or 0)
    uplift_p90 = float(upl.get("uplift_p90") or 0)
    n_history  = int(upl.get("n_history") or 0)

    # ---- Volume forecasts ----
    # Median forecast — same formula as Forecaster page (line 269):
    #   promo_qty = base_avg * n_weeks * predicted_uplift
    promo_qty_median = base_avg * N_WEEKS * uplift_med
    promo_qty_max    = base_avg * N_WEEKS * uplift_p90
    base_qty         = base_avg * N_WEEKS
    incremental      = promo_qty_median - base_qty

    # ---- P&L ----
    promo_revenue = promo_qty_median * promo_price
    base_revenue  = base_qty * price
    rev_delta     = promo_revenue - base_revenue

    promo_ruc_unit = max(0.0, promo_price - cost)
    base_ruc_unit  = max(0.0, price       - cost)
    promo_ruc      = promo_qty_median * promo_ruc_unit
    base_ruc       = base_qty         * base_ruc_unit
    ruc_delta      = promo_ruc - base_ruc

    # Max-quantity scenario P&L (p90 upside)
    promo_ruc_max     = promo_qty_max * promo_ruc_unit
    promo_revenue_max = promo_qty_max * promo_price

    margin_pct = ((promo_price - cost) / promo_price * 100) if promo_price > 0 else None

    result.update({
        "Baseline / tjedan (RCM+WEB)": round(base_avg, 1),
        "Baseline non-promo tjedana":  n_clean_w,
        "Promo history n":             n_history,
        "Predicted uplift (×)":        round(uplift_med, 2),
        "Realistic upside p90 (×)":    round(uplift_p90, 2),

        "Forecast units (median)":  int(round(promo_qty_median)),
        "Max units (p90 upside)":   int(round(promo_qty_max)),
        "Baseline units (no promo)": int(round(base_qty)),
        "Incremental units":         int(round(incremental)),

        "Promo revenue (EUR)":      round(promo_revenue, 2),
        "Max revenue p90 (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 p90 (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,
    })

    if base_avg <= 0:
        result["Note"] = "Baseline = 0 — nema clean non-promo tjedana zadnjih 13w"
    if missing_master:
        prefix = (result.get("Note") + " · ") if result.get("Note") else ""
        result["Note"] = prefix + "Master-data missing: " + ",".join(missing_master)
    result["Master data missing"] = ",".join(missing_master) if missing_master else ""

    return result


# --------------------------------------------------------------------------
# Driver
# --------------------------------------------------------------------------

def _read_sku_list(path: Path) -> list[str]:
    """One SKU per line, accepts either a bare list or a CSV with 'SKU' column.
    Strips BOM, blank lines and lines starting with '#'.
    """
    raw = path.read_text(encoding="utf-8-sig").splitlines()
    out: list[str] = []
    seen: set[str] = set()
    header_skipped = False
    for line in raw:
        s = line.strip().lstrip("﻿").strip()
        if not s or s.startswith("#"):
            continue
        # If the file looks like a CSV, drop the header row (first 'SKU' or 'sku')
        if not header_skipped and s.lower() == "sku":
            header_skipped = True
            continue
        # Tolerate "SKU,extra,cols" — take first column.
        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 per SKU.")
    ap.add_argument("--skus", type=Path, default=None,
                    help="Path to a text/CSV file with one SKU per line. "
                         "When omitted, uses the built-in SKUS list.")
    ap.add_argument("--out", type=Path, default=None,
                    help="Output xlsx path. Defaults to "
                         "data/promo_w22_w26_per_sku.xlsx, or "
                         "data/promo_w22_w26_<stem>.xlsx when --skus is set.")
    args = ap.parse_args()

    if args.skus is not None:
        skus = _read_sku_list(args.skus)
        out_path = args.out or DATA_DIR / f"promo_w22_w26_{args.skus.stem}.xlsx"
    else:
        skus = SKUS
        out_path = args.out or OUT_FILE

    print(f"Loading policies from:\n  {POLICY_PRIMARY.name}\n  {POLICY_FALLBK.name}")
    policy_map = load_policy_map()
    print(f"  → {len(policy_map)} SKUs in combined policy map")

    print(f"\nComputing promo metrics for {len(skus)} SKUs, "
          f"period W{PROMO_WEEKS[0]}–W{PROMO_WEEKS[-1]} {PROMO_YEAR}…")
    rows: list[dict] = []
    missing: list[str] = []
    for sku in skus:
        pol = policy_map.get(sku)
        rows.append(compute_sku(sku, pol))
        if pol is None:
            missing.append(sku)

    df = pd.DataFrame(rows)
    # Sort columns: identity first, then policy, then forecast, then P&L.
    col_order = [
        "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 >)",
        "Baseline / tjedan (RCM+WEB)", "Baseline non-promo tjedana",
        "Promo history n",
        "Predicted uplift (×)", "Realistic upside p90 (×)",
        "Baseline units (no promo)",
        "Forecast units (median)", "Max units (p90 upside)",
        "Incremental units",
        "Promo revenue (EUR)", "Max revenue p90 (EUR)",
        "Revenue delta vs no-promo",
        "Promo unit RUC (EUR)", "Promo RUC total (EUR)",
        "Max RUC p90 (EUR)", "RUC delta vs no-promo",
        "Promo margin %",
        "Master data missing",
        "Note",
    ]
    df = df.reindex(columns=col_order)

    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", index=False)
        # Auto-fit columns (approximate).
        ws = wr.sheets["Promo_W22_W26"]
        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, 40)

    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, '?')})")


if __name__ == "__main__":
    main()
