"""CFO S&OP audit — full analysis package.

Produces six deliverables in data/cfo_audit_outputs/:

  Step 1  lost_sales_projection.csv           — per-SKU stockout walk + lost € (cost + revenue)
  Step 2  locked_cash_overstock.csv           — per-SKU excess vs target stock
          locked_cash_summary_by_category.csv — category-level rollup
  Step 3  contest_july_risk.csv               — CW27-CW31 baseline vs uplift scenarios
  Step 4  slow_mover_analysis.csv             — LT-relative classification, all SKUs
          slow_mover_summary.csv              — pivoted by category + supplier
  Step 5  bridge_analysis.csv                 — per-SKU Vol/Price/Mix/COGS variance
          bridge_summary.csv                  — rolled up by category + total
          bridge_waterfall_data.json          — for visualisation
          bridge_feasibility_note.txt         — what we built + what's missing
  Step 6  cfo_dashboard_data.json             — master headline rollup

Universe: every product in dim_products that has stock OR forecast OR
incoming PO OR transactions. Cost uses the cascade defined in
demand_repo.get_sku_pricing (ERP standard → realized → NPD planned).

Run:
    python scripts/cfo_audit.py
"""
from __future__ import annotations

import json
import math
from collections import defaultdict
from datetime import date, datetime
from pathlib import Path
from typing import Optional

import pandas as pd
from sqlalchemy import text

from backend.models.database import SessionLocal

# ─────────────────────────────────────────────────────────────────────────
# Constants
# ─────────────────────────────────────────────────────────────────────────
ROOT = Path(__file__).resolve().parents[1]
OUTDIR = ROOT / "data" / "cfo_audit_outputs"
OUTDIR.mkdir(parents=True, exist_ok=True)

HORIZON_WEEKS = 13                  # CW20–CW32 = 13 weeks
CONTEST_START_CW = 27               # July
CONTEST_END_CW   = 31

# Tier-based fallbacks for SKUs without a supply_master row
LT_FALLBACK = {"01 GOLD": 5.0, "02 SILVER": 4.0, "03 BRONZE": 3.0}
SAFETY_WEEKS = {"01 GOLD": 2.0, "02 SILVER": 1.5, "03 BRONZE": 1.0}
DEFAULT_LT = 4.0
DEFAULT_SAFETY = 1.0
REVIEW_PERIOD_WEEKS = 1.0

GADGET_CATEGORIES = {"GADGETI", "GADGETS"}
NON_FOOD_CATEGORIES = {
    "GADGETI", "DRINKWARE I HOME", "BORILAČKA OPREMA", "BORILACKA OPREMA",
    "ODJEĆA", "ODJECA", "OBUĆA", "OBUCA", "REKVIZITI", "TEKSTIL",
    "OPREMA", "SUPLEMENTI ALATI",
}


# ─────────────────────────────────────────────────────────────────────────
# Helpers
# ─────────────────────────────────────────────────────────────────────────
def week_series(start_y: int, start_w: int, n: int) -> list[tuple[int, int]]:
    out, y, w = [], start_y, start_w
    for _ in range(n):
        out.append((y, w))
        w += 1
        if w > 52:
            y, w = y + 1, 1
    return out


def to_yw(y: int, w: int) -> int:
    return y * 100 + w


def _safe_div(a, b):
    try:
        return a / b if b else 0.0
    except Exception:
        return 0.0


