"""
ABC 3-scenario PO postpone/cancel analysis using REAL weekly demand
(walk-forward weeks-of-cover, not 8w-avg approximation).

Outputs:
  three_scenarios_per_po.csv
  weekly_inventory_projection.csv
  scenario_summary.json
  abc_cancel_list_scenario_b.csv
  abc_postpone_list_scenario_b.csv
"""
import json
from collections import defaultdict

import numpy as np
import pandas as pd

# ----- constants -----
CURRENT_YEAR = 2026
CURRENT_WEEK = 20
HORIZON_WEEKS = 13                              # CW20..CW32 inclusive
PROJECTION_HORIZON = 14                         # for cash projection CSV: CW20..CW33
PO_WINDOW = [(2026, w) for w in range(22, 28)]
POSTPONE_DELAY = 4
CANCEL_THRESHOLDS = [None, 26, 36]              # A, B, C
POSTPONE_COVER_TRIGGER = 4
CURRENT_STOCK_EUR = 4_616_194

# ----- helpers -----
def yw(y, w):
    return y * 100 + w

def add_weeks(y, w, n):
    w += n
    while w > 52:
        w -= 52
        y += 1
    return y, w

def weeks_after(y, w, n):
    out = []
    cy, cw = y, w
    for _ in range(n):
        out.append((cy, cw))
        cy, cw = add_weeks(cy, cw, 1)
    return out

# ----- load -----
inc = pd.read_csv("data/incoming_supply.csv")
sup = pd.read_csv("data/supply_master.csv")
sup["is_abc"] = sup["supplier"].fillna("").str.contains("ABC NUTRITIONAL", case=False)
stock_wh = pd.read_csv("data/stock.csv")
stores_frames = []
for f in ["stock_stores.csv", "stock_stores_at.csv", "stock_stores_slo.csv"]:
    df = pd.read_csv(f"data/{f}")
    df.columns = [c.lower() for c in df.columns]
    stores_frames.append(df[["sku", "on_hand"]])
stores_df = pd.concat(stores_frames, ignore_index=True).groupby("sku", as_index=False)["on_hand"].sum()
fc = pd.read_csv("data/forecast_for_supply.csv")
costs = pd.read_csv("data/sku_costs.csv")
plan = pd.read_csv("data/sku_plan_list.csv")
plan.columns = [c.lower() for c in plan.columns]
plan = plan.rename(columns={"cat": "category", "oznaka": "tier"})

cost_map = dict(zip(costs["sku"], costs["cost_price"]))
wh_map = dict(zip(stock_wh["sku"], stock_wh["on_hand"]))
stores_map = dict(zip(stores_df["sku"], stores_df["on_hand"]))

fc["yw"] = fc["year"] * 100 + fc["week"]
fc_lookup = {(r.sku, r.yw): r.demand for r in fc.itertuples()}

horizon = weeks_after(CURRENT_YEAR, CURRENT_WEEK, HORIZON_WEEKS)
horizon_labels = [f"CW{w}" for (y, w) in horizon]

proj_weeks = weeks_after(CURRENT_YEAR, CURRENT_WEEK, PROJECTION_HORIZON)
proj_labels = [f"CW{w}" for (y, w) in proj_weeks]
proj_yw_set = {yw(y, w) for (y, w) in proj_weeks}

# ----- cover math -----
def demand_series(sku, start_y, start_w, n=HORIZON_WEEKS):
    return [fc_lookup.get((sku, yw(y, w)), 0) for (y, w) in weeks_after(start_y, start_w, n)]

def real_weeks_cover(stock_now, ds):
    """Walk-forward weeks of cover. Handles peaks correctly."""
    remaining = stock_now
    weeks = 0.0
    for d in ds:
        if d <= 0:
            weeks += 1
            continue
        if remaining >= d:
            remaining -= d
            weeks += 1
        else:
            weeks += remaining / d
            return weeks
    # exhausted horizon with stock left -> extrapolate
    avg = sum(ds) / max(len(ds), 1)
    if avg <= 0:
        return weeks + 999
    return weeks + remaining / avg

