"""
Polleo — €5M Stock Cap Compliance Plan
========================================

Hard constraint surfaced by management: total stock value across
warehouse + retail stores must NOT exceed €5M (at cost).

Current state (live):
  Warehouse:                €3.14M
  Stores:                   €1.45M
  TOTAL:                    €4.59M     headroom €0.41M
  Already-incoming PO pipeline: €2.60M
  Without any action, projected total: ~€7.19M  (€2.19M over cap)

The pipeline by health classification of receiving SKU:
  STAR    POs: €1.44M  ← keep, urgent need
  HEALTHY POs: €0.61M  ← keep
  SLOW    POs: €0.52M  ← CANCEL/DEFER candidates
  DEAD    POs: €0.025M ← CANCEL outright

This script produces a 5-sheet workbook:
  1. Cap Situation         — executive headline, what hits the wall and when
  2. Cancel PO List        — SLOW/DEAD incoming, in priority order
  3. Keep PO List          — STAR/HEALTHY incoming (let them arrive)
  4. Stock Timeline        — week-by-week projection with cap line + 2 scenarios
  5. Cap-Aware Reinvest    — when each new order can be placed without breach

Reads LIVE from PostgreSQL.
Output: polleo_stock_cap_plan_YYYYMMDD.xlsx
"""
from __future__ import annotations

import sys
import math
from datetime import datetime, date, timedelta
from pathlib import Path

import pandas as pd
import numpy as np
from sqlalchemy import text

if hasattr(sys.stdout, "reconfigure"):
    sys.stdout.reconfigure(encoding="utf-8")

from db.connection import get_engine, is_db_available

ROOT      = Path(__file__).resolve().parent
TODAY     = date.today()
OUT_XLSX  = ROOT / f"polleo_stock_cap_plan_{TODAY:%Y%m%d}.xlsx"

# --- constraints ----------------------------------------------------------
STOCK_CAP_EUR          = 5_000_000.0
SALES_WINDOW_DAYS      = 77
SAFETY_STOCK_WEEKS     = 2.0
DEFAULT_LEAD_TIME      = 4.0
REINVEST_BUDGET_EUR    = 400_000.0
LIQUIDATION_TOTAL_EUR  = 1_590_000.0   # Scenario C target stock reduction
LIQUIDATION_WEEKS      = 16            # spread evenly over 16 weeks
LIQUIDATION_START_WEEK = 2             # weeks from today (W1 = setup)
WOS_STAR_MAX           = 8
WOS_HEALTHY_MAX        = 16
WOS_SLOW_MAX           = 52
HORIZON_WEEKS          = 16

URG_COLORS = {1: "F8CBAD", 2: "FFE699", 3: "C6EFCE", 4: "DDEBF7"}
HEALTH_COLOR = {"STAR": "C6EFCE", "HEALTHY": "DDEBF7",
                "SLOW": "FFE699", "DEAD": "F8CBAD"}


