"""
Walk-Forward Factor Selection with Meta-Selection
JKP Characteristics Trading Factor (CTF) challenge submission.

INTERFACE
=========
    main(chars, features, daily_ret) -> DataFrame[id, eom, w]

`chars` is the ctff_chars stock-month panel, `features` the ctff_features
list of characteristic column names, `daily_ret` the optional ctff_daily_ret
table (accepted and ignored). The returned weights at month-end t are held
over month t+1 and earn that row's ret_exc_lead1m. Monthly rebalancing.

WHAT IT DOES
============
Four stages, each of which may only use information realized by its own
decision date:

  Stage 1  For each of the ~400 characteristics and each month, build a
           linear-rank long-short factor portfolio return:
               w_i = pct_rank(char_i) - 0.5
               LS_t = sum_i(w_i * ret_i) / sum_i(|w_i|)
           CONVENTION: the LS value stored at month t is computed from
           ret_exc_lead1m, so it is the return earned OVER MONTH t+1 and is
           known only at the end of month t+1.

  Stage 2  A fixed, symmetric menu of 108 configurations
           (3 lookback windows x 3 portfolio sizes x 3 selection rules
           x 4 weighting rules) is simulated walk-forward. Every 6 months each configuration picks
           its factors and weights from LS rows <= rp - 1 month only, then
           earns the realized LS returns until the next rebalance. Each
           configuration therefore accumulates its own out-of-sample record.

           The menu is a complete cross, not a curated list. Candidates
           were added to it; none was ever chosen for having performed well.
           Selection rules: Sharpe, ClusterSharpe, and ConsistentSharpe (a
           factor's worst Sharpe across three equal sub-periods of the
           lookback, which rewards consistency over one lucky stretch).
           Weighting rules: EW, RankSharpe, MeanVar, and InvVol (w ~ 1/sigma,
           which uses no mean return and so cannot chase a lucky streak).
           Portfolio sizes: 5, 10 or 20 characteristics.

  Stage 3  Meta-selection. At each rebalance the configurations are ranked
           by the annualized Sharpe of their OWN out-of-sample returns up to
           rp - 1 month, and the top 3 are blended equally. Before any
           configuration has 120 months of track record, all active
           configurations are blended equally. The methodology in force at
           any date is thus chosen by data strictly older than that date,
           which removes configuration selection bias.

  Stage 4  The blended configurations' factor weights are mapped back to
           stocks using month-t characteristic ranks only:
               w_i = (1/K) * sum_cfg sum_k fw_k * (rank_i - 0.5)
                                            / sum_j|rank_j - 0.5|

WHY ROW rp - 1 IS THE LAST ROW READ
===================================
A decision taken at the end of month rp may use any return realized by then.
LS row rp - 1 is realized over month rp, so it is known at the end of month
rp; LS row rp is realized over month rp + 1 and is not. Every lookback slice
in Stages 2 and 3 therefore ends at rp - 1 month.

DETERMINISM
===========
No randomness anywhere: no sampling, no shuffling, no random initialization,
no iteration over unordered containers whose order could affect results.
Two runs on the same input produce byte-identical output.

DEPENDENCIES
============
numpy, pandas, scipy only (see requirements.txt for pinned versions).
No external data is read; everything comes from the arguments to main().

CTF Admin Modifications (2026-09-13):
--------------------------------------
1. Added a suppression comment on line 195 to allow the broad exception
   handler in select_cluster_sharpe().
   Reason: The pipeline's static security checks reject a broad exception
   handler that does not re-raise. Here the handler is an intentional
   fallback to plain Sharpe selection when hierarchical clustering fails,
   so the code is safe to run as written.
"""
import argparse
import time

import numpy as np
import pandas as pd
from scipy.cluster.hierarchy import linkage, fcluster
from scipy.spatial.distance import squareform