def avg_weeks_cover(stock_now, sku, start_y, start_w, n=8):
    """8w avg-based cover, used only for the movers diff."""
    ds = demand_series(sku, start_y, start_w, n)
    avg = sum(ds) / max(len(ds), 1)
    return 999 if avg <= 0 else stock_now / avg

def classify(real_cover_now, real_post, cancel_threshold):
    if cancel_threshold is not None and real_post > cancel_threshold:
        return "CANCEL"
    if real_cover_now >= POSTPONE_COVER_TRIGGER:
        return "POSTPONE"
    if real_cover_now >= 2:
        return "REVIEW"
    return "PRODUCE"

# ----- ABC POs in window -----
inc["yw"] = inc["year"] * 100 + inc["week"]
inc_m = inc.merge(sup[["sku", "is_abc"]], on="sku", how="left")
inc_m["is_abc"] = inc_m["is_abc"].fillna(False).astype(bool)
window_yw_set = {yw(y, w) for (y, w) in PO_WINDOW}
abc_window = inc_m[inc_m["is_abc"] & inc_m["yw"].isin(window_yw_set)].copy()
plan_min = plan[["sku", "category", "tier"]].copy() if "category" in plan.columns else None

# ----- per-PO classification -----
rows = []
for _, po in abc_window.iterrows():
    sku = po["sku"]
    qty = float(po["qty"])
    py, pw = int(po["year"]), int(po["week"])

    wh_now = float(wh_map.get(sku, 0))
    stores_now = float(stores_map.get(sku, 0))
    combined_now = wh_now + stores_now

    ds_now = demand_series(sku, CURRENT_YEAR, CURRENT_WEEK)
    real_cov_wh = real_weeks_cover(wh_now, ds_now)
    real_cov_combined = real_weeks_cover(combined_now, ds_now)

    ds_delivery = demand_series(sku, py, pw)
    real_post = real_weeks_cover(wh_now + qty, ds_delivery)

    avg_cov = avg_weeks_cover(wh_now, sku, CURRENT_YEAR, CURRENT_WEEK)
    avg_post = avg_weeks_cover(wh_now + qty, sku, py, pw)

    avg_d = float(np.mean(ds_now)) if ds_now else 0.0
    max_d = float(np.max(ds_now)) if ds_now else 0.0
    peak_idx = int(np.argmax(ds_now)) if max_d > 0 else 0
    peak_lab = horizon_labels[peak_idx]

    scen, new_w, avg_scen = {}, {}, {}
    for letter, thresh in zip(["a", "b", "c"], CANCEL_THRESHOLDS):
        a = classify(real_cov_wh, real_post, thresh)
        scen[letter] = a
        if a == "POSTPONE":
            ny, nw = add_weeks(py, pw, POSTPONE_DELAY)
            new_w[letter] = f"CW{nw}/{ny}"
        else:
            new_w[letter] = ""
        avg_scen[letter] = classify(avg_cov, avg_post, thresh)

    tier, cat = "", ""
    if plan_min is not None:
        pr = plan_min[plan_min["sku"] == sku]
        if not pr.empty:
            tier = pr["tier"].iloc[0]
            cat = pr["category"].iloc[0]

    rows.append({
        "sku": sku, "tier": tier, "category": cat,
        "po_year": py, "po_week": pw, "po_label": f"CW{pw}",
        "qty": int(qty),
        "cost_price": round(cost_map.get(sku, 0), 4),
        "eur_value": round(qty * cost_map.get(sku, 0), 2),
        "wh_stock_now": int(wh_now), "stores_stock_now": int(stores_now),
        "combined_stock_now": int(combined_now),
        "real_weeks_cover_wh": round(real_cov_wh, 2),
        "real_weeks_cover_combined": round(real_cov_combined, 2),
        "real_post_delivery_cover": round(real_post, 2),
        "avg_weekly_demand_13w": round(avg_d, 2),
        "max_weekly_demand_13w": round(max_d, 2),
        "peak_demand_week": peak_lab,
        "scenario_a": scen["a"], "scenario_b": scen["b"], "scenario_c": scen["c"],
        "new_delivery_week_a": new_w["a"],
        "new_delivery_week_b": new_w["b"],
        "new_delivery_week_c": new_w["c"],
        "_avg_a": avg_scen["a"], "_avg_b": avg_scen["b"], "_avg_c": avg_scen["c"],
    })