# ---------------------------------------------------------------------------
def load_live():
    eng = get_engine()

    # Total stock value (WH + stores) at cost
    cur_total_q = """
        SELECT
            CASE WHEN ds.is_warehouse THEN 'Warehouse' ELSE 'Stores' END AS loc,
            SUM(sc.stock_qty * COALESCE(ec.cost_price, 0)) AS value
        FROM erp_stock_current sc
        JOIN dim_stores ds ON sc.store_id = ds.id
        LEFT JOIN erp_costs ec ON ec.product_id = sc.product_id
        GROUP BY ds.is_warehouse
    """
    cur_total = pd.read_sql(cur_total_q, eng)

    # Universe — every SKU with stock OR incoming PO, with health metrics
    sku_q = text("""
        WITH last_tx AS (SELECT MAX(transaction_date) AS d FROM erp_transactions),
        sales AS (
            SELECT t.product_id, SUM(GREATEST(t.quantity,0)) AS qty_sold
            FROM erp_transactions t
            LEFT JOIN lookup_channel_map cm ON t.channel_map_id=cm.id
            WHERE t.transaction_date BETWEEN
                  (SELECT d FROM last_tx) - INTERVAL ':days days'
                  AND (SELECT d FROM last_tx)
              AND cm.channel IN ('retail','webshop','wholesale')
            GROUP BY t.product_id
        ),
        wh_stock AS (
            SELECT sc.product_id, sc.stock_qty
            FROM erp_stock_current sc
            JOIN dim_stores ds ON sc.store_id=ds.id
            WHERE ds.is_warehouse=TRUE
        )
        SELECT
            p.id AS product_id, p.sku, p.name AS sku_name,
            COALESCE(dc.name,'(no cat)') AS category,
            sp.tier,
            COALESCE(ws.stock_qty, 0) AS qty_on_hand,
            COALESCE(ec.cost_price, 0) AS cost_price,
            COALESCE(s.qty_sold, 0) AS qty_sold
        FROM dim_products p
        LEFT JOIN dim_categories dc ON p.category_id=dc.id
        LEFT JOIN sku_planning sp ON sp.product_id=p.id
        LEFT JOIN wh_stock ws ON ws.product_id=p.id
        LEFT JOIN erp_costs ec ON ec.product_id=p.id
        LEFT JOIN sales s     ON s.product_id=p.id
    """.replace(":days days", f"{SALES_WINDOW_DAYS - 1} days"))
    sku = pd.read_sql(sku_q, eng)

    # Incoming POs with arrival week + value
    inc_q = """
        SELECT i.id AS po_id, i.product_id, i.year, i.week,
               i.quantity, i.status,
               COALESCE(ec.cost_price, 0) AS cost_price
        FROM incoming_supply i
        LEFT JOIN erp_costs ec ON ec.product_id = i.product_id
    """
    incoming = pd.read_sql(inc_q, eng)

    # Lead time + MOQ per SKU
    supply = pd.read_sql(
        "SELECT product_id, lead_time_weeks, moq FROM supply_master", eng
    )

    # Latest forecast — for reinvest sizing
    fc_q = """
        WITH latest AS (SELECT MAX(id) AS rid FROM forecast_runs)
        SELECT product_id, year, week, COALESCE(total,0) AS qty
        FROM forecasts, latest
        WHERE run_id = latest.rid
    """
    fc = pd.read_sql(fc_q, eng)

    return {
        "cur_total":  cur_total,
        "sku":        sku,
        "incoming":   incoming,
        "supply":     supply,
        "forecast":   fc,
    }


# ---------------------------------------------------------------------------
def classify(wos, qsold):
    if qsold <= 0: return "DEAD"
    if not np.isfinite(wos): return "DEAD"
    if wos < WOS_STAR_MAX:    return "STAR"
    if wos < WOS_HEALTHY_MAX: return "HEALTHY"
    if wos < WOS_SLOW_MAX:    return "SLOW"
    return "DEAD"


def compute(data):
    sku = data["sku"].copy()
    inc = data["incoming"].copy()
    sup = data["supply"]
    fc  = data["forecast"]

    n_weeks = SALES_WINDOW_DAYS / 7.0
    sku["wkly_run"] = sku["qty_sold"] / n_weeks
    sku["wos"]      = np.where(sku["wkly_run"] > 0,
                                sku["qty_on_hand"] / sku["wkly_run"],
                                np.inf)
    sku["health"]   = [classify(w, q)
                       for w, q in zip(sku["wos"], sku["qty_sold"])]

    inc["value"] = inc["quantity"].astype(float) * inc["cost_price"].astype(float)
    inc = inc.merge(sku[["product_id", "sku", "sku_name", "category",
                          "tier", "qty_on_hand", "wos", "health"]],
                     on="product_id", how="left")
    # cancel candidates
    inc["recommendation"] = np.where(
        inc["health"].isin(["DEAD"]),
        "CANCEL — receiver SKU is dead",
        np.where(
            inc["health"] == "SLOW",
            "DEFER/CANCEL — already 16+ weeks coverage",
            np.where(
                inc["health"] == "STAR",
                "KEEP — needed for STAR replenishment",
                "KEEP — healthy SKU"
            )
        )
    )
    return sku, inc