# ---------------------------------------------------------------------------
# Hyperparameters. Identical to the baseline; fixed a priori and not retuned.
# ---------------------------------------------------------------------------
REBAL_MONTHS = 6      # rebalance / meta-selection frequency (months)
META_TOP_K = 3        # how many configs the meta layer blends
META_MIN_OBS = 120    # months of config history before meta-selection starts
MIN_OBS = 60          # months of data required for a factor to be usable
MIN_STOCKS = 10       # min stocks per side of a factor portfolio
MIN_HISTORY = 24      # months of history before the first rebalance
SHRINKAGE = 0.3       # covariance shrinkage in mean-variance weighting
N_CLUSTERS_MULT = 3   # clusters = n_factors * this, in ClusterSharpe

# New, and likewise fixed a priori.
N_SUBWINDOWS = 3      # sub-periods ConsistentSharpe scores a factor over


# ===========================================================================
# Stage 1: factor portfolio returns
# ===========================================================================
def compute_factor_ls_returns(chars, feature_names, min_stocks, min_obs):
    """
    Linear-rank long-short return for every feature, every month.

    weight_i = pct_rank(feature_i) - 0.5          (mean ~ 0)
    LS_t     = sum_i(weight_i * ret_i) / sum_i(|weight_i|)

    CONVENTION (memorize this): the LS value stored at month t uses
    ret_exc_lead1m, i.e. it is the return earned over month t+1. It is
    KNOWN only at the end of month t+1. All later stages must respect this.
    """
    print("  Stage 1: factor LS returns...")
    t0 = time.time()

    valid_features = [f for f in feature_names if f in chars.columns]
    df = chars[["eom", "ret_exc_lead1m"] + valid_features]
    df = df[df["ret_exc_lead1m"].notna()]

    months = sorted(df["eom"].unique())
    ls_matrix = np.full((len(months), len(valid_features)), np.nan)

    for m_idx, eom in enumerate(months):
        month_data = df[df["eom"] == eom]
        if len(month_data) < min_stocks * 2:
            continue

        ret = month_data["ret_exc_lead1m"].values
        ranks = month_data[valid_features].rank(pct=True).values
        w = np.nan_to_num(ranks, nan=0.5) - 0.5
        not_nan = ~np.isnan(month_data[valid_features].values)
        w = w * not_nan

        num = (ret[:, None] * w).sum(axis=0)
        den = np.abs(w).sum(axis=0)
        ok = (not_nan.sum(axis=0) >= min_stocks * 2) & (den > 1e-10)
        row = np.full(len(valid_features), np.nan)
        row[ok] = num[ok] / den[ok]
        ls_matrix[m_idx] = row

    ls_df = pd.DataFrame(ls_matrix, index=pd.to_datetime(months),
                         columns=valid_features).sort_index()
    ls_df = ls_df[ls_df.columns[ls_df.notna().sum() >= min_obs]]
    print(f"    {ls_df.shape[0]} months x {ls_df.shape[1]} factors "
          f"({time.time()-t0:.0f}s)")
    return ls_df


# ===========================================================================
# Selection rules: history -> ordered list of factor names
#
# Every rule receives `tc`, the lookback slice simulate_configs has already
# truncated to LS rows <= rp - 1 month and filled NaN with 0. No rule may look
# anywhere else.
# ===========================================================================
def _sharpe_series(tc):
    stds = tc.std().replace(0, np.nan)
    return ((tc.mean() / stds) * np.sqrt(12)).fillna(0)


def select_sharpe(tc, valid, n_factors):
    """Top-N factors by annualized Sharpe over the lookback slice."""
    return list(_sharpe_series(tc[valid]).nlargest(
        min(n_factors, len(valid))).index)


