"""Bridge v2 investigation — Problem A (channel mismatch), Problem B (COGS),
Problem C (promo list price), + plan-qty channel split.

Outputs concrete numbers the user needs to see BEFORE we touch any code.

Period: April + May 2026 (CW14-22 by calendar month).
"""
import sys
sys.stdout.reconfigure(encoding="utf-8")
from backend.models.database import SessionLocal
from sqlalchemy import text
import pandas as pd

db = SessionLocal()


def header(title: str) -> None:
    print()
    print("═" * 80)
    print(title)
    print("═" * 80)


# ────────────────────────────────────────────────────────────────────
# PROBLEM A — Channel mismatch in price effect
# ────────────────────────────────────────────────────────────────────
header("PROBLEM A — Channel mismatch in price effect")

# A.1 Wholesale share of volume by month
print("\nA.1 — Volume share by channel (Apr + May 2026)")
print("  channel       Apr qty       Apr rev       May qty       May rev")
print("  ──────────────────────────────────────────────────────────────────")
rows = db.execute(text("""
    SELECT cm.channel,
           SUM(CASE WHEN EXTRACT(MONTH FROM et.transaction_date) = 4 THEN et.quantity ELSE 0 END)::float    AS apr_qty,
           SUM(CASE WHEN EXTRACT(MONTH FROM et.transaction_date) = 4 THEN et.total_value ELSE 0 END)::float AS apr_rev,
           SUM(CASE WHEN EXTRACT(MONTH FROM et.transaction_date) = 5 THEN et.quantity ELSE 0 END)::float    AS may_qty,
           SUM(CASE WHEN EXTRACT(MONTH FROM et.transaction_date) = 5 THEN et.total_value ELSE 0 END)::float AS may_rev
    FROM erp_transactions et
    JOIN lookup_channel_map cm ON cm.id = et.channel_map_id
    WHERE EXTRACT(YEAR FROM et.transaction_date) = 2026
      AND EXTRACT(MONTH FROM et.transaction_date) IN (4, 5)
    GROUP BY cm.channel
    ORDER BY apr_rev DESC
""")).mappings().all()
apr_total_qty = sum(r["apr_qty"] for r in rows)
may_total_qty = sum(r["may_qty"] for r in rows)
apr_total_rev = sum(r["apr_rev"] for r in rows)
may_total_rev = sum(r["may_rev"] for r in rows)
for r in rows:
    print(f"  {r['channel']:<12s}  {r['apr_qty']:>8.0f}     EUR{r['apr_rev']:>10.0f}     {r['may_qty']:>8.0f}     EUR{r['may_rev']:>10.0f}")
print(f"  {'TOTAL':<12s}  {apr_total_qty:>8.0f}     EUR{apr_total_rev:>10.0f}     {may_total_qty:>8.0f}     EUR{may_total_rev:>10.0f}")
print()
for r in rows:
    ws_apr = r["apr_qty"] / apr_total_qty * 100 if apr_total_qty else 0
    ws_may = r["may_qty"] / may_total_qty * 100 if may_total_qty else 0
    print(f"  {r['channel']:<12s}  Apr share: {ws_apr:>5.1f}% qty   May share: {ws_may:>5.1f}% qty")

# A.2 Avg price per unit per channel
print("\nA.2 — Avg unit price by channel (tax_base/quantity, Apr+May)")
rows = db.execute(text("""
    SELECT cm.channel,
           EXTRACT(MONTH FROM et.transaction_date)::int AS mo,
           SUM(et.quantity)::float  AS qty,
           SUM(et.tax_base)::float  AS rev_net,
           SUM(et.total_value)::float AS rev_gross,
           SUM(et.purchase_value)::float AS pv
    FROM erp_transactions et
    JOIN lookup_channel_map cm ON cm.id = et.channel_map_id
    WHERE EXTRACT(YEAR FROM et.transaction_date) = 2026
      AND EXTRACT(MONTH FROM et.transaction_date) IN (4, 5)
    GROUP BY cm.channel, mo
    ORDER BY cm.channel, mo
""")).mappings().all()
for r in rows:
    mo = "Apr" if r["mo"] == 4 else "May"
    avg_net = r["rev_net"] / r["qty"] if r["qty"] else 0
    avg_gross = r["rev_gross"] / r["qty"] if r["qty"] else 0
    print(f"  {r['channel']:<12s} {mo}: avg_unit_price_net=EUR{avg_net:.2f}  gross=EUR{avg_gross:.2f}  (qty={r['qty']:,.0f})")