# ---------------------------------------------------------------------------
def build_timeline(cur_total_eur, incoming, scenario="base"):
    """Week-by-week stock value projection.

    scenario = 'base'        — all POs land, no liquidation
               'with_action' — cancel SLOW/DEAD POs + apply liquidation drawdown
    """
    # Today is week N of year Y — find current ISO week
    today = TODAY
    iso = today.isocalendar()
    cur_y, cur_w = iso.year, iso.week

    # Build week index for horizon
    weeks_idx = []
    y, w = cur_y, cur_w
    for _ in range(HORIZON_WEEKS):
        weeks_idx.append((y, w))
        w += 1
        if w > 52:
            w = 1; y += 1

    # Incoming PO arrivals per week
    inc = incoming.copy()
    if scenario == "with_action":
        # cancel SLOW/DEAD
        inc = inc[~inc["health"].isin(["SLOW", "DEAD"])]

    inc_by_week = (inc.groupby(["year","week"])["value"].sum()
                       .reset_index())

    # Liquidation drawdown
    liq_weekly = (LIQUIDATION_TOTAL_EUR / LIQUIDATION_WEEKS) \
                  if scenario == "with_action" else 0

    rows = []
    running = cur_total_eur
    for i, (yy, ww) in enumerate(weeks_idx):
        arr = float(inc_by_week.loc[
            (inc_by_week["year"] == yy) & (inc_by_week["week"] == ww),
            "value"].sum())
        liq = liq_weekly if i >= LIQUIDATION_START_WEEK else 0
        running = running + arr - liq
        rows.append({
            "week_offset": i,
            "year":        yy,
            "iso_week":    ww,
            "incoming":    arr,
            "liquidation": liq,
            "stock_value": running,
            "over_cap":    max(running - STOCK_CAP_EUR, 0),
        })
    return pd.DataFrame(rows)


# ---------------------------------------------------------------------------
def build_reinvest(sku_df, fc, current_total_stock, cancel_value):
    """Compute reinvest orders, but only allow them when projected stock allows.

    Simple sequencing: assume cancelled POs free headroom immediately;
    liquidation drawdown adds further headroom over time. We pace
    reinvest orders so they arrive after enough headroom exists.
    """
    sku = sku_df.copy()

    # forecast lookup
    fc_cum = {}
    for pid, g in fc.groupby("product_id"):
        fc_cum[pid] = g.sort_values(["year","week"])["qty"].cumsum().tolist()

    def fc_qty(pid, n_weeks):
        if pid not in fc_cum: return None
        n = min(int(math.ceil(n_weeks)), len(fc_cum[pid]))
        if n <= 0: return 0.0
        return float(fc_cum[pid][n - 1])

    sku["lead_time_weeks"] = sku["lead_time_weeks"].fillna(DEFAULT_LEAD_TIME) \
        if "lead_time_weeks" in sku.columns else DEFAULT_LEAD_TIME

    eligible = sku[sku["health"].isin(["STAR","HEALTHY"])].copy()
    eligible["weekly_run_rate"] = eligible["wkly_run"]
    eligible["margin_yield"]    = 0  # already computed elsewhere; not critical here

    # Compute order qty as before
    plan_rows = []
    for r in eligible.itertuples(index=False):
        lead = float(getattr(r, "lead_time_weeks", DEFAULT_LEAD_TIME))
        target = lead + SAFETY_STOCK_WEEKS
        qty_need = fc_qty(r.product_id, target)
        if qty_need is None:
            qty_need = r.wkly_run * target
        net = qty_need - r.qty_on_hand
        if net <= 0:
            continue
        moq = float(getattr(r, "moq", 0) or 0)
        if moq and net < moq:
            net = moq
        cost = float(r.cost_price)
        value = net * cost
        if value <= 0:
            continue
        wos = r.wos
        if not np.isfinite(wos):
            urg = 4
        elif wos <= lead:        urg = 1
        elif wos <= lead + 2:    urg = 2
        elif wos <= lead + 4:    urg = 3
        else:                    urg = 4
        plan_rows.append({
            "sku": r.sku, "sku_name": r.sku_name,
            "category": r.category, "health": r.health,
            "qty_on_hand": r.qty_on_hand, "wos": r.wos,
            "lead_time_weeks": lead,
            "order_qty": net, "cost_price": cost,
            "order_value": value, "urgency_week": urg,
        })
    plan = pd.DataFrame(plan_rows).sort_values(
        ["urgency_week"], ascending=[True]
    )
    plan["cumulative"] = plan["order_value"].cumsum()
    plan["in_budget"]  = plan["cumulative"] <= REINVEST_BUDGET_EUR

    # Headroom check — cumulative new orders must fit into post-liquidation
    # window. Simple model: headroom = cancel_value + liquidation_progress.
    # We assume orders placed in week W arrive at W+lead_time, so for now
    # tag them with arrival week.
    plan["assumed_order_week"] = plan["urgency_week"]
    plan["arrival_week"] = plan["assumed_order_week"] + plan["lead_time_weeks"]
    return plan


