"""
Zarattini & Antonacci (2024/2025) **Timing Industry** long-only trend-following engine.

Implements the authors' published Colab reference logic (Concretum Research blog, 2025).
"""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np
import pandas as pd

from RenTech.strategy_stack.french_industry_loader import (
    align_industry_factors,
    load_french_factors_daily,
    load_industry_daily_returns,
)

N_INDUSTRIES = 48
UP_DAY = 20
DOWN_DAY = 40
ADR_VOL_ADJ = 1.4
KELT_MULT = 2.0 * ADR_VOL_ADJ
TARGET_VOL = 0.015
MAX_LEVERAGE = 2.0
MAX_NOT_TRADE = 1.0


@dataclass
class IndustryTrendTimingResult:
    daily: pd.DataFrame
    yearly: pd.DataFrame
    meta: dict
    weights: pd.DataFrame
    active_mask: pd.DataFrame


def _rolling_mean_abs_change(price: pd.DataFrame, window: int) -> pd.DataFrame:
    chg = price.diff().abs()
    return chg.rolling(window=window, min_periods=window - 1).mean()


def _compute_indicators(ind_rets: pd.DataFrame) -> dict[str, pd.DataFrame]:
    filled = ind_rets.fillna(0.0)
    price = (1.0 + filled).cumprod()

    vol = filled.rolling(window=UP_DAY).std(ddof=0)
    ema_up = price.ewm(span=UP_DAY, adjust=False).mean()
    ema_down = price.ewm(span=DOWN_DAY, adjust=False).mean()

    donc_up = price.rolling(window=UP_DAY).max()
    donc_down = price.rolling(window=DOWN_DAY).min()

    avg_chg = _rolling_mean_abs_change(price, UP_DAY)
    avg_chg_dn = _rolling_mean_abs_change(price, DOWN_DAY)
    kelt_up = ema_up + KELT_MULT * avg_chg
    kelt_down = ema_down - KELT_MULT * avg_chg_dn

    long_band = pd.DataFrame(
        np.minimum(donc_up.to_numpy(), kelt_up.to_numpy()),
        index=price.index,
        columns=price.columns,
    )
    short_band = pd.DataFrame(
        np.maximum(donc_down.to_numpy(), kelt_down.to_numpy()),
        index=price.index,
        columns=price.columns,
    )

    long_band_lag = long_band.shift(1)
    short_band_lag = short_band.shift(1)
    long_signal = (price >= long_band_lag) & (long_band_lag > short_band_lag)

    return {
        "price": price,
        "vol": vol,
        "long_band": long_band,
        "short_band": short_band,
        "long_signal": long_signal,
        "rets_filled": filled,
    }


def _build_meta(
    *,
    usable: pd.DataFrame,
    capital: float,
    universe_label: str,
    strategy_label: str,
    benchmark_label: str,
    n_universe: int,
    target_vol: float,
    max_leverage: float,
    max_not_trade: float,
    extra: dict | None = None,
) -> dict:
    r = usable["daily_ret"].astype(float)
    mkt_u = usable["mkt_ret"].astype(float)
    rf_u = usable["rf_ret"].astype(float)
    excess = r - rf_u
    mkt_excess = mkt_u - rf_u

    n = len(usable)
    years = n / 252.0
    end_eq = float(usable["equity_usd"].iloc[-1])
    total_ret = end_eq / float(capital) - 1.0
    cagr = (end_eq / float(capital)) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    vol_ann = float(r.std(ddof=1) * np.sqrt(252.0)) if n > 1 else float("nan")
    sharpe = (
        float(excess.mean() / excess.std(ddof=1) * np.sqrt(252.0))
        if excess.std(ddof=1) > 1e-12
        else float("nan")
    )
    eq = usable["equity_usd"]
    max_dd = float((eq / eq.cummax() - 1.0).min())

    mkt_end = float((1.0 + mkt_u).cumprod().iloc[-1])
    mkt_cagr = mkt_end ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    mkt_vol = float(mkt_u.std(ddof=1) * np.sqrt(252.0))
    mkt_sharpe = (
        float(mkt_excess.mean() / mkt_excess.std(ddof=1) * np.sqrt(252.0))
        if mkt_excess.std(ddof=1) > 1e-12
        else float("nan")
    )
    mkt_dd = float(((1.0 + mkt_u).cumprod() / (1.0 + mkt_u).cumprod().cummax() - 1.0).min())

    xy = pd.DataFrame({"y": excess, "x": mkt_excess}).dropna()
    beta = float("nan")
    alpha_ann = float("nan")
    if len(xy) > 10 and xy["x"].std(ddof=1) > 1e-12:
        beta = float(np.cov(xy["y"], xy["x"], ddof=1)[0, 1] / np.var(xy["x"], ddof=1))
        alpha_daily = float(xy["y"].mean() - beta * xy["x"].mean())
        alpha_ann = alpha_daily * 252.0

    meta = {
        "strategy": strategy_label,
        "paper": "Zarattini & Antonacci — A Century of Profitable Industry Trends (2025 Dow Award)",
        "implementation": "Concretum Colab reference (long_band>short_band filter, target_vol/available sizing)",
        "universe": universe_label,
        "start": str(usable["date"].iloc[0]),
        "end": str(usable["date"].iloc[-1]),
        "n_sessions": int(n),
        "capital": float(capital),
        "ending_equity_usd": round(end_eq, 2),
        "total_return_pct": round(total_ret * 100.0, 4),
        "cagr_pct": round(cagr * 100.0, 4),
        "vol_ann_pct": round(vol_ann * 100.0, 4),
        "sharpe_excess_rf": round(sharpe, 4),
        "max_drawdown_pct": round(max_dd * 100.0, 4),
        "alpha_ann_pct": round(alpha_ann * 100.0, 4) if np.isfinite(alpha_ann) else None,
        "beta_mkt": round(beta, 4) if np.isfinite(beta) else None,
        "avg_gross_exposure": round(float(usable["gross_exposure"].mean()), 4),
        "avg_n_active": round(float(usable["n_active"].mean()), 2),
        "benchmark": {
            "label": benchmark_label,
            "cagr_pct": round(mkt_cagr * 100.0, 4),
            "vol_ann_pct": round(mkt_vol * 100.0, 4),
            "sharpe_excess_rf": round(mkt_sharpe, 4),
            "max_drawdown_pct": round(mkt_dd * 100.0, 4),
        },
        "parameters": {
            "up_day": UP_DAY,
            "down_day": DOWN_DAY,
            "kelt_mult": KELT_MULT,
            "target_vol": target_vol,
            "max_leverage": max_leverage,
            "max_not_trade": max_not_trade,
            "vol_lookback": UP_DAY,
            "vol_ddof": 0,
            "n_universe": n_universe,
        },
    }
    if extra:
        meta.update(extra)
    return meta