# ─────────────────────────────────────────────────────────────────────────
# Load — pulls everything once into in-memory frames
# ─────────────────────────────────────────────────────────────────────────
def load_context(db) -> dict:
    """One big load — every step works off this dict to avoid round-trips."""
    cur = db.execute(text("""
        SELECT EXTRACT(ISOYEAR FROM now())::int AS y,
               EXTRACT(WEEK    FROM now())::int AS w
    """)).mappings().first()
    cur_year, cur_week = int(cur["y"]), int(cur["w"])
    horizon = week_series(cur_year, cur_week, HORIZON_WEEKS)
    horizon_yws = {to_yw(y, w) for (y, w) in horizon}

    # Master: every product with stock, forecast, incoming, or recent transaction
    products = pd.read_sql(text("""
        SELECT
            p.id              AS product_id,
            p.sku,
            COALESCE(p.name, '')              AS name,
            COALESCE(c.name, 'OTHER')         AS category,
            COALESCE(sp.tier, '')             AS tier,
            COALESCE(sm.lead_time_weeks::float, NULL) AS lead_time_weeks_raw,
            COALESCE(ds.name, '(unknown)')    AS supplier,
            ec.cost_price::float              AS erp_cost,
            ep.avg_sell_price::float          AS avg_sell_price,
            ep.normal_retail_ppp::float       AS retail_ppp,
            ep.normal_webshop_ppp::float      AS webshop_ppp,
            sp.vpc::float                     AS vpc,
            sp.ws_share_26w::float            AS ws_share
        FROM dim_products p
        LEFT JOIN dim_categories c   ON c.id  = p.category_id
        LEFT JOIN sku_planning  sp   ON sp.product_id = p.id
        LEFT JOIN supply_master sm   ON sm.product_id = p.id
        LEFT JOIN dim_suppliers ds   ON ds.id = sm.supplier_id
        LEFT JOIN erp_costs     ec   ON ec.product_id = p.id
        LEFT JOIN erp_prices    ep   ON ep.product_id = p.id
    """), db.bind)
    products["product_id"] = products["product_id"].astype(int)

    # Stock — split WH vs stores via dim_stores.unit_code
    stock_rows = pd.read_sql(text("""
        SELECT p.id AS product_id, p.sku,
               CASE WHEN ds.unit_code = '01' THEN 'wh' ELSE 'store' END AS bucket,
               SUM(esc.stock_qty)::float AS qty
        FROM erp_stock_current esc
        JOIN dim_products p ON p.id = esc.product_id
        JOIN dim_stores  ds ON ds.id = esc.store_id
        GROUP BY p.id, p.sku, bucket
    """), db.bind)
    wh_map    = {int(r.product_id): float(r.qty)
                 for r in stock_rows.itertuples() if r.bucket == "wh"}
    store_map = {int(r.product_id): float(r.qty)
                 for r in stock_rows.itertuples() if r.bucket == "store"}

    # Forecast (latest run)
    fc_df = pd.read_sql(text("""
        SELECT product_id, year, week,
               COALESCE(total, 0)::float    AS demand,
               COALESCE(baseline, 0)::float AS baseline
        FROM forecasts
        WHERE run_id = (SELECT MAX(run_id) FROM forecasts)
    """), db.bind)
    fc_total: dict[tuple[int, int, int], float] = {
        (int(r.product_id), int(r.year), int(r.week)): float(r.demand)
        for r in fc_df.itertuples()
    }
    fc_baseline: dict[tuple[int, int, int], float] = {
        (int(r.product_id), int(r.year), int(r.week)): float(r.baseline)
        for r in fc_df.itertuples()
    }

    # 12-week non-promo avg as fallback for unplanned SKUs (FMCG standard).
    if cur_week > 12:
        avg_cy, avg_cw = cur_year, cur_week - 12
    else:
        avg_cy, avg_cw = cur_year - 1, 52 + cur_week - 12
    max_yw = db.execute(text(
        "SELECT MAX(year*100+week) FROM v_sales_weekly_full"
    )).scalar() or 0
    avg_df = pd.read_sql(text("""
        WITH win AS (
            SELECT v.product_id, v.qty_total,
                   COALESCE(epw.is_erp_promo, FALSE) AS is_promo
            FROM v_sales_weekly_full v
            LEFT JOIN erp_promo_weeks epw
                ON epw.product_id = v.product_id
               AND epw.year = v.year AND epw.week = v.week
            WHERE v.year*100+v.week BETWEEN :cutoff AND :max_yw
        ),
        non_promo AS (
            SELECT product_id, AVG(qty_total)::float AS avg_np
            FROM win WHERE NOT is_promo GROUP BY product_id
        ),
        all_w AS (
            SELECT product_id, AVG(qty_total)::float AS avg_all
            FROM win GROUP BY product_id
        )
        SELECT a.product_id,
               COALESCE(np.avg_np, a.avg_all)::float AS avg_qty
        FROM all_w a LEFT JOIN non_promo np ON np.product_id = a.product_id
    """), db.bind, params={"cutoff": avg_cy*100 + avg_cw, "max_yw": max_yw})
    avg_map = {int(r.product_id): float(r.avg_qty or 0.0)
               for r in avg_df.itertuples()}

    # Realized cost / RUC per SKU (last 13w) for the cost cascade
    realized_df = pd.read_sql(text("""
        SELECT product_id,
               SUM(quantity)::float        AS units,
               SUM(purchase_value)::float  AS pv,
               SUM(ruc_eur)::float         AS ruc,
               SUM(total_value)::float     AS rev,
               SUM(approved_discount)::float AS disc
        FROM erp_transactions
        WHERE transaction_date >= (CURRENT_DATE - INTERVAL '13 weeks')
        GROUP BY product_id
    """), db.bind)
    realized: dict[int, dict] = {}
    for r in realized_df.itertuples():
        u = float(r.units or 0)
        if u <= 0:
            continue
        realized[int(r.product_id)] = {
            "units": u,
            "cost_per_unit":   (r.pv  / u) if r.pv  else None,
            "ruc_per_unit":    (r.ruc / u) if r.ruc else None,
            "rev_per_unit":    (r.rev / u) if r.rev else None,
            "disc_per_unit":   (r.disc / u) if r.disc else 0.0,
        }

    # NPD planned cost (last-resort fallback)
    npd_df = pd.read_sql(text("""
        SELECT sku, cost_price::float AS npd_cost, msrp_hr::float AS msrp_hr
        FROM npd_products
    """), db.bind)
    npd_cost  = {r.sku: float(r.npd_cost) for r in npd_df.itertuples()
                 if r.npd_cost is not None and not math.isnan(r.npd_cost)}
    npd_price = {r.sku: float(r.msrp_hr) for r in npd_df.itertuples()
                 if r.msrp_hr is not None and not math.isnan(r.msrp_hr)}

    # Incoming POs in horizon
    inc_df = pd.read_sql(text("""
        SELECT product_id, year, week, SUM(quantity)::float AS qty
        FROM incoming_supply
        WHERE COALESCE(LOWER(status), '') <> 'cancelled'
        GROUP BY product_id, year, week
    """), db.bind)
    inc_map: dict[tuple[int, int, int], float] = {
        (int(r.product_id), int(r.year), int(r.week)): float(r.qty)
        for r in inc_df.itertuples()
    }

    # Promo / contest signal — `erp_promo_weeks.is_erp_promo` flags
    promo_df = pd.read_sql(text("""
        SELECT product_id, year, week
        FROM erp_promo_weeks
        WHERE is_erp_promo = TRUE
    """), db.bind)
    promo_set = {(int(r.product_id), int(r.year), int(r.week))
                 for r in promo_df.itertuples()}

    # sku_uplift from data/sku_uplift.csv (loaded once)
    uplift_path = ROOT / "data" / "sku_uplift.csv"
    uplift_map: dict[str, dict] = {}
    if uplift_path.exists():
        up = pd.read_csv(uplift_path)
        for r in up.itertuples():
            uplift_map[r.sku] = {
                "promo_uplift": float(getattr(r, "promo_uplift", 1.0) or 1.0),
                "ws_uplift":    float(getattr(r, "ws_uplift", 1.0) or 1.0),
            }

    print(f"  context loaded — cur={cur_year}W{cur_week:02d}  horizon={[f'CW{w}' for _,w in horizon]}")
    print(f"  products={len(products)}  wh_skus={len(wh_map)}  forecast_pids={fc_df['product_id'].nunique()}")
    print(f"  realized 13w pids={len(realized)}  incoming pids={inc_df['product_id'].nunique()}")
    print(f"  promo (pid,yw) flags={len(promo_set)}  uplift rows={len(uplift_map)}")

    return {
        "cur_year": cur_year, "cur_week": cur_week,
        "horizon": horizon, "horizon_yws": horizon_yws,
        "products": products,
        "wh_map": wh_map, "store_map": store_map,
        "fc_total": fc_total, "fc_baseline": fc_baseline,
        "avg_map": avg_map,
        "realized": realized,
        "npd_cost": npd_cost, "npd_price": npd_price,
        "inc_map": inc_map,
        "promo_set": promo_set, "uplift_map": uplift_map,
    }


# ─────────────────────────────────────────────────────────────────────────
# Per-SKU helpers
# ─────────────────────────────────────────────────────────────────────────
def resolve_cost(pid: int, sku: str, prod_row, ctx) -> tuple[Optional[float], str]:
    """Cost cascade: erp_costs → realized → npd_planned. Returns (cost, source)."""
    erp = prod_row.erp_cost
    if erp is not None and not (isinstance(erp, float) and math.isnan(erp)):
        return float(erp), "erp_standard"
    realized = ctx["realized"].get(pid)
    if realized and realized.get("cost_per_unit"):
        return float(realized["cost_per_unit"]), "realized"
    npd = ctx["npd_cost"].get(sku)
    if npd is not None:
        return float(npd), "npd_planned"
    return None, ""


def resolve_price(pid: int, sku: str, prod_row, ctx) -> tuple[Optional[float], str]:
    """Selling price cascade: avg_sell_price → realized rev/qty → MSRP HR.
    Used for the 'lost sales revenue' figure."""
    asp = prod_row.avg_sell_price
    if asp is not None and not (isinstance(asp, float) and math.isnan(asp)) and asp > 0:
        return float(asp), "erp_avg_sell"
    realized = ctx["realized"].get(pid)
    if realized and realized.get("rev_per_unit"):
        return float(realized["rev_per_unit"]), "realized"
    msrp = ctx["npd_price"].get(sku)
    if msrp is not None:
        return float(msrp), "msrp"
    return None, ""


