"""Parent (proizvod) → child SKU grouping.

Two-stage resolution:

1. **Primary**: `data/Polleo Help svi artikli.xlsx` (sheet `Artikli`,
   columns `Naziv`, `SKU`). When a SKU is present in this file, its
   `Naziv` is used as the parent key — same `Naziv` across multiple
   SKUs = variants of one product. This is the source of truth.

2. **Fallback heuristic** for SKUs not in the Excel:
   - SKU with >=3 dash-segments (VENUM-03813-449-M, 1357719-001-XS):
     parent = SKU minus last segment.
   - Otherwise: parent = first N tokens of name (variant tokens stripped).

Single-SKU "parents" are kept but only useful as size-1 groups; the
planner filters those out.
"""

from __future__ import annotations

import re
import unicodedata
from functools import lru_cache
from pathlib import Path
from typing import Dict, List, Optional

import pandas as pd


# Variant tokens stripped from name tail when inferring parent stem.
# Conservative list — better to miss a grouping than to over-merge distinct
# products. Lowercase, no diacritics.
_VARIANT_TOKENS = {
    # apparel sizes
    "xs", "s", "sm", "m", "md", "l", "lg", "xl", "xxl", "xxxl",
    "small", "medium", "large",
    # colors EN
    "black", "white", "blue", "red", "green", "yellow", "orange",
    "purple", "pink", "grey", "gray", "brown", "silver", "gold",
    "navy", "olive", "beige", "crimson", "rosemary", "halo",
    # colors HR (no diacritics)
    "crna", "crno", "crni", "bijela", "bijelo", "bijeli",
    "plava", "plavo", "plavi", "zelena", "zeleno", "zeleni",
    "zuta", "zuto", "zuti", "crvena", "crveno", "crveni",
    "siva", "sivo", "sivi", "smeda", "smedi", "smedo",
    "ljubicasta", "narancasta", "zlatna", "srebrna", "zlatno", "srebrno",
    # flavors
    "chocolate", "vanilla", "strawberry", "banana", "coconut", "hazelnut",
    "coffee", "mango", "lemon", "watermelon", "peach", "raspberry",
    "blueberry", "cookie", "cookies", "cream", "milk", "plain", "neutral",
    "cinnamon", "caramel", "mint",
    "cokolada", "cokoladni", "vanilija", "vanilijev", "jagoda", "jagodin",
    "kokos", "kokosov", "ljesnjak", "kava", "lubenica", "breskva",
    "malina", "marakuja", "mandarina", "naranca", "limun", "neutralan",
}

_SIZE_UNIT_RE = re.compile(
    r"^\d+(?:[.,]\d+)?\s*(?:g|kg|ml|l|cm|mm|oz|caps|tabs|kom|stick|stickova|pack)?$"
)


def _normalize(s: str) -> str:
    """Lowercase, strip diacritics, collapse whitespace, drop punctuation."""
    s = unicodedata.normalize("NFKD", str(s)).encode("ascii", "ignore").decode("ascii")
    s = s.lower()
    s = re.sub(r"[^a-z0-9., ]+", " ", s)
    s = re.sub(r"\s+", " ", s).strip()
    return s


def _name_stem(name: str, n_tokens: int = 4) -> str:
    """Keep leading tokens until a variant token (color/flavor) is hit, or
    n_tokens reached. Size-unit tokens (25g, 2kg, 414g) are kept — they
    typically belong to the parent identity, not the variant."""
    norm = _normalize(name)
    tokens = norm.split()
    keep: List[str] = []
    for t in tokens:
        if len(keep) >= n_tokens:
            break
        if t in _VARIANT_TOKENS:
            break
        keep.append(t)
    return " ".join(keep) if keep else norm


def infer_parent_key(sku: str, name: str) -> str:
    """Heuristic fallback: stable parent key when no Excel mapping exists.
    Prefixed so SKU-derived and name-derived keys never collide with
    XLSX-derived keys."""
    sku_s = str(sku).strip()
    parts = sku_s.split("-")
    if len(parts) >= 3:
        return "SKU::" + "-".join(parts[:-1])
    return "NAME::" + _name_stem(name)