# A.3 What does the bridge currently use for plan_price?
print("\nA.3 — plan_price source (erp_prices.avg_sell_price = retail+webshop 12w blend, wholesale EXCLUDED)")
r = db.execute(text("""
    SELECT
        AVG(ep.avg_sell_price)::float        AS avg_avg_sell,
        AVG(ep.normal_retail_ppp)::float     AS avg_retail_ppp,
        AVG(ep.normal_webshop_ppp)::float    AS avg_webshop_ppp,
        AVG(sp.vpc)::float                   AS avg_vpc,
        COUNT(*) AS n
    FROM erp_prices ep
    LEFT JOIN sku_planning sp ON sp.product_id = ep.product_id
    WHERE ep.avg_sell_price > 0
""")).mappings().first()
print(f"  Catalog-wide avg of erp_prices.avg_sell_price (PLAN PRICE SOURCE): EUR {r['avg_avg_sell']:.2f}")
print(f"  Catalog-wide avg of normal_retail_ppp:                              EUR {r['avg_retail_ppp']:.2f}")
print(f"  Catalog-wide avg of normal_webshop_ppp:                             EUR {r['avg_webshop_ppp']:.2f}")
print(f"  Catalog-wide avg of sku_planning.vpc (wholesale list):              EUR {r['avg_vpc']:.2f}")

# A.4 Simulate price effect IF actual_price was B2C-only vs current full-blend
print("\nA.4 — Simulated price effect: full-blend (current) vs B2C-only")
# Replicate the current bridge: blended actual_price per SKU per month
sim = pd.read_sql(text("""
    WITH txn_all AS (
        SELECT p.sku,
               EXTRACT(MONTH FROM et.transaction_date)::int AS mo,
               SUM(et.quantity)::float                   AS qty,
               SUM(et.tax_base)::float                   AS rev_net,
               SUM(CASE WHEN cm.channel IN ('retail','webshop')
                        THEN et.quantity ELSE 0 END)::float AS qty_b2c,
               SUM(CASE WHEN cm.channel IN ('retail','webshop')
                        THEN et.tax_base ELSE 0 END)::float AS rev_net_b2c,
               SUM(CASE WHEN cm.channel = 'wholesale'
                        THEN et.quantity ELSE 0 END)::float AS qty_ws,
               SUM(CASE WHEN cm.channel = 'wholesale'
                        THEN et.tax_base ELSE 0 END)::float AS rev_net_ws
        FROM erp_transactions et
        JOIN dim_products p ON p.id = et.product_id
        JOIN lookup_channel_map cm ON cm.id = et.channel_map_id
        WHERE EXTRACT(YEAR FROM et.transaction_date) = 2026
          AND EXTRACT(MONTH FROM et.transaction_date) IN (4, 5)
        GROUP BY p.sku, mo
    )
    SELECT t.*,
           ep.avg_sell_price::float AS plan_price
    FROM txn_all t
    LEFT JOIN dim_products p ON p.sku = t.sku
    LEFT JOIN erp_prices ep ON ep.product_id = p.id
"""), db.bind)

def _price_effect(df, qty_col, rev_col, label):
    df = df[df[qty_col] > 0].copy()
    df["actual_unit_price"] = df[rev_col] / df[qty_col]
    df["plan_price"] = df["plan_price"].fillna(0)
    df["effect"] = (df["actual_unit_price"] - df["plan_price"]) * df[qty_col]
    return float(df["effect"].sum())