def resolve_lead_time(prod_row) -> tuple[float, str]:
    lt = prod_row.lead_time_weeks_raw
    if lt is not None and not (isinstance(lt, float) and math.isnan(lt)) and lt > 0:
        return float(lt), "supply_master"
    tier = prod_row.tier or ""
    lt = LT_FALLBACK.get(tier, DEFAULT_LT)
    return float(lt), "tier_fallback"


def resolve_safety_weeks(tier: str) -> float:
    return SAFETY_WEEKS.get(tier, DEFAULT_SAFETY)


def weekly_demand(pid: int, y: int, w: int, ctx) -> float:
    """Forecast first, 4w-avg fallback."""
    v = ctx["fc_total"].get((pid, y, w))
    if v is not None:
        return float(v)
    return float(ctx["avg_map"].get(pid, 0.0))


# ─────────────────────────────────────────────────────────────────────────
# Step 1 — Lost sales projection
# ─────────────────────────────────────────────────────────────────────────
def step1_lost_sales(ctx) -> dict:
    print("\n[Step 1] Lost sales projection ...")
    horizon = ctx["horizon"]
    rows: list[dict] = []

    for prod in ctx["products"].itertuples():
        pid, sku = int(prod.product_id), prod.sku
        # Universe filter: only include SKUs with any meaningful activity
        wh    = ctx["wh_map"].get(pid, 0.0)
        store = ctx["store_map"].get(pid, 0.0)
        avg   = ctx["avg_map"].get(pid, 0.0)
        has_fc = any((pid, y, w) in ctx["fc_total"] for (y, w) in horizon)
        if wh + store < 1 and avg < 0.1 and not has_fc:
            continue   # truly dormant — skip

        cost, _  = resolve_cost(pid, sku, prod, ctx)
        price, _ = resolve_price(pid, sku, prod, ctx)
        cost  = cost  or 0.0
        price = price or (cost * 2 if cost else 0.0)  # 100 % markup fallback only if cost known

        stock_now = wh + store
        total_demand = 0.0
        total_incoming = 0.0
        lost_units = 0.0
        stockout_yw: Optional[int] = None

        s = stock_now
        weeks_out = 0
        for (y, w) in horizon:
            d = weekly_demand(pid, y, w, ctx)
            i = ctx["inc_map"].get((pid, y, w), 0.0)
            total_demand += d
            total_incoming += i
            opening = s
            closing = opening - d + i
            if closing < 0:
                lost_units += -closing
                weeks_out += 1
                if stockout_yw is None:
                    stockout_yw = to_yw(y, w)
                s = 0.0
            else:
                s = closing

        rows.append({
            "sku": sku,
            "name": prod.name,
            "tier": prod.tier,
            "category": prod.category,
            "supplier": prod.supplier,
            "stock_now": round(stock_now, 1),
            "wh_stock": round(wh, 1),
            "store_stock": round(store, 1),
            "total_incoming_h": round(total_incoming, 1),
            "total_demand_h":   round(total_demand, 1),
            "stockout_week":    (f"CW{stockout_yw % 100} {stockout_yw // 100}"
                                  if stockout_yw else "SAFE"),
            "weeks_out_of_stock": weeks_out,
            "lost_demand_units":   round(lost_units, 1),
            "cost_price":          round(cost, 4),
            "selling_price":       round(price, 4),
            "lost_sales_cost_eur": round(lost_units * cost,  2),
            "lost_sales_revenue_eur": round(lost_units * price, 2),
        })

    df = pd.DataFrame(rows)
    df = df.sort_values("lost_sales_revenue_eur", ascending=False)
    df.to_csv(OUTDIR / "lost_sales_projection.csv", index=False)

    # Summary by category + tier
    at_risk = df[df["stockout_week"] != "SAFE"]
    by_cat = at_risk.groupby("category").agg(
        skus_at_risk=("sku", "count"),
        lost_units=("lost_demand_units", "sum"),
        lost_cost_eur=("lost_sales_cost_eur", "sum"),
        lost_revenue_eur=("lost_sales_revenue_eur", "sum"),
    ).reset_index().sort_values("lost_revenue_eur", ascending=False)
    by_cat.to_csv(OUTDIR / "lost_sales_by_category.csv", index=False)

    by_tier = at_risk.groupby("tier").agg(
        skus_at_risk=("sku", "count"),
        lost_units=("lost_demand_units", "sum"),
        lost_cost_eur=("lost_sales_cost_eur", "sum"),
        lost_revenue_eur=("lost_sales_revenue_eur", "sum"),
    ).reset_index().sort_values("lost_revenue_eur", ascending=False)
    by_tier.to_csv(OUTDIR / "lost_sales_by_tier.csv", index=False)

    print(f"  → {len(df)} SKUs; {len(at_risk)} at risk; "
          f"€{at_risk['lost_sales_revenue_eur'].sum():,.0f} revenue lost "
          f"(€{at_risk['lost_sales_cost_eur'].sum():,.0f} cost)")
    return {
        "df": df,
        "total_cost_eur":    float(at_risk["lost_sales_cost_eur"].sum()),
        "total_revenue_eur": float(at_risk["lost_sales_revenue_eur"].sum()),
        "n_affected": int(len(at_risk)),
        "top5": at_risk.head(5)[["sku", "name", "category",
                                   "lost_sales_revenue_eur"]].to_dict("records"),
    }