def run_trend_timing_panel(
    asset_rets: pd.DataFrame,
    *,
    rf: pd.Series,
    mkt_rets: pd.Series,
    start: str,
    end: str,
    capital: float = 100_000.0,
    n_universe: int,
    universe_label: str,
    strategy_label: str = "industry_trend_timing",
    benchmark_label: str = "Market",
    target_vol: float = TARGET_VOL,
    max_leverage: float = MAX_LEVERAGE,
    max_not_trade: float = MAX_NOT_TRADE,
    extra_meta: dict | None = None,
) -> IndustryTrendTimingResult:
    """Run timing model on any daily-return panel (decimal returns)."""
    t0, t1 = pd.Timestamp(start), pd.Timestamp(end)
    idx = asset_rets.index.intersection(rf.index).intersection(mkt_rets.index)
    asset_rets = asset_rets.loc[idx].copy()
    rf = rf.reindex(idx)
    mkt_rets = mkt_rets.reindex(idx)

    asset_rets = asset_rets.loc[(asset_rets.index >= t0) & (asset_rets.index <= t1)]
    rf = rf.reindex(asset_rets.index)
    mkt_rets = mkt_rets.reindex(asset_rets.index)

    ind = _compute_indicators(asset_rets)
    price = ind["price"]
    vol = ind["vol"]
    long_signal = ind["long_signal"]
    short_band = ind["short_band"]
    rets_raw = asset_rets
    rets_fill = ind["rets_filled"]

    tickers = list(asset_rets.columns)
    n_assets = len(tickers)
    dates = asset_rets.index
    T = len(dates)

    px = price.to_numpy(dtype=np.float64)
    sig = long_signal.to_numpy(dtype=bool)
    sb = short_band.to_numpy(dtype=np.float64)
    vol_arr = vol.to_numpy(dtype=np.float64)
    ret_arr = rets_fill.to_numpy(dtype=np.float64)
    ret_nan = rets_raw.to_numpy(dtype=np.float64)
    rf_arr = rf.to_numpy(dtype=np.float64)
    mkt_arr = mkt_rets.to_numpy(dtype=np.float64)
    available = np.sum(np.isfinite(ret_nan), axis=1).astype(np.float64)
    available = np.where(available > 0, available, float(n_universe))

    exposure = np.zeros((T, n_assets), dtype=np.int8)
    ind_weight = np.zeros((T, n_assets), dtype=np.float64)
    trail = np.full((T, n_assets), np.nan, dtype=np.float64)

    for t in range(1, T):
        valid = np.isfinite(ret_nan[t]) & np.isfinite(sb[t])
        prev_exp = exposure[t - 1]
        cur_exp = np.zeros(n_assets, dtype=np.int8)
        cur_trail = trail[t].copy()

        new_long = valid & (prev_exp <= 0) & sig[t]
        confirm_long = valid & (prev_exp == 1) & (
            px[t] > np.maximum(trail[t - 1], sb[t])
        )
        exit_long = valid & (prev_exp == 1) & (
            px[t] <= np.maximum(trail[t - 1], sb[t])
        )

        cur_exp[new_long] = 1
        cur_trail[new_long] = sb[t][new_long]

        cur_exp[confirm_long] = 1
        cur_trail[confirm_long] = np.maximum(
            trail[t - 1, confirm_long], sb[t, confirm_long]
        )

        cur_exp[exit_long] = 0
        ind_weight[t, exit_long] = 0.0

        active = cur_exp == 1
        lev = np.divide(
            target_vol,
            vol_arr[t],
            out=np.zeros_like(vol_arr[t]),
            where=np.isfinite(vol_arr[t]) & (vol_arr[t] > 0),
        )
        ind_weight[t, active] = lev[active]
        exposure[t] = cur_exp
        trail[t] = cur_trail

    port_w = ind_weight / available[:, None]
    port_w = np.minimum(port_w, max_not_trade)
    gross = port_w.sum(axis=1)
    over = gross > max_leverage
    if np.any(over):
        port_w[over] *= (max_leverage / gross[over])[:, None]
    gross = port_w.sum(axis=1)

    port_ret = np.zeros(T, dtype=np.float64)
    for t in range(1, T):
        w_prev = port_w[t - 1]
        risky = np.nansum(w_prev * np.nan_to_num(ret_arr[t], nan=0.0))
        port_ret[t] = risky + (1.0 - gross[t - 1]) * rf_arr[t]

    daily = pd.DataFrame(
        {
            "date": dates.strftime("%Y-%m-%d"),
            "daily_ret": port_ret,
            "daily_pnl_usd": port_ret * float(capital),
            "equity_unit": (1.0 + pd.Series(port_ret, index=dates)).cumprod().to_numpy(),
            "equity_usd": float(capital) * (1.0 + pd.Series(port_ret, index=dates)).cumprod().to_numpy(),
            "gross_exposure": gross,
            "n_active": exposure.sum(axis=1),
            "rf_ret": rf_arr,
            "mkt_ret": mkt_arr,
        }
    )

    warm = max(UP_DAY, DOWN_DAY) + 2
    usable = daily.iloc[warm:].copy()

    yearly_rows = []
    for yr, g in usable.groupby(pd.to_datetime(usable["date"]).dt.year):
        yr_ret = float((1.0 + g["daily_ret"]).prod() - 1.0)
        mkt_yr = float((1.0 + g["mkt_ret"]).prod() - 1.0)
        yearly_rows.append(
            {
                "year": int(yr),
                "return_pct": round(yr_ret * 100.0, 4),
                "mkt_return_pct": round(mkt_yr * 100.0, 4),
                "avg_gross_exposure": round(float(g["gross_exposure"].mean()), 4),
                "avg_n_active": round(float(g["n_active"].mean()), 2),
            }
        )
    yearly = pd.DataFrame(yearly_rows)

    meta = _build_meta(
        usable=usable,
        capital=capital,
        universe_label=universe_label,
        strategy_label=strategy_label,
        benchmark_label=benchmark_label,
        n_universe=n_universe,
        target_vol=target_vol,
        max_leverage=max_leverage,
        max_not_trade=max_not_trade,
        extra=extra_meta,
    )

    w_df = pd.DataFrame(port_w, index=dates, columns=tickers)
    active_df = exposure.astype(bool)

    return IndustryTrendTimingResult(
        daily=daily,
        yearly=yearly,
        meta=meta,
        weights=w_df,
        active_mask=active_df,
    )


def run_industry_trend_timing(
    *,
    start: str = "1926-07-01",
    end: str = "2024-03-29",
    capital: float = 100_000.0,
    n_industries: int = N_INDUSTRIES,
    industry_weighting: str = "value",
    target_vol: float = TARGET_VOL,
    max_leverage: float = MAX_LEVERAGE,
    max_not_trade: float = MAX_NOT_TRADE,
) -> IndustryTrendTimingResult:
    ind_rets = load_industry_daily_returns(weighting=industry_weighting)
    factors = load_french_factors_daily()
    ind_rets, mkt_rets, rf = align_industry_factors(ind_rets, factors)

    return run_trend_timing_panel(
        ind_rets,
        rf=rf,
        mkt_rets=mkt_rets,
        start=start,
        end=end,
        capital=capital,
        n_universe=n_industries,
        universe_label=f"Ken French 48 industries ({industry_weighting}-weighted daily returns)",
        strategy_label="industry_trend_timing",
        benchmark_label="French VW market (Mkt-RF + RF)",
        target_vol=target_vol,
        max_leverage=max_leverage,
        max_not_trade=max_not_trade,
    )