for mo, mo_name in [(4, "April"), (5, "May")]:
    df = sim[sim["mo"] == mo].copy()
    full_blend  = _price_effect(df, "qty",     "rev_net",     "full")
    b2c_only    = _price_effect(df, "qty_b2c", "rev_net_b2c", "b2c")
    ws_only     = _price_effect(df, "qty_ws",  "rev_net_ws",  "ws")
    structural = full_blend - b2c_only
    print(f"\n  {mo_name} 2026:")
    print(f"    Full-blend price effect (current bridge):  EUR{full_blend:>12,.0f}")
    print(f"    B2C-only price effect:                     EUR{b2c_only:>12,.0f}")
    print(f"    Wholesale-only price effect:               EUR{ws_only:>12,.0f}")
    print(f"    'Structural channel' component (full-B2C): EUR{structural:>12,.0f}   ← currently buried in Price")


# ────────────────────────────────────────────────────────────────────
# PROBLEM B — COGS +€343K in May
# ────────────────────────────────────────────────────────────────────
header("PROBLEM B — COGS +EUR343K in May 2026")

# B.1 Top 20 SKUs in May by volume — compare plan_cost vs actual_cost per unit
# plan_cost = NabavneCijene prior-3-month avg → already loaded into ctx.cost_history
# actual_cost = purchase_value/quantity from erp_transactions
print("\nB.1 — Top 20 SKUs by May volume — plan vs actual unit cost")

# Load NabavneCijene for plan_cost lookup
from pathlib import Path
nc_path = Path("data/NabavneCijene.xlsx")
nc = pd.read_excel(nc_path)
nc = nc.rename(columns={"Šifra": "sku", "Datum": "date", "Količina": "qty",
                         "Nabavna vrijednost €": "value"})
nc["sku"] = nc["sku"].astype(str).str.strip()
nc["date"] = pd.to_datetime(nc["date"], errors="coerce")
nc = nc[nc["sku"].notna() & nc["qty"].notna() & (nc["qty"] > 0) & nc["date"].notna()]
nc["year"] = nc["date"].dt.year
nc["month"] = nc["date"].dt.month
nc_monthly = (nc.groupby(["sku", "year", "month"])
              .agg(qty=("qty", "sum"), value=("value", "sum"))
              .reset_index())
nc_monthly["unit_cost"] = nc_monthly["value"] / nc_monthly["qty"]
cost_hist = {(r.sku, int(r.year), int(r.month)): float(r.unit_cost)
             for r in nc_monthly.itertuples()}

def prior_3m_avg(sku, year, month):
    vals = []
    y, m = year, month
    for _ in range(3):
        m -= 1
        if m == 0:
            m = 12; y -= 1
        v = cost_hist.get((sku, y, m))
        if v is not None and v > 0:
            vals.append(v)
    return sum(vals) / len(vals) if vals else None

may_skus = pd.read_sql(text("""
    SELECT p.sku,
           SUM(et.quantity)::float       AS qty,
           SUM(et.total_value)::float    AS rev,
           SUM(et.purchase_value)::float AS pv,
           COUNT(et.id)                  AS n_lines
    FROM erp_transactions et
    JOIN dim_products p ON p.id = et.product_id
    WHERE et.transaction_date >= '2026-05-01'
      AND et.transaction_date <  '2026-06-01'
    GROUP BY p.sku
    ORDER BY qty DESC LIMIT 20
"""), db.bind)

print(f"  {'sku':<14s}  {'qty':>8s}  {'actual/u':>10s}  {'plan/u':>10s}  {'diff/u':>10s}  {'cogs_eff_eur':>14s}")
print(f"  {'─'*78}")
total_cogs = 0.0
for _, r in may_skus.iterrows():
    actual_u = r["pv"] / r["qty"] if r["qty"] else 0
    plan_u = prior_3m_avg(r["sku"], 2026, 5)
    if plan_u is None:
        print(f"  {r['sku']:<14s}  {r['qty']:>8.0f}  {actual_u:>10.4f}  {'NULL':>10s}  {'(no hist)':>10s}  {'':>14s}")
        continue
    diff = plan_u - actual_u
    cogs_eff = diff * r["qty"]  # positive if actual < plan
    total_cogs += cogs_eff
    print(f"  {r['sku']:<14s}  {r['qty']:>8.0f}  {actual_u:>10.4f}  {plan_u:>10.4f}  {diff:>+10.4f}  {cogs_eff:>+14,.0f}")