# ─────────────────────────────────────────────────────────────────────────
# Step 2 — Locked cash in overstock
# ─────────────────────────────────────────────────────────────────────────
def step2_locked_cash(ctx) -> dict:
    print("\n[Step 2] Locked cash in overstock ...")
    rows: list[dict] = []

    for prod in ctx["products"].itertuples():
        pid, sku = int(prod.product_id), prod.sku
        wh = ctx["wh_map"].get(pid, 0.0)
        store = ctx["store_map"].get(pid, 0.0)
        stock_now = wh + store
        if stock_now <= 0:
            continue

        avg = ctx["avg_map"].get(pid, 0.0)
        if avg <= 0:
            # If no recent sales, use forecast horizon avg
            fc_vals = [ctx["fc_total"].get((pid, y, w), 0.0)
                       for (y, w) in ctx["horizon"]]
            avg = sum(fc_vals) / len(fc_vals) if fc_vals else 0.0

        cost, _  = resolve_cost(pid, sku, prod, ctx)
        price, _ = resolve_price(pid, sku, prod, ctx)
        cost  = cost  or 0.0
        price = price or (cost * 2 if cost else 0.0)

        lt_weeks, lt_source = resolve_lead_time(prod)
        safety = resolve_safety_weeks(prod.tier or "")
        target_weeks = lt_weeks + REVIEW_PERIOD_WEEKS + safety
        target_units = target_weeks * avg
        excess_units = max(0.0, stock_now - target_units)
        weeks_cover  = (stock_now / avg) if avg > 0 else None

        rows.append({
            "sku": sku,
            "name": prod.name,
            "tier": prod.tier,
            "category": prod.category,
            "supplier": prod.supplier,
            "stock_now":        round(stock_now, 1),
            "wh_stock":         round(wh, 1),
            "store_stock":      round(store, 1),
            "weekly_demand":    round(avg, 2),
            "weeks_cover":      round(weeks_cover, 1) if weeks_cover is not None else None,
            "lead_time_weeks":  round(lt_weeks, 1),
            "lt_source":        lt_source,
            "target_stock_weeks": round(target_weeks, 1),
            "target_stock_units": round(target_units, 1),
            "excess_units":     round(excess_units, 1),
            "cost_price":       round(cost, 4),
            "selling_price":    round(price, 4),
            "locked_cash_cost_eur":    round(excess_units * cost, 2),
            "locked_cash_revenue_eur": round(excess_units * price, 2),
            "category_manager_notes": "",
        })

    df = pd.DataFrame(rows).sort_values("locked_cash_cost_eur", ascending=False)
    df.to_csv(OUTDIR / "locked_cash_overstock.csv", index=False)

    # Category summary
    cat_sum = df.groupby("category").agg(
        sku_count=("sku", "count"),
        total_stock_eur=("locked_cash_cost_eur", lambda s: 0.0),  # placeholder
        total_excess_eur=("locked_cash_cost_eur", "sum"),
    ).reset_index()
    # recompute total_stock_eur properly (stock_now × cost across all SKUs)
    stock_eur = (df["stock_now"] * df["cost_price"]).groupby(df["category"]).sum()
    cat_sum["total_stock_eur"] = cat_sum["category"].map(stock_eur).fillna(0.0)
    cat_sum["excess_pct"] = (cat_sum["total_excess_eur"] /
                              cat_sum["total_stock_eur"].replace(0, 1)).round(3)
    # Top-3 offenders per category
    def _top3(cat: str) -> str:
        sub = df[df["category"] == cat].nlargest(3, "locked_cash_cost_eur")
        return ", ".join(f"{r.sku} (€{r.locked_cash_cost_eur:,.0f})" for r in sub.itertuples())
    cat_sum["top3_offenders"] = cat_sum["category"].apply(_top3)
    cat_sum = cat_sum.sort_values("total_excess_eur", ascending=False)
    cat_sum.to_csv(OUTDIR / "locked_cash_summary_by_category.csv", index=False)

    # Supplier summary (extra — for sourcing decisions)
    sup_sum = df.groupby("supplier").agg(
        sku_count=("sku", "count"),
        total_excess_eur=("locked_cash_cost_eur", "sum"),
    ).reset_index().sort_values("total_excess_eur", ascending=False)
    sup_sum.to_csv(OUTDIR / "locked_cash_summary_by_supplier.csv", index=False)

    print(f"  → €{df['locked_cash_cost_eur'].sum():,.0f} excess cost "
          f"across {len(df[df['excess_units']>0])} overstocked SKUs")
    return {
        "df": df,
        "total_excess_cost_eur":    float(df["locked_cash_cost_eur"].sum()),
        "total_excess_revenue_eur": float(df["locked_cash_revenue_eur"].sum()),
        "by_category": cat_sum.set_index("category")[["total_excess_eur"]]
                              .to_dict()["total_excess_eur"],
        "by_supplier": sup_sum.set_index("supplier")[["total_excess_eur"]]
                              .to_dict()["total_excess_eur"],
    }


# ─────────────────────────────────────────────────────────────────────────
# Step 3 — Contest / July stock-out risk
# ─────────────────────────────────────────────────────────────────────────
def step3_contest(ctx) -> dict:
    print("\n[Step 3] Contest July risk ...")
    horizon_july = [(y, w) for (y, w) in ctx["horizon"]
                    if CONTEST_START_CW <= w <= CONTEST_END_CW]
    if not horizon_july:
        # If horizon doesn't reach July, take last 5 weeks
        horizon_july = ctx["horizon"][-5:]

    rows: list[dict] = []
    for prod in ctx["products"].itertuples():
        pid, sku = int(prod.product_id), prod.sku
        # Focus: Gold + Silver tier (contest typically pushes hero SKUs)
        if prod.tier not in ("01 GOLD", "02 SILVER"):
            continue
        stock_now = ctx["wh_map"].get(pid, 0.0) + ctx["store_map"].get(pid, 0.0)
        if stock_now <= 0:
            continue

        cost, _  = resolve_cost(pid, sku, prod, ctx)
        price, _ = resolve_price(pid, sku, prod, ctx)
        cost  = cost  or 0.0
        price = price or (cost * 2 if cost else 0.0)

        # Sum baseline demand pre-contest weeks + during contest
        pre_contest_demand = sum(
            weekly_demand(pid, y, w, ctx) for (y, w) in ctx["horizon"]
            if (y, w) < horizon_july[0]
        )
        pre_contest_incoming = sum(
            ctx["inc_map"].get((pid, y, w), 0.0) for (y, w) in ctx["horizon"]
            if (y, w) < horizon_july[0]
        )
        available_at_contest = stock_now - pre_contest_demand + pre_contest_incoming

        # Baseline demand during contest
        baseline_july = sum(weekly_demand(pid, y, w, ctx) for (y, w) in horizon_july)
        incoming_during = sum(ctx["inc_map"].get((pid, y, w), 0.0) for (y, w) in horizon_july)

        # Uplift: sku_uplift.csv if available, otherwise 1.5×
        uplift = ctx["uplift_map"].get(sku, {})
        promo_uplift = uplift.get("promo_uplift", 1.5)
        if not promo_uplift or promo_uplift <= 1.0:
            promo_uplift = 1.5

        contest_july = baseline_july * promo_uplift

        # Stockout simulation under each scenario
        def _stockout_label(start_stock: float, dem: float) -> str:
            net = start_stock + incoming_during - dem
            if net >= 0:
                return "SAFE"
            return "AT RISK"

        base_label    = _stockout_label(available_at_contest, baseline_july)
        contest_label = _stockout_label(available_at_contest, contest_july)

        gap_units = max(0.0, contest_july - (available_at_contest + incoming_during))

        rows.append({
            "sku": sku,
            "name": prod.name,
            "tier": prod.tier,
            "category": prod.category,
            "supplier": prod.supplier,
            "stock_now": round(stock_now, 1),
            "incoming_before_contest": round(pre_contest_incoming, 1),
            "available_at_contest":    round(available_at_contest, 1),
            "incoming_during_contest": round(incoming_during, 1),
            "baseline_demand_july":    round(baseline_july, 1),
            "promo_uplift_factor":     round(promo_uplift, 2),
            "contest_demand_july":     round(contest_july, 1),
            "stockout_baseline":       base_label,
            "stockout_contest":        contest_label,
            "gap_units":               round(gap_units, 1),
            "gap_cost_eur":            round(gap_units * cost,  2),
            "gap_revenue_eur":         round(gap_units * price, 2),
        })

    df = pd.DataFrame(rows).sort_values("gap_revenue_eur", ascending=False)
    df.to_csv(OUTDIR / "contest_july_risk.csv", index=False)

    at_risk = df[df["stockout_contest"] == "AT RISK"]
    print(f"  → {len(at_risk)} Gold/Silver SKUs at risk under {df['promo_uplift_factor'].mean():.1f}× "
          f"contest uplift; €{at_risk['gap_revenue_eur'].sum():,.0f} revenue at stake")
    return {
        "df": df,
        "at_risk_skus":           int(len(at_risk)),
        "gap_revenue_eur_baseline": float(df[df["stockout_baseline"] == "AT RISK"]["gap_revenue_eur"].sum()),
        "gap_revenue_eur_contest":  float(at_risk["gap_revenue_eur"].sum()),
        "critical": at_risk.head(10)[["sku", "name", "category",
                                        "gap_revenue_eur"]].to_dict("records"),
    }