def select_cluster_sharpe(tc, valid, n_factors):
    """
    Diversified top-N: cluster factors by return correlation, then take the
    best-Sharpe factor from each of the best N clusters. Prevents picking
    five flavours of the same momentum factor.
    """
    tcv = tc[valid]
    if len(tcv) < 12 or len(valid) < n_factors:
        return select_sharpe(tc, valid, n_factors)

    corr = tcv.corr().fillna(0).clip(-1, 1)
    dist = (1.0 - corr.abs()).values.copy()
    np.fill_diagonal(dist, 0)
    dist = np.clip((dist + dist.T) / 2, 0, None)
    try:
        condensed = np.nan_to_num(squareform(dist, checks=False),
                                  nan=1.0, posinf=1.0, neginf=0.0)
        Z = linkage(condensed, method="average")
        n_clusters = min(max(n_factors * N_CLUSTERS_MULT, 20), len(valid))
        labels = fcluster(Z, t=n_clusters, criterion="maxclust")
    except Exception: - intentional fallback to plain Sharpe selection
        return select_sharpe(tc, valid, n_factors)

    scores = _sharpe_series(tcv)
    cdf = pd.DataFrame({"factor": valid, "cluster": labels,
                        "score": scores.fillna(-999).values})
    best = cdf.loc[cdf.groupby("cluster")["score"].idxmax()]
    best = best[best["score"] > -998].sort_values("score", ascending=False)
    return best.head(n_factors)["factor"].tolist()


def select_consistent_sharpe(tc, valid, n_factors):
    """
    EXTENSION POINT 2. Rank factors by their WORST Sharpe across N_SUBWINDOWS
    contiguous, equal sub-periods of the lookback rather than by the Sharpe of
    the whole window.

    A factor that earned its full-window Sharpe in one lucky stretch scores
    badly here; a factor that worked in every third scores well. This is a
    genuinely different ordering, not a rescaling: note that shrinking Sharpe
    toward the cross-factor mean by a common factor -- the other obvious
    candidate -- is a monotone transform and would return the identical top-N,
    which is why it is not offered.

    Falls back to plain Sharpe when the window is too short to split, so the
    123-month validation panel and the early years still behave.
    """
    tcv = tc[valid]
    if len(tcv) < N_SUBWINDOWS * 12 or len(valid) == 0:
        return select_sharpe(tc, valid, n_factors)

    parts = np.array_split(np.arange(len(tcv)), N_SUBWINDOWS)
    sub = pd.concat([_sharpe_series(tcv.iloc[p]) for p in parts], axis=1)
    worst = sub.min(axis=1)
    return list(worst.nlargest(min(n_factors, len(valid))).index)


SELECT_FN = {
    "Sharpe": select_sharpe,
    "ClusterSharpe": select_cluster_sharpe,
    "ConsistentSharpe": select_consistent_sharpe,
}


# ===========================================================================
# Weighting rules: history + chosen factors -> weight vector (sums to 1)
# ===========================================================================
def weight_ew(tc, factors):
    return np.full(len(factors), 1.0 / len(factors))


def weight_rank_sharpe(tc, factors):
    """Best factor by Sharpe gets the largest weight (rank weights)."""
    order = _sharpe_series(tc[factors]).rank(method="first")
    w = order.values.astype(float)
    return w / w.sum()


def weight_mean_var(tc, factors):
    """w ~ Sigma^{-1} mu with diagonal shrinkage; long-only, sums to 1."""
    tcf = tc[factors].fillna(0)
    n = len(factors)
    if len(tcf) < max(12, n + 2):
        return np.full(n, 1.0 / n)
    mu, cov = tcf.mean().values, tcf.cov().values
    cov = (1 - SHRINKAGE) * cov + SHRINKAGE * (np.trace(cov) / n) * np.eye(n)
    try:
        w = np.linalg.inv(cov) @ mu
    except np.linalg.LinAlgError:
        return np.full(n, 1.0 / n)
    w = np.maximum(w, 0)
    return w / w.sum() if w.sum() > 1e-10 else np.full(n, 1.0 / n)