# ---------------------------------------------------------------------------
def fmt_eur(x):
    try: return f"€{x:,.0f}"
    except Exception: return str(x)


def write_excel(sku, inc, timeline_base, timeline_action, reinvest,
                cur_total_eur, cancel_total, keep_total, out_path):
    from openpyxl import Workbook
    from openpyxl.styles import Font, PatternFill, Alignment, Border, Side
    from openpyxl.utils import get_column_letter
    from openpyxl.chart import LineChart, BarChart, Reference
    from openpyxl.chart.label import DataLabelList

    wb = Workbook()
    H1 = Font(bold=True, size=14, color="FFFFFF")
    H2 = Font(bold=True, size=11, color="FFFFFF")
    H3 = Font(bold=True, size=11)
    H_RED = Font(bold=True, size=11, color="C00000")
    FILL_DARK = PatternFill("solid", fgColor="1F4E79")
    FILL_MID  = PatternFill("solid", fgColor="2E75B6")
    FILL_RED  = PatternFill("solid", fgColor="F8CBAD")
    FILL_AMB  = PatternFill("solid", fgColor="FFE699")
    FILL_GRN  = PatternFill("solid", fgColor="C6EFCE")

    def title(ws, text, span=8):
        c = ws.cell(row=1, column=1, value=text)
        c.font = H1; c.fill = FILL_DARK
        ws.merge_cells(start_row=1, start_column=1, end_row=1, end_column=span)
        ws.row_dimensions[1].height = 22

    def autosize(ws, mx=52):
        for col in ws.columns:
            try: letter = get_column_letter(col[0].column)
            except AttributeError: continue
            length = 12
            for c in col:
                if c.value is None: continue
                length = max(length, min(mx, len(str(c.value)) + 2))
            ws.column_dimensions[letter].width = length

    # ====================================================================
    # SHEET 1: Cap Situation
    # ====================================================================
    ws = wb.active
    ws.title = "Cap Situation"
    title(ws, f"€{STOCK_CAP_EUR/1e6:.0f}M STOCK CAP — current position + projected breach", span=8)
    ws.cell(row=2, column=1,
            value=f"Generated {datetime.now():%Y-%m-%d %H:%M}    "
                  f"Cap applies to WAREHOUSE + STORES combined, at cost").font = Font(italic=True, color="595959")
    ws.merge_cells("A2:H2")

    row = 4
    ws.cell(row=row, column=1, value="CURRENT POSITION (live)").font = H3
    row += 1
    block = [
        ("Stock cap",                        fmt_eur(STOCK_CAP_EUR)),
        ("Current total stock (WH+stores)",  fmt_eur(cur_total_eur)),
        ("Headroom right now",               fmt_eur(STOCK_CAP_EUR - cur_total_eur)),
        ("Already-incoming PO pipeline",     fmt_eur(inc["value"].sum())),
        ("Projected stock if all POs land",  fmt_eur(cur_total_eur + inc["value"].sum())),
        ("Projected BREACH (over cap)",      fmt_eur(max(0, cur_total_eur + inc["value"].sum() - STOCK_CAP_EUR))),
    ]
    for label, val in block:
        ws.cell(row=row, column=1, value=label).font = Font(bold=True)
        c = ws.cell(row=row, column=4, value=val)
        if "BREACH" in label or "Headroom" in label:
            c.font = H_RED
        row += 1
    row += 1

    ws.cell(row=row, column=1, value="INCOMING PO BREAKDOWN").font = H3
    row += 1
    by_health = inc.groupby("health")["value"].agg(['count','sum']).reindex(
        ["STAR","HEALTHY","SLOW","DEAD"]).fillna(0)
    for i, h in enumerate(["STAR","HEALTHY","SLOW","DEAD"]):
        cnt = int(by_health.loc[h,"count"])
        val = float(by_health.loc[h,"sum"])
        rec = "KEEP" if h in ("STAR","HEALTHY") else "CANCEL/DEFER"
        c1 = ws.cell(row=row, column=1, value=h); c1.fill = PatternFill("solid", fgColor=HEALTH_COLOR[h])
        ws.cell(row=row, column=2, value=f"{cnt} POs")
        ws.cell(row=row, column=4, value=fmt_eur(val))
        ws.cell(row=row, column=6, value=rec).font = Font(bold=True,
            color=("006100" if rec == "KEEP" else "C00000"))
        row += 1
    row += 1

    ws.cell(row=row, column=1, value="ACTION SUMMARY").font = H3
    row += 1
    actions = [
        f"CANCEL/DEFER {(inc['health'].isin(['SLOW','DEAD'])).sum()} POs worth "
        f"{fmt_eur(cancel_total)} — frees headroom WITHOUT losing sales (these SKUs already have 16-50w stock).",
        f"KEEP {(inc['health'].isin(['STAR','HEALTHY'])).sum()} POs worth "
        f"{fmt_eur(keep_total)} — needed to prevent stockouts on top sellers.",
        f"After cancellation: projected stock = "
        f"{fmt_eur(cur_total_eur + keep_total)}  "
        f"(still {'OVER' if cur_total_eur + keep_total > STOCK_CAP_EUR else 'under'} cap by "
        f"{fmt_eur(abs(cur_total_eur + keep_total - STOCK_CAP_EUR))}).",
        f"Liquidation cadence (Scenario C): ~"
        f"{fmt_eur(LIQUIDATION_TOTAL_EUR / LIQUIDATION_WEEKS)}/week starting W{LIQUIDATION_START_WEEK} "
        f"for {LIQUIDATION_WEEKS} weeks — brings projection back under cap by ~W"
        f"{LIQUIDATION_WEEKS - 2}.",
        "Reinvestment orders MUST be sequenced so they arrive AFTER liquidation "
        f"has opened headroom — see 'Cap-Aware Reinvest' sheet.",
    ]
    for a in actions:
        ws.cell(row=row, column=1, value="• " + a)
        ws.merge_cells(start_row=row, start_column=1,
                       end_row=row, end_column=8)
        row += 1
    autosize(ws)

    # ====================================================================
    # SHEET 2: Cancel PO List
    # ====================================================================
    ws = wb.create_sheet("Cancel POs")
    title(ws, "PO CANCEL / DEFER list — SLOW + DEAD receivers (act THIS WEEK)", span=10)
    cancel_df = inc[inc["health"].isin(["SLOW","DEAD"])].copy()
    cancel_df = cancel_df.sort_values("value", ascending=False)
    cols = ["sku","sku_name","category","tier","health",
            "qty_on_hand","wos","year","week","quantity",
            "cost_price","value","recommendation"]
    hdr = 3
    for i, c in enumerate(cols, 1):
        cell = ws.cell(row=hdr, column=i, value=c)
        cell.font = H2; cell.fill = FILL_MID

    for i, r in enumerate(cancel_df.itertuples(index=False), start=hdr+1):
        ws.cell(row=i, column=1,  value=r.sku)
        ws.cell(row=i, column=2,  value=r.sku_name)
        ws.cell(row=i, column=3,  value=r.category)
        ws.cell(row=i, column=4,  value=r.tier)
        ws.cell(row=i, column=5,  value=r.health)
        ws.cell(row=i, column=5).fill = PatternFill("solid", fgColor=HEALTH_COLOR[r.health])
        ws.cell(row=i, column=6,  value=float(r.qty_on_hand))
        ws.cell(row=i, column=7,
                value=(None if not np.isfinite(r.wos) else float(r.wos)))
        ws.cell(row=i, column=8,  value=int(r.year))
        ws.cell(row=i, column=9,  value=int(r.week))
        ws.cell(row=i, column=10, value=float(r.quantity))
        ws.cell(row=i, column=11, value=float(r.cost_price))
        ws.cell(row=i, column=12, value=float(r.value))
        ws.cell(row=i, column=13, value=r.recommendation)

    last = hdr + len(cancel_df)
    for col_idx, fmt in [(6,"#,##0.0"),(7,"0.0"),(10,"#,##0.0"),
                          (11,"#,##0.00"),(12,"#,##0")]:
        for r in range(hdr+1, last+1):
            ws.cell(row=r, column=col_idx).number_format = fmt

    sub = last + 2
    ws.cell(row=sub, column=1, value="TOTAL CANCELLABLE").font = Font(bold=True)
    ws.cell(row=sub, column=12, value=float(cancel_df["value"].sum())).number_format = "#,##0"
    for c in range(1,14):
        ws.cell(row=sub, column=c).fill = FILL_AMB

    ws.freeze_panes = ws.cell(row=hdr+1, column=1)
    autosize(ws)

    # ====================================================================
    # SHEET 3: Keep PO List
    # ====================================================================
    ws = wb.create_sheet("Keep POs")
    title(ws, "PO KEEP list — STAR + HEALTHY receivers (let them arrive)", span=10)
    keep_df = inc[inc["health"].isin(["STAR","HEALTHY"])].copy()
    keep_df = keep_df.sort_values("value", ascending=False)
    hdr = 3
    for i, c in enumerate(cols, 1):
        cell = ws.cell(row=hdr, column=i, value=c)
        cell.font = H2; cell.fill = FILL_MID
    for i, r in enumerate(keep_df.itertuples(index=False), start=hdr+1):
        ws.cell(row=i, column=1,  value=r.sku)
        ws.cell(row=i, column=2,  value=r.sku_name)
        ws.cell(row=i, column=3,  value=r.category)
        ws.cell(row=i, column=4,  value=r.tier)
        ws.cell(row=i, column=5,  value=r.health)
        ws.cell(row=i, column=5).fill = PatternFill("solid", fgColor=HEALTH_COLOR[r.health])
        ws.cell(row=i, column=6,  value=float(r.qty_on_hand))
        ws.cell(row=i, column=7,
                value=(None if not np.isfinite(r.wos) else float(r.wos)))
        ws.cell(row=i, column=8,  value=int(r.year))
        ws.cell(row=i, column=9,  value=int(r.week))
        ws.cell(row=i, column=10, value=float(r.quantity))
        ws.cell(row=i, column=11, value=float(r.cost_price))
        ws.cell(row=i, column=12, value=float(r.value))
        ws.cell(row=i, column=13, value=r.recommendation)
    last = hdr + len(keep_df)
    for col_idx, fmt in [(6,"#,##0.0"),(7,"0.0"),(10,"#,##0.0"),
                          (11,"#,##0.00"),(12,"#,##0")]:
        for r in range(hdr+1, last+1):
            ws.cell(row=r, column=col_idx).number_format = fmt
    sub = last + 2
    ws.cell(row=sub, column=1, value="TOTAL TO KEEP").font = Font(bold=True)
    ws.cell(row=sub, column=12, value=float(keep_df["value"].sum())).number_format = "#,##0"
    for c in range(1,14):
        ws.cell(row=sub, column=c).fill = FILL_GRN
    ws.freeze_panes = ws.cell(row=hdr+1, column=1)
    autosize(ws)

    # ====================================================================
    # SHEET 4: Stock Timeline
    # ====================================================================
    ws = wb.create_sheet("Stock Timeline")
    title(ws, "16-week stock value projection — base vs cap-compliant scenarios", span=8)

    hdr = 3
    cols = ["Week offset", "Year", "ISO week",
            "Base: incoming (€)", "Base: stock value (€)",
            "Cap-compliant: incoming (€)", "Cap-compliant: liquidation (€)",
            "Cap-compliant: stock value (€)", "Cap (€)"]
    for i, c in enumerate(cols, 1):
        cell = ws.cell(row=hdr, column=i, value=c)
        cell.font = H2; cell.fill = FILL_MID

    for i in range(len(timeline_base)):
        r = timeline_base.iloc[i]
        r2 = timeline_action.iloc[i]
        row = hdr + 1 + i
        ws.cell(row=row, column=1, value=int(r["week_offset"]))
        ws.cell(row=row, column=2, value=int(r["year"]))
        ws.cell(row=row, column=3, value=int(r["iso_week"]))
        ws.cell(row=row, column=4, value=float(r["incoming"])).number_format = "#,##0"
        ws.cell(row=row, column=5, value=float(r["stock_value"])).number_format = "#,##0"
        ws.cell(row=row, column=6, value=float(r2["incoming"])).number_format = "#,##0"
        ws.cell(row=row, column=7, value=float(r2["liquidation"])).number_format = "#,##0"
        ws.cell(row=row, column=8, value=float(r2["stock_value"])).number_format = "#,##0"
        ws.cell(row=row, column=9, value=STOCK_CAP_EUR).number_format = "#,##0"
        if r["stock_value"] > STOCK_CAP_EUR:
            ws.cell(row=row, column=5).fill = FILL_RED
        if r2["stock_value"] > STOCK_CAP_EUR:
            ws.cell(row=row, column=8).fill = FILL_RED

    last = hdr + len(timeline_base)
    chart = LineChart()
    chart.title = "Stock value projection — base vs cap-compliant (€)"
    chart.y_axis.title = "EUR"
    chart.x_axis.title = "Week offset from today"
    # add 3 series: base, cap-compliant, cap line
    data_ref = Reference(ws, min_col=5, min_row=hdr,
                         max_row=last, max_col=5)
    chart.add_data(data_ref, titles_from_data=True)
    data_ref2 = Reference(ws, min_col=8, min_row=hdr,
                          max_row=last, max_col=8)
    chart.add_data(data_ref2, titles_from_data=True)
    cap_ref = Reference(ws, min_col=9, min_row=hdr,
                        max_row=last, max_col=9)
    chart.add_data(cap_ref, titles_from_data=True)
    cats = Reference(ws, min_col=1, min_row=hdr+1, max_row=last)
    chart.set_categories(cats)
    chart.height = 11; chart.width = 18
    ws.add_chart(chart, "K3")
    autosize(ws)

    # ====================================================================
    # SHEET 5: Cap-Aware Reinvest Plan
    # ====================================================================
    ws = wb.create_sheet("Cap-Aware Reinvest")
    title(ws, "Reinvest order timing — when each new PO can be placed without breach", span=10)
    ws.cell(row=2, column=1,
            value=f"Total reinvest budget: €{REINVEST_BUDGET_EUR:,.0f}    "
                  f"Cap-compliant only if liquidation runs in parallel — "
                  f"first orders arrive after week ~{int(reinvest['arrival_week'].min()) if not reinvest.empty else 'N/A'}"
            ).font = Font(italic=True, color="595959")
    ws.merge_cells("A2:J2")

    hdr = 4
    cols = ["sku","sku_name","category","health","qty_on_hand","wos",
            "lead_time_weeks","order_qty","cost_price","order_value",
            "urgency_week","arrival_week","cumulative","in_budget"]
    for i, c in enumerate(cols, 1):
        cell = ws.cell(row=hdr, column=i, value=c)
        cell.font = H2; cell.fill = FILL_MID

    for i, r in enumerate(reinvest.itertuples(index=False), start=hdr+1):
        ws.cell(row=i, column=1,  value=r.sku)
        ws.cell(row=i, column=2,  value=r.sku_name)
        ws.cell(row=i, column=3,  value=r.category)
        ws.cell(row=i, column=4,  value=r.health)
        ws.cell(row=i, column=4).fill = PatternFill("solid", fgColor=HEALTH_COLOR[r.health])
        ws.cell(row=i, column=5,  value=float(r.qty_on_hand))
        ws.cell(row=i, column=6,
                value=(None if not np.isfinite(r.wos) else float(r.wos)))
        ws.cell(row=i, column=7,  value=float(r.lead_time_weeks))
        ws.cell(row=i, column=8,  value=float(r.order_qty))
        ws.cell(row=i, column=9,  value=float(r.cost_price))
        ws.cell(row=i, column=10, value=float(r.order_value))
        ws.cell(row=i, column=11, value=int(r.urgency_week))
        ws.cell(row=i, column=11).fill = PatternFill("solid", fgColor=URG_COLORS[int(r.urgency_week)])
        ws.cell(row=i, column=12, value=float(r.arrival_week))
        ws.cell(row=i, column=13, value=float(r.cumulative))
        ws.cell(row=i, column=14, value=bool(r.in_budget))

    last = hdr + len(reinvest)
    for col_idx, fmt in [(5,"#,##0.0"),(6,"0.0"),(7,"0.0"),(8,"#,##0.0"),
                          (9,"#,##0.00"),(10,"#,##0"),(12,"0.0"),(13,"#,##0")]:
        for r in range(hdr+1, last+1):
            ws.cell(row=r, column=col_idx).number_format = fmt
    ws.freeze_panes = ws.cell(row=hdr+1, column=1)
    autosize(ws)

    wb.save(out_path)


