"""Retail / web per-region forecaster (spec §4.1).

For a cleaned retail (or web) weekly series of one (SKU, region):
  1. clean (clean_retail: deflate promos, impute OOS) -> baseline,
  2. pick the best of a few light, proven models via a holdout backtest,
  3. forecast h weeks, then clamp with the sanity cap/floor.

Models are pure-numpy (recency-weighted SES, capped WMA, seasonal-index) — the
forms validated in the legacy engine, no heavy dependency. Holt/Holt-Winters is
deliberately excluded (project guardrail).
"""
from __future__ import annotations

import numpy as np

from forecast_v4.cleaning_retail import clean_retail

ALPHA_LO, ALPHA_HI = 0.15, 0.75


def _ses(y, h):
    y = np.asarray(y, float); n = len(y)
    if n < 3:
        return np.full(h, y.mean() if n else 0.0)
    a = 0.5; lvl = y[0]
    for t in range(1, n):
        lvl = a * y[t] + (1 - a) * lvl
    nz = y[-min(8, n):]; nz = nz[nz > 0]
    floor = np.median(nz) * 0.5 if len(nz) >= 3 else 0.0
    return np.full(h, max(floor, lvl))


def _wma(y, h):
    y = np.asarray(y, float); n = len(y)
    if n < 4:
        return np.full(h, y.mean() if n else 0.0)
    k = min(6, n); w = np.arange(1, k + 1, dtype=float)
    return np.full(h, max(0.0, np.average(y[-k:], weights=w)))


def _seasonal(y, weeks, start_cw, h):
    """26-wk base x weekly seasonal index x trend (damped)."""
    y = np.asarray(y, float); n = len(y)
    if n < 8 or weeks is None:
        return _wma(y, h)
    tail = y[-min(26, n):]; nzt = tail[tail > 0]
    base = float(nzt.mean()) if len(nzt) else float(tail.mean())
    if base <= 0:
        return np.zeros(h)
    overall = float(y[y > 0].mean()) if (y > 0).any() else base
    sums, cnts = {}, {}
    for i in range(n):
        if y[i] > 0:
            wk = int(weeks[i])
            sums[wk] = sums.get(wk, 0) + y[i]; cnts[wk] = cnts.get(wk, 0) + 1
    idx = {wk: 1.0 + (sums[wk] / cnts[wk] / overall - 1.0) * min(cnts[wk] / 3.0, 1.0)
           for wk in sums}
    tail8 = y[-min(8, n):]; nz8 = tail8[tail8 > 0]
    trend = float(np.clip((nz8.mean() if len(nz8) else base) / base, 0.75, 1.4))
    return np.array([max(0.0, base * idx.get(((start_cw + i - 1) % 52) + 1, 1.0) * trend)
                     for i in range(h)])


def _apply_bounds(fc, y):
    """Sanity cap (2x recent 13-wk avg) and floor (50% of 8-wk median)."""
    y = np.asarray(y, float)
    rec = y[-min(13, len(y)):]; rec = rec[rec > 0]
    if len(rec):
        fc = np.minimum(fc, 2.0 * rec.mean())
    nz8 = y[-min(8, len(y)):]; nz8 = nz8[nz8 > 0]
    if len(nz8) >= 3:
        fc = np.maximum(fc, 0.5 * np.median(nz8))
    return np.maximum(fc, 0.0)


def _backtest_pick(y, weeks, start_cw, test=4):
    """Hold out the last `test` weeks; pick the model with lowest MAE."""
    n = len(y)
    cands = {"ses": _ses, "wma": _wma}
    if n >= test + 12:
        ytr, yte = y[:-test], y[-test:]
        wtr = weeks[:-test] if weeks is not None else None
        best, berr = "wma", np.inf
        for name, fn in cands.items():
            pred = fn(ytr, test)
            err = np.mean(np.abs(pred - yte))
            if err < berr:
                berr, best = err, name
        if weeks is not None:
            pred = _seasonal(ytr, wtr, int(weeks[-test]) if len(weeks) else 1, test)
            if np.mean(np.abs(pred - yte)) < berr:
                best = "seasonal"
        return best
    return "wma"


def forecast_retail(qty, weeks=None, start_cw=1, promo_mask=None, price=None,
                    h=13, test=4):
    """Returns (forecast[h], model_name, uplift). `weeks` = ISO week number per
    point (enables the seasonal model); `start_cw` = first forecast ISO week."""
    baseline, uplift = clean_retail(qty, promo_mask=promo_mask, price=price)
    model = _backtest_pick(baseline, weeks, start_cw, test=test)
    if model == "seasonal":
        fc = _seasonal(baseline, weeks, start_cw, h)
    elif model == "ses":
        fc = _ses(baseline, h)
    else:
        fc = _wma(baseline, h)
    return _apply_bounds(fc, baseline), model, float(uplift)