_XLSX_REL_PATH = "data/Polleo Help svi artikli.xlsx"


@lru_cache(maxsize=1)
def _load_xlsx_parent_map() -> Optional[Dict[str, str]]:
    """Load `Polleo Help svi artikli.xlsx` → {sku_str: naziv_str}.
    Returns None if the file is missing. Cached across calls.

    Path resolution: tries cwd first, then the parent_map.py module directory.
    Both styles are used by PromoTool depending on how Streamlit is launched.
    """
    candidates = [
        Path(_XLSX_REL_PATH),
        Path(__file__).resolve().parent / _XLSX_REL_PATH,
    ]
    xlsx_path = next((p for p in candidates if p.exists()), None)
    if xlsx_path is None:
        return None
    try:
        df = pd.read_excel(xlsx_path, sheet_name="Artikli")
    except Exception:
        return None
    df = df.dropna(subset=["SKU", "Naziv"])
    df["SKU"] = df["SKU"].astype(str).str.strip()
    df["Naziv"] = df["Naziv"].astype(str).str.strip()
    df = df[(df["SKU"] != "") & (df["Naziv"] != "")]
    # Same SKU appearing twice → first Naziv wins (rare in this file).
    return dict(zip(df["SKU"], df["Naziv"]))


def build_parent_map(
    plan_df: pd.DataFrame,
    *,
    sku_col: str = "sku",
    name_col: str = "name",
) -> Dict[str, Dict]:
    """Return {parent_key: {"display": str, "skus": [sku, ...]}}.

    For each SKU in plan_df:
      - if it's in the XLSX → group key = `XLSX::<Naziv>`, display = Naziv
      - else → fall back to heuristic (`SKU::...` or `NAME::...`)

    Display name preference:
      - XLSX-keyed: the Naziv string itself
      - SKU-keyed: SKU prefix + first-seen plan name
      - NAME-keyed: title-cased name stem
    """
    xlsx_map = _load_xlsx_parent_map() or {}

    parents: Dict[str, Dict] = {}
    first_seen: Dict[str, str] = {}      # parent_key -> first plan name
    naziv_display: Dict[str, str] = {}   # XLSX-key -> Naziv (display source)

    for _, row in plan_df.iterrows():
        sku = str(row.get(sku_col, "") or "").strip()
        name = str(row.get(name_col, "") or "").strip()
        if not sku:
            continue
        naziv = xlsx_map.get(sku)
        if naziv:
            key = "XLSX::" + naziv
            naziv_display.setdefault(key, naziv)
        else:
            key = infer_parent_key(sku, name)
        if key not in parents:
            parents[key] = {"display": "", "skus": []}
            first_seen[key] = name
        parents[key]["skus"].append(sku)

    # Build display name per parent
    for key, info in parents.items():
        if key.startswith("XLSX::"):
            info["display"] = naziv_display.get(key, key[len("XLSX::"):])
        elif key.startswith("SKU::"):
            sku_prefix = key[len("SKU::"):]
            base_name = first_seen.get(key, "")
            info["display"] = f"{sku_prefix} — {base_name}" if base_name else sku_prefix
        else:
            stem = key[len("NAME::"):]
            info["display"] = stem.title() if stem else first_seen.get(key, "(unnamed)")
    return parents