print(f"\n  Subtotal of top 20 COGS effect: EUR{total_cogs:>+14,.0f}")
print(f"  (Bridge reported total May COGS effect = EUR +342,594 — see how much top-20 explains)")

# B.2 How many SKUs have NULL plan_cost (fall back to erp_costs)?
print("\nB.2 — Fallback frequency in May")
all_may_skus = pd.read_sql(text("""
    SELECT p.sku,
           SUM(et.quantity)::float AS qty,
           SUM(et.purchase_value)::float AS pv
    FROM erp_transactions et
    JOIN dim_products p ON p.id = et.product_id
    WHERE et.transaction_date >= '2026-05-01'
      AND et.transaction_date <  '2026-06-01'
    GROUP BY p.sku
"""), db.bind)
n_total = len(all_may_skus)
n_with_pv = (all_may_skus["pv"] > 0).sum()
n_with_plan = sum(1 for sku in all_may_skus["sku"]
                   if prior_3m_avg(sku, 2026, 5) is not None)
n_fallback = n_total - n_with_plan
print(f"  Total SKUs sold in May:           {n_total:,}")
print(f"  ...with non-zero purchase_value:  {n_with_pv:,}  ({100*n_with_pv/n_total:.1f}%)")
print(f"  ...with NabavneCijene plan_cost:  {n_with_plan:,}  ({100*n_with_plan/n_total:.1f}%)")
print(f"  ...falling back to erp_costs:     {n_fallback:,}  ({100*n_fallback/n_total:.1f}%)")

# B.3 Distribution of actual vs plan cost gap
print("\nB.3 — Distribution of (actual − plan) cost gap across all SKUs with both signals")
gaps = []
for _, r in all_may_skus.iterrows():
    if r["qty"] <= 0 or r["pv"] <= 0:
        continue
    actual_u = r["pv"] / r["qty"]
    plan_u = prior_3m_avg(r["sku"], 2026, 5)
    if plan_u is None or plan_u <= 0:
        continue
    gaps.append({"sku": r["sku"], "qty": r["qty"],
                 "actual_u": actual_u, "plan_u": plan_u,
                 "gap": plan_u - actual_u,  # positive = actual lower = COGS savings
                 "gap_pct": (plan_u - actual_u) / plan_u * 100,
                 "cogs_eff": (plan_u - actual_u) * r["qty"]})
gaps_df = pd.DataFrame(gaps)
print(f"  SKUs with both plan + actual:  {len(gaps_df):,}")
print(f"  Actual < plan (savings):       {(gaps_df['gap'] > 0).sum():,}")
print(f"  Actual > plan (over-spend):    {(gaps_df['gap'] < 0).sum():,}")
print(f"  Actual = plan (within 1%):     {((gaps_df['gap_pct'].abs()) < 1).sum():,}")
print(f"  Median gap_pct:                {gaps_df['gap_pct'].median():.2f}%")
print(f"  Total COGS effect from gaps:   EUR{gaps_df['cogs_eff'].sum():+,.0f}")


# ────────────────────────────────────────────────────────────────────
# PROBLEM C — Promo investment uses avg_sell_price (which includes promos)
# ────────────────────────────────────────────────────────────────────
header("PROBLEM C — Promo list-price source")

# C.1 For promo SKUs in May: compare avg_sell_price vs normal_retail_ppp
print("\nC.1 — Promo SKUs in May 2026: avg_sell_price vs normal_retail_ppp")
promo_skus = pd.read_sql(text("""
    SELECT DISTINCT p.sku
    FROM promo_policy_items ppi
    JOIN dim_products p ON p.id = ppi.product_id
    WHERE ppi.valid_from <= '2026-05-31'
      AND ppi.valid_to   >= '2026-05-01'
"""), db.bind)
sku_list = promo_skus["sku"].tolist()