# ─────────────────────────────────────────────────────────────────────────
# Step 4 — Slow mover analysis
# ─────────────────────────────────────────────────────────────────────────
def step4_slow_movers(ctx) -> dict:
    print("\n[Step 4] Slow mover analysis ...")
    rows: list[dict] = []

    for prod in ctx["products"].itertuples():
        pid, sku = int(prod.product_id), prod.sku
        stock = ctx["wh_map"].get(pid, 0.0) + ctx["store_map"].get(pid, 0.0)
        if stock <= 0:
            continue   # need stock to classify

        cost, _  = resolve_cost(pid, sku, prod, ctx)
        price, _ = resolve_price(pid, sku, prod, ctx)
        cost  = cost  or 0.0
        price = price or (cost * 2 if cost else 0.0)

        avg = ctx["avg_map"].get(pid, 0.0)
        # Next-13w forecast demand for dead/near-dead test
        fc_13w = sum(ctx["fc_total"].get((pid, y, w), 0.0) for (y, w) in ctx["horizon"])

        lt_weeks, _ = resolve_lead_time(prod)
        safety = resolve_safety_weeks(prod.tier or "")
        target_weeks = lt_weeks + REVIEW_PERIOD_WEEKS + safety
        overstock_thr = lt_weeks * 2
        slow_thr      = lt_weeks * 3

        weeks_cover = (stock / avg) if avg > 0 else None

        # Classification
        if fc_13w == 0 and avg == 0:
            klass = "DEAD_STOCK"
        elif weeks_cover is None:
            klass = "DEAD_STOCK"
        elif weeks_cover > slow_thr:
            klass = "SLOW_MOVER"
        elif weeks_cover > overstock_thr:
            klass = "OVERSTOCK"
        elif weeks_cover < lt_weeks:
            klass = "UNDERSTOCK"
        else:
            klass = "HEALTHY"

        excess_units = max(0.0, stock - target_weeks * avg) if avg > 0 else stock

        rows.append({
            "sku": sku, "name": prod.name,
            "tier": prod.tier, "category": prod.category, "supplier": prod.supplier,
            "stock_units":         round(stock, 1),
            "stock_cost_eur":      round(stock * cost,  2),
            "stock_revenue_eur":   round(stock * price, 2),
            "weekly_demand":       round(avg, 2),
            "weeks_cover":         round(weeks_cover, 1) if weeks_cover is not None else None,
            "lead_time_weeks":     round(lt_weeks, 1),
            "overstock_thr_weeks": round(overstock_thr, 1),
            "slow_thr_weeks":      round(slow_thr, 1),
            "fc_13w_demand":       round(fc_13w, 1),
            "classification":      klass,
            "excess_units":        round(excess_units, 1),
            "excess_cost_eur":     round(excess_units * cost,  2),
            "excess_revenue_eur":  round(excess_units * price, 2),
        })

    df = pd.DataFrame(rows).sort_values("stock_cost_eur", ascending=False)
    df.to_csv(OUTDIR / "slow_mover_analysis.csv", index=False)

    # Pivots — by category and by supplier
    def _pivot(group_col: str) -> pd.DataFrame:
        agg = df.groupby([group_col, "classification"]).agg(
            stock_eur=("stock_cost_eur", "sum"),
        ).unstack(fill_value=0)
        agg.columns = [c[1] for c in agg.columns]
        agg["total_stock_eur"] = agg.sum(axis=1)
        for k in ("DEAD_STOCK", "SLOW_MOVER", "OVERSTOCK", "HEALTHY", "UNDERSTOCK"):
            if k not in agg.columns:
                agg[k] = 0.0
        agg["unhealthy_eur"] = agg["DEAD_STOCK"] + agg["SLOW_MOVER"] + agg["OVERSTOCK"]
        agg["unhealthy_pct"] = (agg["unhealthy_eur"]
                                 / agg["total_stock_eur"].replace(0, 1)).round(3)
        agg = agg.reset_index().sort_values("unhealthy_eur", ascending=False)
        return agg

    by_cat = _pivot("category")
    by_cat["group_by"] = "category"
    by_cat = by_cat.rename(columns={"category": "group_value"})
    by_sup = _pivot("supplier")
    by_sup["group_by"] = "supplier"
    by_sup = by_sup.rename(columns={"supplier": "group_value"})
    summary = pd.concat([by_cat, by_sup], ignore_index=True)
    cols = ["group_by", "group_value", "total_stock_eur", "DEAD_STOCK",
            "SLOW_MOVER", "OVERSTOCK", "HEALTHY", "UNDERSTOCK",
            "unhealthy_eur", "unhealthy_pct"]
    summary[cols].to_csv(OUTDIR / "slow_mover_summary.csv", index=False)

    # Special call-outs
    gadgets_eur = float(df[df["category"].isin(GADGET_CATEGORIES)]["stock_cost_eur"].sum())
    non_food_eur = float(df[df["category"].isin(NON_FOOD_CATEGORIES)]["stock_cost_eur"].sum())

    dead_eur     = float(df[df["classification"] == "DEAD_STOCK"]["stock_cost_eur"].sum())
    slow_eur     = float(df[df["classification"] == "SLOW_MOVER"]["stock_cost_eur"].sum())
    over_eur     = float(df[df["classification"] == "OVERSTOCK"]["stock_cost_eur"].sum())
    health_eur   = float(df[df["classification"] == "HEALTHY"]["stock_cost_eur"].sum())
    under_eur    = float(df[df["classification"] == "UNDERSTOCK"]["stock_cost_eur"].sum())
    print(f"  → dead=€{dead_eur:,.0f}  slow=€{slow_eur:,.0f}  "
          f"over=€{over_eur:,.0f}  healthy=€{health_eur:,.0f}  "
          f"under=€{under_eur:,.0f}")
    print(f"  → gadgets=€{gadgets_eur:,.0f}  non_food=€{non_food_eur:,.0f}")

    return {
        "df": df,
        "dead_stock_eur": dead_eur, "slow_mover_eur": slow_eur,
        "overstock_eur": over_eur, "healthy_eur": health_eur,
        "understock_eur": under_eur,
        "gadgets_eur": gadgets_eur, "non_food_eur": non_food_eur,
    }