# ---------------------------------------------------------------------------
def print_summary(cur_total, inc, t_base, t_action, reinvest,
                  cancel_total, keep_total):
    print("\n=== POLLEO €5M STOCK CAP COMPLIANCE ===\n")
    print(f"Current stock (WH+stores):  €{cur_total:,.0f}")
    print(f"Cap:                        €{STOCK_CAP_EUR:,.0f}")
    print(f"Headroom right now:         €{STOCK_CAP_EUR - cur_total:,.0f}")
    print(f"Incoming PO pipeline:       €{inc['value'].sum():,.0f}")
    print(f"Projected breach (no action): "
          f"€{max(0, cur_total + inc['value'].sum() - STOCK_CAP_EUR):,.0f}\n")
    print("PO by health classification:")
    print(inc.groupby("health")["value"].agg(['count','sum']).reindex(
        ["STAR","HEALTHY","SLOW","DEAD"]).to_string())
    print()
    print(f"CANCEL/DEFER SLOW+DEAD POs total: €{cancel_total:,.0f}")
    print(f"KEEP STAR+HEALTHY POs total:      €{keep_total:,.0f}")
    print()
    print("Stock timeline summary (base scenario):")
    print(f"  Peak stock value (no action):    €{t_base['stock_value'].max():,.0f} "
          f"at week offset {int(t_base['stock_value'].idxmax())}")
    print(f"  Peak with cancel + liquidation:  €{t_action['stock_value'].max():,.0f} "
          f"at week offset {int(t_action['stock_value'].idxmax())}")
    print(f"  Final week stock (with action):  €{t_action['stock_value'].iloc[-1]:,.0f}")
    print()
    print("Cap-aware reinvest plan:")
    print(f"  Eligible SKUs:        {len(reinvest):,}")
    print(f"  Within €400k budget:  {int(reinvest['in_budget'].sum()):,}")
    print(f"  Total order value:    €{reinvest.loc[reinvest['in_budget'],'order_value'].sum():,.0f}")