def _unused_synth_parent_rows(parent_key: str, skus, sales: pd.DataFrame,
                       erp: pd.DataFrame):
    """Build synthetic sales/ERP rows aggregating across child SKUs.

    Returned dataframes have ``sku == parent_key`` for every row. Sums hold
    qty/RUC columns; OR is applied to promo flags (any child on promo →
    parent flagged); max-abs holds discount %. ERP rows merge promo_types
    by union of strings within the same (year, week).

    Cheap to call — pure pandas groupby on the child subset. Cache with
    session_state keyed by ``(parent_key, tuple(sorted(skus)))`` if calling
    repeatedly with identical inputs.
    """
    skus_set = set(skus)
    cs = sales[sales["sku"].isin(skus_set)] if "sku" in sales.columns else \
        sales.iloc[0:0]
    if len(cs) == 0:
        empty_sales = sales.iloc[0:0].copy()
        empty_erp = erp.iloc[0:0].copy() if len(erp) else pd.DataFrame()
        return empty_sales, empty_erp

    sum_cols = [c for c in _SUM_COLS if c in cs.columns]
    flag_cols = [c for c in _MAX_COLS if c in cs.columns]
    disc_cols = [c for c in _ABS_MAX_COLS if c in cs.columns]

    g_sum = cs.groupby(["year", "week"], as_index=False)[sum_cols].sum() \
        if sum_cols else cs.groupby(["year", "week"], as_index=False).size()[["year", "week"]]
    parent_sales = g_sum

    # Promo flags: majority rule — flag parent week only when >=50% of the
    # children present that week are flagged. Avoids the failure mode where
    # a single child on promo (out of 20) torpedoes the parent baseline.
    if flag_cols:
        children_per_wk = cs.groupby(["year", "week"], as_index=False)["sku"].nunique()
        children_per_wk = children_per_wk.rename(columns={"sku": "_n"})
        flag_sums = cs.groupby(["year", "week"], as_index=False)[flag_cols].sum()
        flag_sums = flag_sums.merge(children_per_wk, on=["year", "week"])
        for c in flag_cols:
            flag_sums[c] = (flag_sums[c] >= 0.5 * flag_sums["_n"]).astype(int)
        parent_sales = parent_sales.merge(
            flag_sums[["year", "week"] + flag_cols],
            on=["year", "week"], how="left",
        )

    # Discount %: qty-weighted mean (not max-abs). 1 of 27 variants at -20%
    # should not look like a parent-wide -20% promo.
    if disc_cols and "qty_retail" in cs.columns:
        cs_w = cs.copy()
        cs_w["_w"] = cs_w["qty_retail"].fillna(0).clip(lower=0)
        for c in disc_cols:
            cs_w[f"_p_{c}"] = cs_w[c].fillna(0) * cs_w["_w"]
        agg_dict = {f"_p_{c}": "sum" for c in disc_cols}
        agg_dict["_w"] = "sum"
        g_disc = cs_w.groupby(["year", "week"], as_index=False).agg(agg_dict)
        for c in disc_cols:
            g_disc[c] = g_disc[f"_p_{c}"] / g_disc["_w"].replace(0, 1)
        parent_sales = parent_sales.merge(
            g_disc[["year", "week"] + disc_cols],
            on=["year", "week"], how="left",
        )
    elif disc_cols:
        # No weights available — fall back to plain mean.
        g_disc = cs.groupby(["year", "week"], as_index=False)[disc_cols].mean()
        parent_sales = parent_sales.merge(g_disc, on=["year", "week"], how="left")

    parent_sales["sku"] = parent_key
    # Reorder to match input column order (where possible)
    ordered = [c for c in sales.columns if c in parent_sales.columns]
    parent_sales = parent_sales[ordered]

    # ── ERP aggregation ──
    if len(erp) > 0 and "sku" in erp.columns:
        ce = erp[erp["sku"].isin(skus_set)]
        if len(ce):
            agg = {}
            for c in ce.columns:
                if c in ("sku", "year", "week"):
                    continue
                if c == "promo_types":
                    agg[c] = lambda s: ";".join(sorted({
                        t.strip()
                        for entries in s.astype(str)
                        for t in str(entries).split(";")
                        if t.strip()
                    }))
                else:
                    agg[c] = "first"
            parent_erp = ce.groupby(["year", "week"], as_index=False).agg(agg)
            parent_erp["sku"] = parent_key
            ordered_e = [c for c in erp.columns if c in parent_erp.columns]
            parent_erp = parent_erp[ordered_e]
        else:
            parent_erp = erp.iloc[0:0].copy()
    else:
        parent_erp = pd.DataFrame()

    return parent_sales, parent_erp
