"""Load all data/VrstaRabatnePolitike*.xlsx files into promo_policies +
promo_policy_items. Idempotent: deletes and re-inserts each policy by name.

Each file has one sheet ('RabatnaPolitika') with one row per (policy × SKU).
The header row may be the policy summary (no SKU); we skip those.
"""
import sys
sys.stdout.reconfigure(encoding="utf-8")
from pathlib import Path
import pandas as pd
from backend.models.database import SessionLocal
from sqlalchemy import text


def load_one_file(path: Path, db) -> int:
    """Returns count of items loaded."""
    print(f"\n=== {path.name} ===")
    df = pd.read_excel(path, sheet_name="RabatnaPolitika")
    print(f"  rows in file: {len(df)}")

    # Drop the summary/header rows that have no Artikal (SKU)
    df = df[df["Artikal"].notna() & (df["Artikal"].astype(str).str.strip() != "")]
    print(f"  rows with SKU:  {len(df)}")
    if df.empty:
        return 0

    # Each file may contain >1 policy_name — group by policy
    n_items = 0
    for policy_name, grp in df.groupby("Rabatna politika"):
        policy_name = str(policy_name).strip()
        # Earliest/latest valid_from/to across items = policy window
        from_dates = pd.to_datetime(grp["Trajanje od"], errors="coerce")
        to_dates   = pd.to_datetime(grp["Trajanje do"], errors="coerce")
        valid_from = from_dates.min().date() if from_dates.notna().any() else None
        valid_to   = to_dates.max().date()   if to_dates.notna().any()   else None

        # Upsert header
        db.execute(text("DELETE FROM promo_policies WHERE policy_name = :n"),
                   {"n": policy_name})
        pid = db.execute(text("""
            INSERT INTO promo_policies (policy_name, valid_from, valid_to, source_file)
            VALUES (:n, :vf, :vt, :sf)
            RETURNING id
        """), {"n": policy_name, "vf": valid_from, "vt": valid_to,
               "sf": path.name}).scalar()

        # Resolve product_id via dim_products.sku — leave NULL if not in catalog
        skus = grp["Artikal"].astype(str).str.strip().tolist()
        prod_map = {
            row["sku"]: row["id"]
            for row in db.execute(text(
                "SELECT id, sku FROM dim_products WHERE sku = ANY(:s)"
            ), {"s": skus}).mappings()
        }

        items_inserted = 0
        unmatched_skus = 0
        for _, r in grp.iterrows():
            sku = str(r["Artikal"]).strip()
            product_id = prod_map.get(sku)
            if product_id is None:
                unmatched_skus += 1

            def _num(col):
                v = r.get(col)
                if v is None or pd.isna(v):
                    return None
                try:
                    return float(v)
                except (TypeError, ValueError):
                    return None

            def _date(col):
                v = r.get(col)
                if v is None or pd.isna(v):
                    return None
                try:
                    return pd.to_datetime(v).date()
                except Exception:
                    return None

            db.execute(text("""
                INSERT INTO promo_policy_items
                    (policy_id, product_id, sku, sku_name, description,
                     rabat_pct, promo_price, dev_price, min_qty,
                     vp_margin_pct, vp_margin_eur, mp_margin_pct, mp_margin_eur,
                     valid_from, valid_to, datalink)
                VALUES
                    (:pid, :pr, :sku, :nm, :ds,
                     :rb, :cp, :dp, :mq,
                     :vmp, :vme, :mmp, :mme,
                     :vf, :vt, :dl)
            """), {
                "pid": pid, "pr": product_id, "sku": sku,
                "nm": str(r.get("Naziv artikla") or "")[:255] or None,
                "ds": str(r.get("Opis") or "")[:255] or None,
                "rb":  _num("Rabat %"),
                "cp":  _num("Cijena"),
                "dp":  _num("Dev.cijena"),
                "mq":  _num("Količina >"),
                "vmp": _num("VP Marža %"),
                "vme": _num("VP Marža"),
                "mmp": _num("MP Marža %"),
                "mme": _num("MP Marža"),
                "vf":  _date("Trajanje od"),
                "vt":  _date("Trajanje do"),
                "dl":  int(r["DataLink"]) if pd.notna(r.get("DataLink")) else None,
            })
            items_inserted += 1
            n_items += 1
        print(f"  policy {policy_name!r}:  inserted {items_inserted} items, "
              f"unmatched SKUs={unmatched_skus}, window={valid_from}…{valid_to}")
    return n_items


def main():
    data_dir = Path("data")
    files = sorted(data_dir.glob("VrstaRabatnePolitike*.xlsx"))
    print(f"Files to load: {len(files)}")
    for p in files:
        print(f"  - {p.name}")

    db = SessionLocal()
    try:
        total = 0
        for p in files:
            total += load_one_file(p, db)
        db.commit()
        print(f"\n✓ loaded {total} total policy items across {len(files)} files")

        # Summary
        r = db.execute(text("""
            SELECT pp.policy_name, pp.valid_from, pp.valid_to,
                   COUNT(ppi.id) AS n_items
            FROM promo_policies pp
            LEFT JOIN promo_policy_items ppi ON ppi.policy_id = pp.id
            GROUP BY pp.id, pp.policy_name, pp.valid_from, pp.valid_to
            ORDER BY pp.policy_name
        """)).mappings().all()
        print("\nPolicies in DB:")
        for x in r:
            print(f"  {x['policy_name']:<24s}  {x['valid_from']}…{x['valid_to']}  ({x['n_items']} items)")
    except Exception:
        db.rollback()
        raise
    finally:
        db.close()


if __name__ == "__main__":
    main()