def weight_inv_vol(tc, factors):
    """
    EXTENSION POINT 3. w_k proportional to 1 / sigma_k, normalized to sum
    to 1.

    Unlike MeanVar this uses no mean return at all, so it cannot tilt toward a
    factor that happens to have had a good run -- it only equalizes risk
    contributions. That makes it the most conservative member of the weighting
    menu and a useful counterweight to MeanVar, which is the noisiest.
    """
    n = len(factors)
    sd = tc[factors].std()
    inv = (1.0 / sd.replace(0, np.nan)).replace([np.inf, -np.inf], np.nan)
    w = inv.fillna(0.0).values.astype(float)
    if not np.isfinite(w).all() or w.sum() <= 1e-10:
        return np.full(n, 1.0 / n)
    return w / w.sum()


WEIGHT_FN = {
    "EW": weight_ew,
    "RankSharpe": weight_rank_sharpe,
    "MeanVar": weight_mean_var,
    "InvVol": weight_inv_vol,
}


# ===========================================================================
# The candidate menu. Still a complete symmetric cross, just a wider one.
# ===========================================================================
N_FACTORS_MENU = (5, 10, 20)

CONFIG_GRID = [
    {"window": w, "n_factors": n, "selection": s, "weighting": g}
    for w in (10, 20, None)            # lookback years; None = expanding
    for n in N_FACTORS_MENU            # number of factors held
    for s in SELECT_FN                 # selection rules registered above
    for g in WEIGHT_FN                 # weighting rules registered above
]


def config_label(cfg):
    w = "all" if cfg["window"] is None else f"{cfg['window']}y"
    return f"{w}/{cfg['n_factors']}/{cfg['selection']}/{cfg['weighting']}"