# ─────────────────────────────────────────────────────────────────────────
# Step 5 — Volume/Price/Mix/COGS bridge
# ─────────────────────────────────────────────────────────────────────────
def step5_bridge(db, ctx) -> dict:
    print("\n[Step 5] Volume/Price/Mix/COGS bridge ...")

    # Use backtest_fa.csv as plan-vs-actual on the quantity axis
    bt_path = ROOT / "data" / "backtest_fa.csv"
    if not bt_path.exists():
        note = ("backtest_fa.csv missing — bridge cannot be computed.\n"
                "Required columns: sku, year, week, forecast, actual.")
        (OUTDIR / "bridge_feasibility_note.txt").write_text(note)
        return {"available": False, "type": "missing", "components": [],
                "total_margin_variance": 0.0}

    bt = pd.read_csv(bt_path)
    # Filter to a usable window — most recent N weeks with both forecast & actual
    recent_yws = (bt[(bt["forecast"].notna()) & (bt["actual"].notna())]
                  .assign(yw=lambda d: d["year"]*100 + d["week"])
                  .groupby("yw").size().sort_index().tail(8).index.tolist())
    bt = bt[bt["year"]*100 + bt["week"].isin(recent_yws)
            if False else (bt["year"]*100 + bt["week"]).isin(recent_yws)]
    print(f"  bridge window: {len(recent_yws)} weeks, {bt['sku'].nunique()} SKUs")

    # Actual unit price + unit cost from erp_transactions over same window
    if not recent_yws:
        note = "backtest_fa.csv has no rows with both forecast and actual."
        (OUTDIR / "bridge_feasibility_note.txt").write_text(note)
        return {"available": False, "type": "no_overlap", "components": [],
                "total_margin_variance": 0.0}

    min_yw, max_yw = min(recent_yws), max(recent_yws)
    txn_df = pd.read_sql(text("""
        SELECT p.sku,
               EXTRACT(ISOYEAR FROM et.transaction_date)::int AS year,
               EXTRACT(WEEK    FROM et.transaction_date)::int AS week,
               SUM(et.quantity)::float       AS qty,
               SUM(et.total_value)::float    AS rev,
               SUM(et.purchase_value)::float AS pv,
               SUM(et.ruc_eur)::float        AS ruc,
               SUM(et.tax_base)::float       AS tax_base
        FROM erp_transactions et
        JOIN dim_products p ON p.id = et.product_id
        WHERE EXTRACT(ISOYEAR FROM et.transaction_date)::int * 100
            + EXTRACT(WEEK    FROM et.transaction_date)::int
            BETWEEN :min_yw AND :max_yw
        GROUP BY p.sku, year, week
    """), db.bind, params={"min_yw": min_yw, "max_yw": max_yw})

    # Use tax_base (net of VAT) as revenue — matches how erp_prices is set
    txn_df = txn_df.assign(yw=lambda d: d["year"]*100 + d["week"])
    txn_df = txn_df[txn_df["yw"].isin(recent_yws)]

    # Plan price/cost = current snapshot (we don't have historical snapshots)
    cost_map  = {r.sku: float(r.erp_cost or 0.0)
                 for r in ctx["products"].itertuples()
                 if r.erp_cost is not None and not (isinstance(r.erp_cost, float) and math.isnan(r.erp_cost))}
    price_map = {r.sku: float(r.avg_sell_price or 0.0)
                 for r in ctx["products"].itertuples()
                 if r.avg_sell_price is not None and not (isinstance(r.avg_sell_price, float) and math.isnan(r.avg_sell_price))}

    # Merge: bt provides plan_qty (forecast) + actual_qty; txn_df provides realized revenue/cost
    merged = bt.merge(
        txn_df[["sku", "year", "week", "qty", "rev", "pv", "tax_base"]],
        on=["sku", "year", "week"], how="left",
        suffixes=("_bt", "_txn"),
    )
    # Use forecast.actual as the qty fallback (txn might miss some); prefer txn qty
    merged["actual_qty"] = merged["qty"].fillna(merged["actual"])
    merged["plan_qty"]   = merged["forecast"]
    # Actual revenue ex-VAT (preferred); fallback = actual_qty × plan_price
    merged["actual_revenue"] = merged["tax_base"]
    merged["actual_unit_price"] = merged["actual_revenue"] / merged["actual_qty"]
    merged["actual_unit_cost"]  = merged["pv"] / merged["actual_qty"]

    merged["plan_price"] = merged["sku"].map(price_map).fillna(0.0)
    merged["plan_cost"]  = merged["sku"].map(cost_map).fillna(0.0)
    merged["plan_revenue"] = merged["plan_qty"] * merged["plan_price"]
    merged["plan_cogs"]    = merged["plan_qty"] * merged["plan_cost"]
    merged["actual_cogs"]  = merged["actual_qty"] * merged["actual_unit_cost"].fillna(merged["plan_cost"])

    # Decomposition (per SKU-week)
    merged["volume_effect_eur"] = (merged["actual_qty"]   - merged["plan_qty"]  ) * merged["plan_price"]
    merged["price_effect_eur"]  = (merged["actual_unit_price"].fillna(merged["plan_price"])
                                     - merged["plan_price"]) * merged["actual_qty"]
    merged["mix_effect_eur"]    = (merged["actual_revenue"].fillna(0)
                                     - merged["plan_revenue"]
                                     - merged["volume_effect_eur"]
                                     - merged["price_effect_eur"])
    merged["cogs_effect_eur"]   = -(merged["actual_cogs"] - merged["plan_cogs"])
    merged["margin_effect_total_eur"] = (merged["volume_effect_eur"]
                                           + merged["price_effect_eur"]
                                           + merged["mix_effect_eur"]
                                           + merged["cogs_effect_eur"])

    # Attach tier + category
    prod_map = ctx["products"][["sku", "tier", "category"]].set_index("sku")
    merged = merged.join(prod_map, on="sku")

    # Per-SKU per-week output
    period_label = lambda y, w: f"{int(y)}W{int(w):02d}"
    merged["period"] = [period_label(y, w)
                         for y, w in zip(merged["year"], merged["week"])]
    out_cols = ["sku", "tier", "category", "period",
                 "plan_qty", "actual_qty",
                 "plan_price", "actual_unit_price",
                 "plan_revenue", "actual_revenue",
                 "volume_effect_eur", "price_effect_eur", "mix_effect_eur",
                 "plan_cogs", "actual_cogs", "cogs_effect_eur",
                 "margin_effect_total_eur"]
    bridge_df = merged[out_cols].copy()
    bridge_df = bridge_df.round({c: 2 for c in out_cols
                                  if c not in ("sku", "tier", "category", "period")})
    bridge_df.to_csv(OUTDIR / "bridge_analysis.csv", index=False)

    # Summary by category + total
    by_cat = merged.groupby("category").agg(
        volume_effect=("volume_effect_eur", "sum"),
        price_effect=("price_effect_eur",  "sum"),
        mix_effect=("mix_effect_eur",      "sum"),
        cogs_effect=("cogs_effect_eur",    "sum"),
        total_margin_effect=("margin_effect_total_eur", "sum"),
    ).reset_index().round(0)
    by_cat["level"] = "category"; by_cat = by_cat.rename(columns={"category": "name"})

    total_row = pd.DataFrame([{
        "level": "TOTAL", "name": "all",
        "volume_effect": merged["volume_effect_eur"].sum(),
        "price_effect":  merged["price_effect_eur"].sum(),
        "mix_effect":    merged["mix_effect_eur"].sum(),
        "cogs_effect":   merged["cogs_effect_eur"].sum(),
        "total_margin_effect": merged["margin_effect_total_eur"].sum(),
    }]).round(0)
    summary = pd.concat([total_row, by_cat], ignore_index=True)
    summary = summary[["level", "name", "volume_effect", "price_effect",
                        "mix_effect", "cogs_effect", "total_margin_effect"]]
    summary.to_csv(OUTDIR / "bridge_summary.csv", index=False)

    # Waterfall JSON
    plan_margin = float(merged["plan_revenue"].sum() - merged["plan_cogs"].sum())
    actual_margin = float(merged["actual_revenue"].fillna(0).sum()
                           - merged["actual_cogs"].sum())
    waterfall = {
        "window_weeks": [int(yw) for yw in recent_yws],
        "n_skus": int(merged["sku"].nunique()),
        "bars": [
            {"label": "Plan margin",   "value": round(plan_margin, 0), "type": "base"},
            {"label": "Volume",        "value": round(float(merged["volume_effect_eur"].sum()), 0),
                "type": "increase" if merged["volume_effect_eur"].sum() >= 0 else "decrease"},
            {"label": "Price",         "value": round(float(merged["price_effect_eur"].sum()), 0),
                "type": "increase" if merged["price_effect_eur"].sum() >= 0 else "decrease"},
            {"label": "Mix",           "value": round(float(merged["mix_effect_eur"].sum()), 0),
                "type": "increase" if merged["mix_effect_eur"].sum() >= 0 else "decrease"},
            {"label": "COGS",          "value": round(float(merged["cogs_effect_eur"].sum()), 0),
                "type": "increase" if merged["cogs_effect_eur"].sum() >= 0 else "decrease"},
            {"label": "Actual margin", "value": round(actual_margin, 0), "type": "total"},
        ],
        "by_category": {
            r["name"]: {
                "volume": float(r["volume_effect"]),
                "price":  float(r["price_effect"]),
                "mix":    float(r["mix_effect"]),
                "cogs":   float(r["cogs_effect"]),
                "total":  float(r["total_margin_effect"]),
            }
            for r in summary[summary["level"] == "category"].to_dict("records")
        },
    }
    with open(OUTDIR / "bridge_waterfall_data.json", "w", encoding="utf-8") as fh:
        json.dump(waterfall, fh, indent=2)

    note = (
        "Bridge built: Volume / Price / Mix / COGS over last 8 backtest weeks.\n"
        f"  Window: {recent_yws}\n"
        f"  SKUs:   {merged['sku'].nunique()}\n"
        "Plan = forecast.total (latest backtest row per SKU-week).\n"
        "Plan price/cost = current erp_prices.avg_sell_price / erp_costs.cost_price\n"
        "  (we don't have historical price/cost snapshots — only one valid_from per SKU).\n"
        "Actual qty = erp_transactions.quantity; falls back to backtest_fa.actual.\n"
        "Actual revenue = erp_transactions.tax_base (net of VAT).\n"
        "Actual unit cost = erp_transactions.purchase_value / quantity.\n"
        "\nLimitations:\n"
        "  - No historical price snapshots → price_effect is realized-vs-current,\n"
        "    not realized-vs-period-plan.\n"
        "  - No historical cost snapshots → COGS effect compares realized cost\n"
        "    vs the latest standard cost, not vs a prior-period cost.\n"
        "  - Mix effect is the residual; assumes plan price/cost reflects 'plan mix'.\n"
        "\nFor next S&OP, Finance should provide:\n"
        "  - monthly plan revenue + qty per SKU\n"
        "  - cost snapshots (cost_price, valid_from) — one per cost change\n"
        "  - price snapshots (avg_sell_price, valid_from) — one per price change\n"
    )
    (OUTDIR / "bridge_feasibility_note.txt").write_text(note, encoding="utf-8")

    total_var = float(merged["margin_effect_total_eur"].sum())
    print(f"  → bridge built; total margin variance €{total_var:,.0f} "
          f"over {len(recent_yws)} weeks")
    return {
        "available": True, "type": "full",
        "components": ["volume", "price", "mix", "cogs"],
        "total_margin_variance": total_var,
        "plan_margin": plan_margin,
        "actual_margin": actual_margin,
        "window_weeks": recent_yws,
    }