if sku_list:
    prices = pd.read_sql(text("""
        SELECT p.sku,
               ep.avg_sell_price::float    AS avg_sell,
               ep.normal_retail_ppp::float AS retail_ppp,
               ep.normal_webshop_ppp::float AS webshop_ppp
        FROM dim_products p
        LEFT JOIN erp_prices ep ON ep.product_id = p.id
        WHERE p.sku = ANY(:sk)
    """), db.bind, params={"sk": sku_list})
    prices["delta_retail"] = prices["retail_ppp"] - prices["avg_sell"]
    prices["delta_pct_retail"] = (prices["delta_retail"] / prices["avg_sell"] * 100)
    print(f"  {len(prices)} promo SKUs in May")
    print(f"  Avg avg_sell_price:        EUR {prices['avg_sell'].mean():.2f}")
    print(f"  Avg normal_retail_ppp:     EUR {prices['retail_ppp'].mean():.2f}")
    print(f"  Avg normal_webshop_ppp:    EUR {prices['webshop_ppp'].mean():.2f}")
    print(f"  Avg delta (retail_ppp - avg_sell): EUR {prices['delta_retail'].mean():+.2f}  ({prices['delta_pct_retail'].mean():+.1f}%)")
    print()
    print("  Top 10 promo SKUs by gap (retail_ppp > avg_sell — these are most affected):")
    top = prices.nlargest(10, "delta_retail")
    for _, r in top.iterrows():
        print(f"    {r['sku']:<12s}  avg_sell=EUR{r['avg_sell']:>6.2f}  retail_ppp=EUR{r['retail_ppp']:>6.2f}  delta=+EUR{r['delta_retail']:.2f}")

    # C.2 Recompute promo investment with normal_retail_ppp
    print("\nC.2 — Recompute promo investment with normal_retail_ppp as list (vs avg_sell)")
    promo_tx = pd.read_sql(text("""
        SELECT p.sku, et.quantity::float AS qty, et.total_value::float AS rev,
               cm.channel
        FROM erp_transactions et
        JOIN dim_products p ON p.id = et.product_id
        JOIN lookup_channel_map cm ON cm.id = et.channel_map_id
        JOIN promo_policy_items ppi
            ON ppi.product_id = et.product_id
           AND et.transaction_date BETWEEN ppi.valid_from AND ppi.valid_to
        WHERE et.transaction_date >= '2026-05-01'
          AND et.transaction_date <  '2026-06-01'
          AND cm.channel IN ('retail', 'webshop')
    """), db.bind)
    promo_tx = promo_tx.merge(prices[["sku", "avg_sell", "retail_ppp", "webshop_ppp"]], on="sku", how="left")
    # Use channel-appropriate list price
    promo_tx["list_per_u"] = promo_tx.apply(
        lambda r: r["retail_ppp"] if r["channel"] == "retail" else r["webshop_ppp"],
        axis=1,
    )
    # Old method: avg_sell_price as list
    inv_old = -float(((promo_tx["avg_sell"] * promo_tx["qty"]) - promo_tx["rev"]).sum())
    # New method: normal_retail/webshop_ppp as list
    inv_new = -float(((promo_tx["list_per_u"] * promo_tx["qty"]) - promo_tx["rev"]).sum())
    print(f"  Promo qty B2C in May: {promo_tx['qty'].sum():,.0f}")
    print(f"  Promo revenue in May: EUR {promo_tx['rev'].sum():,.0f}")
    print(f"  Investment with avg_sell_price (current): EUR {inv_old:>+12,.0f}")
    print(f"  Investment with normal_*_ppp (corrected): EUR {inv_new:>+12,.0f}")
    print(f"  Delta (deeper-discount finding):           EUR {inv_new - inv_old:>+12,.0f}")
    # Promo uplift was +€156K in May per existing bridge
    print(f"\n  Current bridge: promo_uplift=EUR+156,140  promo_investment=EUR-47,244  ROI=EUR+108,897")
    new_roi = 156140 + inv_new
    print(f"  Corrected:      promo_uplift=EUR+156,140  promo_investment=EUR{inv_new:>+,.0f}  ROI=EUR{new_roi:>+,.0f}")