per_po = pd.DataFrame(rows)

# ----- multi-PO sanity check per scenario -----
# For each SKU with 2+ POSTPONE in window: simulate stacked postponements
# (each PO shifted +4w). Require >= 1 week real cover (stock >= next week's
# demand) at every week through CW32. If violated, un-postpone the
# latest-week POSTPONE, re-check.
abc_skus = set(abc_window["sku"])
sku_all_pos = {sku: inc_m[inc_m["sku"] == sku].copy() for sku in abc_skus}

def simulate_sku(sku, action_map_for_window):
    """Walk WH-only stock. action_map_for_window: {(year, week): action} for window POs.
    Out-of-window POs of this SKU flow as-is. Returns list of (year, week, closing_stock).
    """
    s = float(wh_map.get(sku, 0))
    inflows = defaultdict(float)
    for _, po in sku_all_pos[sku].iterrows():
        py, pw, q = int(po["year"]), int(po["week"]), float(po["qty"])
        key = (py, pw)
        if key in action_map_for_window:
            a = action_map_for_window[key]
            if a == "CANCEL":
                continue
            if a == "POSTPONE":
                ny, nw = add_weeks(py, pw, POSTPONE_DELAY)
                inflows[(ny, nw)] += q
            else:
                inflows[key] += q
        else:
            inflows[key] += q
    out = []
    for (y, w) in weeks_after(CURRENT_YEAR, CURRENT_WEEK, HORIZON_WEEKS):
        s = max(0.0, s + inflows.get((y, w), 0) - fc_lookup.get((sku, yw(y, w)), 0))
        out.append((y, w, s))
    return out

def has_violations(sku, stocks):
    for i, (y, w, s) in enumerate(stocks[:-1]):
        ny, nw, _ = stocks[i + 1]
        next_d = fc_lookup.get((sku, yw(ny, nw)), 0)
        if next_d > 0 and s < next_d:
            return True
    return False

sanity_unpostpones = defaultdict(list)  # scenario -> [(sku, po_week)]
for letter in ["a", "b", "c"]:
    col = f"scenario_{letter}"
    new_col = f"new_delivery_week_{letter}"
    for sku, g in per_po.groupby("sku"):
        postponed_rows = g[g[col] == "POSTPONE"]
        if len(postponed_rows) < 2:
            continue
        action_map = {(int(r["po_year"]), int(r["po_week"])): r[col] for _, r in g.iterrows()}
        stocks = simulate_sku(sku, action_map)
        while has_violations(sku, stocks):
            postponed = [(y, w) for (y, w), a in action_map.items() if a == "POSTPONE"]
            if not postponed:
                break
            latest = max(postponed)
            action_map[latest] = "PRODUCE"
            sanity_unpostpones[letter].append(f"{sku}@CW{latest[1]}")
            mask = (per_po["sku"] == sku) & (per_po["po_year"] == latest[0]) & (per_po["po_week"] == latest[1])
            per_po.loc[mask, col] = "PRODUCE"
            per_po.loc[mask, new_col] = ""
            stocks = simulate_sku(sku, action_map)