# ─────────────────────────────────────────────────────────────────────────
# Step 6 — Master dashboard JSON
# ─────────────────────────────────────────────────────────────────────────
def step6_dashboard(ctx, r1, r2, r3, r4, r5) -> None:
    print("\n[Step 6] Master dashboard JSON ...")
    out = {
        "generated_at": datetime.now().isoformat(timespec="seconds"),
        "anchor_year_week": f"{ctx['cur_year']}W{ctx['cur_week']:02d}",
        "horizon_weeks": HORIZON_WEEKS,
        "n_products_in_scope": int(len(ctx["products"])),

        "lost_sales": {
            "total_cost_eur":    r1["total_cost_eur"],
            "total_revenue_eur": r1["total_revenue_eur"],
            "affected_skus":     r1["n_affected"],
            "top5_by_revenue":   r1["top5"],
        },
        "locked_cash": {
            "total_excess_cost_eur":    r2["total_excess_cost_eur"],
            "total_excess_revenue_eur": r2["total_excess_revenue_eur"],
            "by_category": dict(sorted(r2["by_category"].items(),
                                        key=lambda kv: -kv[1])[:15]),
            "by_supplier": dict(sorted(r2["by_supplier"].items(),
                                        key=lambda kv: -kv[1])[:15]),
        },
        "contest_risk": {
            "at_risk_skus": r3["at_risk_skus"],
            "gap_revenue_eur_baseline": r3["gap_revenue_eur_baseline"],
            "gap_revenue_eur_contest":  r3["gap_revenue_eur_contest"],
            "critical_skus": r3["critical"],
        },
        "slow_movers": {
            "dead_stock_eur":  r4["dead_stock_eur"],
            "slow_mover_eur":  r4["slow_mover_eur"],
            "overstock_eur":   r4["overstock_eur"],
            "healthy_eur":     r4["healthy_eur"],
            "understock_eur":  r4["understock_eur"],
            "gadgets_eur":     r4["gadgets_eur"],
            "non_food_eur":    r4["non_food_eur"],
        },
        "bridge": {
            "available": r5["available"],
            "type": r5["type"],
            "components": r5["components"],
            "total_margin_variance": r5.get("total_margin_variance", 0.0),
            "plan_margin":   r5.get("plan_margin"),
            "actual_margin": r5.get("actual_margin"),
            "window_weeks":  r5.get("window_weeks", []),
        },
        "feasibility_matrix": [
            {"request": "13w inventory projection (€)", "status": "DONE",
             "note": "Built earlier — Stock Projection chart + Scenario Planner (now aligned)."},
            {"request": "Lost sales (units + cost + revenue)", "status": "DONE",
             "note": "Walk-forward simulation, all SKUs, WH+stores."},
            {"request": "Locked cash in overstock", "status": "DONE",
             "note": "Per-SKU + by-category + by-supplier; real LT where available, tier fallback otherwise."},
            {"request": "Contest July risk", "status": "PARTIAL",
             "note": "No contest SKU list available. Used sku_uplift.csv promo_uplift factors; fallback 1.5× on Gold/Silver. Real contest volumes still needed from sales team."},
            {"request": "Vol/Price/Mix/COGS bridge", "status": "PARTIAL" if r5["available"] else "BLOCKED",
             "note": ("Full 4-component decomposition over last 8 backtest weeks. Plan price/cost = current snapshot — historical snapshots would tighten the price effect."
                      if r5["available"] else "backtest_fa.csv missing — see bridge_feasibility_note.txt")},
            {"request": "Slow movers", "status": "DONE",
             "note": "LT-relative thresholds (2×/3× LT), all suppliers, full pivot by category + supplier."},
        ],
    }
    with open(OUTDIR / "cfo_dashboard_data.json", "w", encoding="utf-8") as fh:
        json.dump(out, fh, indent=2, default=str)
    print(f"  → {OUTDIR / 'cfo_dashboard_data.json'}")