else:
    print("  (no promo SKUs found for May)")


# ────────────────────────────────────────────────────────────────────
# PLAN-QTY CHANNEL SPLIT — Option A (12w historical share) stability check
# ────────────────────────────────────────────────────────────────────
header("PLAN-QTY CHANNEL SPLIT — Option A stability check")

# For top 20 SKUs by May volume, compare trailing 12w wholesale share at two points
print("\nFor top 20 May SKUs: trailing 12w wholesale share at end-Apr vs end-May")
top_skus = may_skus["sku"].tolist()
shares = pd.read_sql(text("""
    SELECT p.sku,
           SUM(CASE WHEN et.transaction_date BETWEEN '2026-02-01' AND '2026-04-30'
                    AND cm.channel = 'wholesale' THEN et.quantity ELSE 0 END)::float AS apr_ws_qty,
           SUM(CASE WHEN et.transaction_date BETWEEN '2026-02-01' AND '2026-04-30'
                    THEN et.quantity ELSE 0 END)::float AS apr_total_qty,
           SUM(CASE WHEN et.transaction_date BETWEEN '2026-03-01' AND '2026-05-22'
                    AND cm.channel = 'wholesale' THEN et.quantity ELSE 0 END)::float AS may_ws_qty,
           SUM(CASE WHEN et.transaction_date BETWEEN '2026-03-01' AND '2026-05-22'
                    THEN et.quantity ELSE 0 END)::float AS may_total_qty
    FROM erp_transactions et
    JOIN dim_products p ON p.id = et.product_id
    JOIN lookup_channel_map cm ON cm.id = et.channel_map_id
    WHERE p.sku = ANY(:s)
    GROUP BY p.sku
"""), db.bind, params={"s": top_skus})
shares["share_apr"] = (shares["apr_ws_qty"] / shares["apr_total_qty"].replace(0, pd.NA) * 100).fillna(0)
shares["share_may"] = (shares["may_ws_qty"] / shares["may_total_qty"].replace(0, pd.NA) * 100).fillna(0)
shares["delta"] = (shares["share_may"] - shares["share_apr"]).abs()
print(f"  Median |Δshare| across top 20: {shares['delta'].median():.1f}pp")
print(f"  Max |Δshare|:                  {shares['delta'].max():.1f}pp")
print(f"  Mean wholesale share (top 20): apr={shares['share_apr'].mean():.1f}%   may={shares['share_may'].mean():.1f}%")
print()
print("  Sample (SKUs where wholesale share shifted most):")
for _, r in shares.nlargest(5, "delta").iterrows():
    print(f"    {r['sku']:<12s}  apr_ws={r['share_apr']:>5.1f}%   may_ws={r['share_may']:>5.1f}%   Δ={r['delta']:+.1f}pp")

# Option B: forecast.on_top_wholesale
print("\nOption B — forecasts table on_top_wholesale check")
fc = pd.read_sql(text("""
    SELECT p.sku,
           SUM(f.baseline)::float    AS baseline,
           SUM(f.on_top_wholesale)::float AS on_top_ws,
           SUM(f.on_top_retail)::float    AS on_top_retail,
           SUM(f.promo_uplift)::float     AS promo,
           SUM(f.total)::float            AS total
    FROM forecasts f JOIN dim_products p ON p.id = f.product_id
    WHERE f.run_id = (SELECT MAX(run_id) FROM forecasts)
    GROUP BY p.sku
"""), db.bind)
print(f"  forecasts rows: {len(fc):,} SKUs")
fc["implied_ws_share"] = (fc["on_top_ws"] / fc["total"].replace(0, pd.NA) * 100).fillna(0)
print(f"  Mean on_top_wholesale / total: {fc['implied_ws_share'].mean():.1f}%")
print(f"  (Compare with trailing 12w avg around {shares['share_may'].mean():.1f}% for top 20)")
print(f"  ⚠ NOTE: on_top_wholesale is INCREMENTAL on top of baseline. Doesn't tell us")
print(f"  the baseline's own channel split — so Option B alone is insufficient.")

db.close()
print()
print("DONE.")