def main():
    if not is_db_available():
        print("ERROR: DB unreachable", file=sys.stderr); sys.exit(1)

    print("Loading live data…")
    data = load_live()
    cur_total = float(data["cur_total"]["value"].sum())

    sku, inc = compute(data)

    # Merge supply (lead_time, moq) into sku for reinvest math
    sku2 = sku.merge(data["supply"], on="product_id", how="left")
    sku2["lead_time_weeks"] = sku2["lead_time_weeks"].fillna(DEFAULT_LEAD_TIME)
    sku2["moq"] = sku2["moq"].fillna(0)

    cancel_total = float(inc.loc[inc["health"].isin(["SLOW","DEAD"]), "value"].sum())
    keep_total   = float(inc.loc[inc["health"].isin(["STAR","HEALTHY"]), "value"].sum())

    print("Building timelines…")
    t_base   = build_timeline(cur_total, inc, scenario="base")
    t_action = build_timeline(cur_total, inc, scenario="with_action")

    print("Building cap-aware reinvest plan…")
    reinvest = build_reinvest(sku2, data["forecast"], cur_total, cancel_total)

    print_summary(cur_total, inc, t_base, t_action, reinvest,
                  cancel_total, keep_total)

    print(f"\nWriting {OUT_XLSX.name} …")
    write_excel(sku, inc, t_base, t_action, reinvest,
                cur_total, cancel_total, keep_total, OUT_XLSX)
    print(f"Excel saved: {OUT_XLSX}")


if __name__ == "__main__":
    main()