# ----- cash projection -----
def build_cash_projection(scenario):
    """scenario: 'baseline' | 'a' | 'b' | 'c'. Returns list of weekly closing EUR."""
    inflow_eur = defaultdict(float)
    action_lookup = {}
    if scenario != "baseline":
        col = f"scenario_{scenario}"
        for _, r in per_po.iterrows():
            action_lookup[(r["sku"], int(r["po_year"]), int(r["po_week"]))] = r[col]
    for _, po in inc_m.iterrows():
        sku = po["sku"]
        py, pw, q = int(po["year"]), int(po["week"]), float(po["qty"])
        cost = cost_map.get(sku, 0)
        eur = q * cost
        if scenario == "baseline" or not po["is_abc"] or (py, pw) not in window_yw_set:
            inflow_eur[(py, pw)] += eur
            continue
        a = action_lookup.get((sku, py, pw), "PRODUCE")
        if a == "CANCEL":
            continue
        if a == "POSTPONE":
            ny, nw = add_weeks(py, pw, POSTPONE_DELAY)
            inflow_eur[(ny, nw)] += eur
        else:
            inflow_eur[(py, pw)] += eur

    outflow_eur = defaultdict(float)
    for (sku, ywv), d in fc_lookup.items():
        if ywv not in proj_yw_set:
            continue
        outflow_eur[(ywv // 100, ywv % 100)] += d * cost_map.get(sku, 0)

    s = float(CURRENT_STOCK_EUR)
    series = []
    for (y, w) in proj_weeks:
        s = max(0.0, s + inflow_eur.get((y, w), 0) - outflow_eur.get((y, w), 0))
        series.append(s)
    return series

baseline_series = build_cash_projection("baseline")
scen_series = {letter: build_cash_projection(letter) for letter in ["a", "b", "c"]}

# ----- outputs -----
# File 1
per_po.drop(columns=["_avg_a", "_avg_b", "_avg_c"]).to_csv("three_scenarios_per_po.csv", index=False)

# File 2
proj_df = pd.DataFrame({
    "week_label": proj_labels,
    "week_iso": [f"{y}-W{w:02d}" for (y, w) in proj_weeks],
    "baseline_eur": [round(v) for v in baseline_series],
    "scenario_a_eur": [round(v) for v in scen_series["a"]],
    "scenario_b_eur": [round(v) for v in scen_series["b"]],
    "scenario_c_eur": [round(v) for v in scen_series["c"]],
})
proj_df.to_csv("weekly_inventory_projection.csv", index=False)

# File 3
def totals(letter):
    col = f"scenario_{letter}"
    cancel = per_po[per_po[col] == "CANCEL"]
    postpone = per_po[per_po[col] == "POSTPONE"]
    return {
        "cancel_eur": round(float(cancel["eur_value"].sum())),
        "postpone_eur": round(float(postpone["eur_value"].sum())),
        "total_impact_eur": round(float(cancel["eur_value"].sum() + postpone["eur_value"].sum())),
    }

def counts(letter):
    col = f"scenario_{letter}"
    cancel = per_po[per_po[col] == "CANCEL"]
    postpone = per_po[per_po[col] == "POSTPONE"]
    return {
        "cancel_pos": int(len(cancel)),
        "postpone_pos": int(len(postpone)),
        "cancel_skus": int(cancel["sku"].nunique()),
        "postpone_skus": int(postpone["sku"].nunique()),
    }

def peak_info(series):
    mx = max(series)
    return {"eur": round(mx), "week": proj_labels[series.index(mx)], "under_5m": mx < 5_000_000}

movers_to_postpone, movers_to_produce = [], []
for _, r in per_po.iterrows():
    key = f"{r['sku']}@CW{r['po_week']}"
    if r["scenario_a"] == "POSTPONE" and r["_avg_a"] != "POSTPONE":
        movers_to_postpone.append(key)
    if r["scenario_a"] == "PRODUCE" and r["_avg_a"] in ("POSTPONE", "CANCEL"):
        movers_to_produce.append(key)

summary = {
    "weeks": proj_labels,
    "scenarios": {
        "baseline": [round(v) for v in baseline_series],
        "scenario_a": [round(v) for v in scen_series["a"]],
        "scenario_b": [round(v) for v in scen_series["b"]],
        "scenario_c": [round(v) for v in scen_series["c"]],
    },
    "totals": {f"scenario_{l}": totals(l) for l in ["a", "b", "c"]},
    "counts": {f"scenario_{l}": counts(l) for l in ["a", "b", "c"]},
    "peak": {
        "baseline": peak_info(baseline_series),
        "scenario_a": peak_info(scen_series["a"]),
        "scenario_b": peak_info(scen_series["b"]),
        "scenario_c": peak_info(scen_series["c"]),
    },
    "movers": {
        "to_postpone_after_real_demand": movers_to_postpone,
        "to_produce_after_real_demand": movers_to_produce,
    },
    "sanity_un_postpones": dict(sanity_unpostpones),
}
with open("scenario_summary.json", "w", encoding="utf-8") as f:
    json.dump(summary, f, indent=2)

# File 4: supplier-facing scenario B lists
def supplier_export(action_label, scenario_letter):
    col = f"scenario_{scenario_letter}"
    new_col = f"new_delivery_week_{scenario_letter}"
    sel = per_po[per_po[col] == action_label].copy()
    sel["original_delivery_week"] = sel.apply(lambda r: f"CW{int(r['po_week'])}/{int(r['po_year'])}", axis=1)
    sel["action"] = action_label
    sel["new_delivery_week_or_blank"] = sel[new_col] if action_label == "POSTPONE" else ""
    sel["our_notes"] = ""
    out = sel[["sku", "tier", "original_delivery_week", "qty", "cost_price",
               "eur_value", "action", "new_delivery_week_or_blank", "our_notes"]]
    return out.sort_values(["tier", "eur_value"], ascending=[True, False])

supplier_export("CANCEL", "b").to_csv("abc_cancel_list_scenario_b.csv", index=False)
supplier_export("POSTPONE", "b").to_csv("abc_postpone_list_scenario_b.csv", index=False)

# ----- console summary -----
print("=" * 72)
print(f"Baseline current stock anchor: EUR {CURRENT_STOCK_EUR:,}")
bp = peak_info(baseline_series)
print(f"Baseline peak:                 EUR {bp['eur']:>10,}  @ {bp['week']}  under_5m={bp['under_5m']}")
print("=" * 72)
for letter in ["a", "b", "c"]:
    p = peak_info(scen_series[letter])
    t = totals(letter)
    c = counts(letter)
    print(f"\nScenario {letter.upper()} (cancel_threshold = {CANCEL_THRESHOLDS[['a','b','c'].index(letter)]})")
    print(f"  Peak: EUR {p['eur']:>10,}  @ {p['week']}  under_5m={p['under_5m']}")
    print(f"  CANCEL:   {c['cancel_pos']:>3} POs / {c['cancel_skus']:>3} SKUs / EUR {t['cancel_eur']:>10,}")
    print(f"  POSTPONE: {c['postpone_pos']:>3} POs / {c['postpone_skus']:>3} SKUs / EUR {t['postpone_eur']:>10,}")
    print(f"  Total impact (cancel+postpone, EUR off the books in window): EUR {t['total_impact_eur']:,}")

print(f"\nMovers (real vs avg, scenario A):")
print(f"  flipped TO postpone: {len(movers_to_postpone)} POs")
print(f"  flipped TO produce:  {len(movers_to_produce)} POs")
print(f"\nSanity un-postpones (multi-PO violations):")
for k, v in sanity_unpostpones.items():
    print(f"  Scenario {k.upper()}: {len(v)} un-postpones -> {v[:5]}{'...' if len(v) > 5 else ''}")

print("\nFiles written:")
for f in ["three_scenarios_per_po.csv", "weekly_inventory_projection.csv",
          "scenario_summary.json", "abc_cancel_list_scenario_b.csv",
          "abc_postpone_list_scenario_b.csv"]:
    print(f"  {f}")