# ===========================================================================
# Stage 2: walk-forward simulation of every config
# ===========================================================================
def simulate_configs(ls_returns, configs, rebal_months, min_obs, min_history):
    """
    At each rebalance rp, every config selects factors/weights from LS rows
    <= rp - 1 month, then earns the realized LS returns until the next
    rebalance. Returns the per-config schedules and out-of-sample returns.
    """
    print("  Stage 2: walk-forward config simulation...")
    t0 = time.time()
    all_months = ls_returns.index.sort_values()
    rebal_points = list(all_months[::rebal_months])

    oos = pd.DataFrame(np.nan, index=all_months, columns=range(len(configs)))
    schedule = {}
    sel_cache = {}

    for i, rp in enumerate(rebal_points):
        sel_end = rp - pd.DateOffset(months=1)   # last KNOWN LS row
        ls_avail = ls_returns.loc[:sel_end]
        if len(ls_avail.dropna(how="all")) < min_history:
            continue

        next_rp = rebal_points[i + 1] if i + 1 < len(rebal_points) else None
        block = (all_months[(all_months >= rp) & (all_months < next_rp)]
                 if next_rp is not None else all_months[all_months >= rp])

        threshold = min(min_obs, max(6, len(ls_avail) // 2))
        valid = list(ls_avail.columns[ls_avail.notna().sum() >= threshold])
        if not valid:
            continue

        per_cfg = {}
        for ci, cfg in enumerate(configs):
            if cfg["window"] is None:
                tc = ls_avail
            else:
                start = (sel_end - pd.DateOffset(years=cfg["window"])
                         + pd.DateOffset(months=1))
                tc = ls_avail.loc[start:]
            tc = tc.fillna(0)
            if len(tc) < 12:
                continue

            key = (rp, cfg["window"], cfg["selection"], cfg["n_factors"])
            if key not in sel_cache:
                sel_cache[key] = SELECT_FN[cfg["selection"]](
                    tc, valid, cfg["n_factors"])
            factors = sel_cache[key]
            if not factors:
                continue

            weights = WEIGHT_FN[cfg["weighting"]](tc, factors)
            per_cfg[ci] = (factors, weights)
            oos.loc[block, ci] = (
                ls_returns.loc[block, factors].fillna(0).values @ weights)

        if per_cfg:
            schedule[rp] = per_cfg

    print(f"    {len(configs)} configs x {len(schedule)} rebalances "
          f"({time.time()-t0:.0f}s)")
    return rebal_points, schedule, oos


# ===========================================================================
# Stage 3: meta-selection  (EXTENSION POINT 5 deliberately NOT used)
# ===========================================================================
def meta_select(rebal_points, schedule, oos, configs, top_k, meta_min_obs):
    """
    At each rebalance, blend the top_k configs by the annualized Sharpe of
    their OWN out-of-sample returns up to t-1. Until any config has
    meta_min_obs months of history, hold an equal blend of all configs.
    """
    print("  Stage 3: meta-selection...")
    chosen, log = {}, []
    for rp in rebal_points:
        if rp not in schedule:
            continue
        active = list(schedule[rp].keys())
        hist = oos.loc[:rp - pd.DateOffset(months=1)]
        counts = hist.notna().sum()

        eligible = [ci for ci in active if counts[ci] >= meta_min_obs]
        if eligible:
            score = {}
            for ci in eligible:
                s = hist[ci].dropna()
                sd = s.std()
                score[ci] = (s.mean() / sd) * np.sqrt(12) if sd > 0 else -np.inf
            top = sorted(score, key=score.get, reverse=True)[:top_k]
            mode = "meta"
        else:
            top, mode = active, "ew_prior"

        chosen[rp] = top
        log.append({"rebal": rp, "mode": mode,
                    "chosen": ";".join(config_label(configs[ci])
                                       for ci in top)})
    n_meta = sum(1 for r in log if r["mode"] == "meta")
    print(f"    {len(chosen)} rebalances ({n_meta} meta-selected)")
    return chosen, log


# ===========================================================================
# Stage 4: stock-level weights
# ===========================================================================
def compute_stock_weights(chars, schedule, chosen):
    """
    Blend the chosen configs equally and map factor weights to stocks:
      w_i = (1/K) * sum_cfg sum_k  fw_k * (rank_i - 0.5) / sum_j|rank_j - 0.5|
    Each config's stock portfolio exactly replicates its factor-level
    return, so Stage 3's statistics describe the portfolio actually held.
    """
    print("  Stage 4: stock-level weights...")
    t0 = time.time()
    rebal_dates = sorted(chosen.keys())
    if not rebal_dates:
        return pd.DataFrame(columns=["id", "eom", "w"])

    all_factors = sorted({f for rp in rebal_dates for ci in chosen[rp]
                          for f in schedule[rp][ci][0]})
    available = [f for f in all_factors if f in chars.columns]
    df_work = chars[["id", "eom"] + available]
    months = sorted(df_work["eom"].unique())

    results = []
    for eom in months:
        applicable = [d for d in rebal_dates if d <= eom]
        if not applicable:
            continue
        rp = applicable[-1]
        cfg_ids = chosen[rp]

        month_data = df_work[df_work["eom"] == eom]
        if len(month_data) < 50:
            continue

        needed = sorted({f for ci in cfg_ids for f in schedule[rp][ci][0]
                         if f in month_data.columns})
        if not needed:
            continue
        ranks_all = month_data[needed].rank(pct=True)

        stock_w = np.zeros(len(month_data))
        blend = 1.0 / len(cfg_ids)
        for ci in cfg_ids:
            factors, weights = schedule[rp][ci]
            for factor, fw in zip(factors, weights):
                if factor not in ranks_all.columns:
                    continue
                r = ranks_all[factor].values
                nan_mask = np.isnan(r)
                rw = np.where(nan_mask, 0.0, r - 0.5)
                denom = np.abs(rw).sum()
                if denom < 1e-10 or (~nan_mask).sum() < MIN_STOCKS * 2:
                    continue
                stock_w += blend * fw * rw / denom

        nz = np.abs(stock_w) > 1e-15
        if nz.any():
            results.append(pd.DataFrame({"id": month_data["id"].values[nz],
                                         "eom": eom, "w": stock_w[nz]}))

    output = pd.concat(results, ignore_index=True)
    print(f"    {len(output):,} rows, {output['eom'].nunique()} months "
          f"({time.time()-t0:.0f}s)")
    return output


# ===========================================================================
# Entry point (course interface): main(chars, features, daily_ret)
# ===========================================================================
def main(chars, features, daily_ret):
    feature_names = features["features"].tolist()
    chars = chars.copy()
    chars["eom"] = pd.to_datetime(chars["eom"])
    n_months = chars["eom"].nunique()
    print(f"Data: {len(chars):,} stock-months, {len(feature_names)} "
          f"features, {n_months} months")

    if n_months < 200:   # small validation datasets: loosen thresholds
        min_obs, min_stocks = max(6, n_months // 5), 5
        min_history, meta_min_obs = 12, max(24, n_months // 3)
    else:
        min_obs, min_stocks = MIN_OBS, MIN_STOCKS
        min_history, meta_min_obs = MIN_HISTORY, META_MIN_OBS

    ls = compute_factor_ls_returns(chars, feature_names, min_stocks, min_obs)
    rebal_points, schedule, oos = simulate_configs(
        ls, CONFIG_GRID, REBAL_MONTHS, min_obs, min_history)
    chosen, meta_log = meta_select(
        rebal_points, schedule, oos, CONFIG_GRID, META_TOP_K, meta_min_obs)

    main.meta_log = meta_log          # diagnostics for graders/notebooks
    main.config_oos = oos
    main.config_labels = [config_label(c) for c in CONFIG_GRID]

    output = compute_stock_weights(chars, schedule, chosen)
    assert set(output.columns) == {"id", "eom", "w"}
    assert output["w"].notna().all() and len(output) > 0
    return output


# ===========================================================================
# Stand-alone runner with evaluation
# ===========================================================================
if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--data", required=True,
                    help="directory containing ctff_*.parquet")
    ap.add_argument("--out", default=None,
                    help="path for the weights CSV "
                         "(default: <data>/weights.csv)")
    args = ap.parse_args()

    chars = pd.read_parquet(f"{args.data}/ctff_chars.parquet")
    features = pd.read_parquet(f"{args.data}/ctff_features.parquet")
    weights = main(chars, features,
                   pd.DataFrame(columns=["id", "date", "ret_exc"]))

    # Gross portfolio returns: weights at eom t earn ret_exc_lead1m of t
    chars["eom"] = pd.to_datetime(chars["eom"])
    weights["eom"] = pd.to_datetime(weights["eom"])
    merged = weights.merge(
        chars[["id", "eom", "ret_exc_lead1m"]].dropna(),
        on=["id", "eom"], how="inner")
    port = (merged["w"] * merged["ret_exc_lead1m"]).groupby(
        merged["eom"]).sum().sort_index()

    def sharpe(r):
        return r.mean() / r.std() * np.sqrt(12) if len(r) >= 12 else np.nan

    print("\n  SHARPE RATIOS BY PERIOD (gross)")
    for name, (a, b) in {
        "1960-1989": ("1960", "1989"), "1990-2003": ("1990", "2003"),
        "2004-2013": ("2004", "2013"), "2014-2023": ("2014", "2023"),
        "1990-2023": ("1990", "2023"),
    }.items():
        sub = port.loc[a:b]
        print(f"  {name}: {sharpe(sub):7.3f}  ({len(sub)} months)")

    # One-way turnover (per month, as fraction of gross book)
    piv = weights.pivot_table(index="eom", columns="id", values="w",
                              fill_value=0.0)
    to = piv.diff().abs().sum(axis=1) / 2
    print(f"\n  Avg monthly one-way turnover: {to.iloc[1:].mean():.3f}")

    out_path = args.out or f"{args.data}/weights.csv"
    submission = weights[["id", "eom", "w"]].copy()
    submission["eom"] = submission["eom"].dt.strftime("%Y-%m-%d")
    submission.to_csv(out_path, index=False)
    print(f"  Wrote {len(submission):,} weight rows -> {out_path}")