# ─────────────────────────────────────────────────────────────────────────
# Main
# ─────────────────────────────────────────────────────────────────────────
def step5b_bridge_monthly(db) -> dict:
    """Reuse the monthly-bridge service to write the additional output files
    the user expects: bridge_by_month.csv, bridge_summary_by_month.csv,
    bridge_summary_by_category.csv, bridge_top_movers.csv,
    bridge_waterfall_data.json (with monthly structure)."""
    print("\n[Step 5b] Monthly bridge (V/P/M/COGS per month) ...")
    import json as _json
    from backend.services.finance_service import (
        FinanceContext, report_bridge_monthly,
    )
    fctx = FinanceContext(db)
    res = report_bridge_monthly(fctx, db)
    if not res["available"]:
        print("  → no monthly bridge data available")
        return res

    # Per-SKU per-month detail
    pd.DataFrame(res["per_sku_rows"]).to_csv(
        OUTDIR / "bridge_by_month.csv", index=False,
    )
    # Summary by month
    summary_rows = [{
        "month": m["month"], "tier": m["tier"], "n_skus": m["n_skus"],
        "plan_margin_total":   m["plan_margin_total"],
        "actual_margin_total": m["actual_margin_total"],
        "total_variance":      m["total_variance"],
        "volume_effect":       m["volume_effect"],
        "price_effect":        m["price_effect"],
        "mix_effect":          m["mix_effect"],
        "cogs_effect":         m["cogs_effect"],
    } for m in res["months"]]
    pd.DataFrame(summary_rows).to_csv(
        OUTDIR / "bridge_summary_by_month.csv", index=False,
    )
    # By category
    by_cat_rows: list[dict] = []
    for m in res["months"]:
        for c in m["by_category"]:
            by_cat_rows.append({"month": m["month"], **c})
    pd.DataFrame(by_cat_rows).to_csv(
        OUTDIR / "bridge_summary_by_category.csv", index=False,
    )
    # Top movers
    pd.DataFrame(res["top_movers_rows"]).to_csv(
        OUTDIR / "bridge_top_movers.csv", index=False,
    )
    # Waterfall JSON (one entry per month + YTD)
    waterfall = {
        "months": [m["month"] for m in res["months"]],
        "by_month": {},
    }
    for m in res["months"]:
        waterfall["by_month"][m["month"]] = {
            "tier": m["tier"],
            "n_skus": m["n_skus"],
            "bars": [
                {"label": "Plan margin",   "value": m["plan_margin_total"],   "type": "base"},
                {"label": "Volume",        "value": m["volume_effect"],
                  "type": "increase" if m["volume_effect"] >= 0 else "decrease"},
                {"label": "Price",         "value": m["price_effect"],
                  "type": "increase" if m["price_effect"]  >= 0 else "decrease"},
                {"label": "Mix",           "value": m["mix_effect"],
                  "type": "increase" if m["mix_effect"]    >= 0 else "decrease"},
                {"label": "COGS",          "value": m["cogs_effect"],
                  "type": "increase" if m["cogs_effect"]   >= 0 else "decrease"},
                {"label": "Actual margin", "value": m["actual_margin_total"], "type": "total"},
            ],
            "by_category": {
                c["category"]: {
                    "volume": c["volume_effect"], "price": c["price_effect"],
                    "mix":    c["mix_effect"],    "cogs":  c["cogs_effect"],
                    "total":  c["total_variance"],
                } for c in m["by_category"]
            },
            "top_positive_skus": m["top_positive_skus"],
            "top_negative_skus": m["top_negative_skus"],
        }
    if res["ytd"]:
        waterfall["ytd"] = {
            "year": res["ytd"]["year"],
            "n_months": res["ytd"]["n_months"],
            "bars": [
                {"label": "Plan margin",   "value": res["ytd"]["plan_margin_total"],   "type": "base"},
                {"label": "Volume",        "value": res["ytd"]["volume_effect"],
                  "type": "increase" if res["ytd"]["volume_effect"] >= 0 else "decrease"},
                {"label": "Price",         "value": res["ytd"]["price_effect"],
                  "type": "increase" if res["ytd"]["price_effect"]  >= 0 else "decrease"},
                {"label": "Mix",           "value": res["ytd"]["mix_effect"],
                  "type": "increase" if res["ytd"]["mix_effect"]    >= 0 else "decrease"},
                {"label": "COGS",          "value": res["ytd"]["cogs_effect"],
                  "type": "increase" if res["ytd"]["cogs_effect"]   >= 0 else "decrease"},
                {"label": "Actual margin", "value": res["ytd"]["actual_margin_total"], "type": "total"},
            ],
            "trend": res["ytd"]["trend"],
        }
    with open(OUTDIR / "bridge_waterfall_data.json", "w", encoding="utf-8") as fh:
        _json.dump(waterfall, fh, indent=2, default=str)
    print(f"  → {len(res['months'])} months written ("
          f"{sum(1 for m in res['months'] if m['tier'] == 'FULL')} full, "
          f"{sum(1 for m in res['months'] if m['tier'] == 'VOLUME_ONLY')} volume-only)")
    return res


def main() -> None:
    db = SessionLocal()
    try:
        print(f"CFO audit — writing to {OUTDIR}")
        ctx = load_context(db)
        r1 = step1_lost_sales(ctx)
        r2 = step2_locked_cash(ctx)
        r3 = step3_contest(ctx)
        r4 = step4_slow_movers(ctx)
        r5 = step5_bridge(db, ctx)
        r5b = step5b_bridge_monthly(db)
        step6_dashboard(ctx, r1, r2, r3, r4, r5)

        # Index of outputs
        outputs = sorted(OUTDIR.glob("*"))
        print(f"\n--- Output inventory ({len(outputs)} files) ---")
        for p in outputs:
            print(f"  {p.name}  ({p.stat().st_size / 1024:.1f} KB)")
    finally:
        db.close()


if __name__ == "__main__":
    main()
