#!/usr/bin/env python3
"""
OmmAlpha ML Phase 2A
====================
Automated, chronological, leakage-conscious training for three objectives:

1) trigger      : will the setup trigger an entry?
2) direct_target: from today's setup, will it eventually reach TARGET?
3) success      : after trigger, will TARGET occur before STOP/TIMEOUT?

The script intentionally compares Logistic Regression and XGBoost rather than
assuming XGBoost must win. A champion is selected using VALIDATION PR-AUC;
the TEST period remains untouched for final reporting. The production champion
is saved together with its selected calibration for Phase 2B inference.
"""
from __future__ import annotations

import argparse
import hashlib
import json
import math
import os
import platform
import sys
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List, Tuple

# Shared hosting safety: prevent BLAS/XGBoost from trying to consume all CPUs.
os.environ.setdefault("OMP_NUM_THREADS", "1")
os.environ.setdefault("OPENBLAS_NUM_THREADS", "1")
os.environ.setdefault("MKL_NUM_THREADS", "1")
os.environ.setdefault("NUMEXPR_NUM_THREADS", "1")

import joblib
import numpy as np
import pandas as pd
import sklearn
import xgboost
from sklearn.impute import SimpleImputer
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import (
    accuracy_score,
    average_precision_score,
    brier_score_loss,
    f1_score,
    log_loss,
    precision_score,
    recall_score,
    roc_auc_score,
)
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
from xgboost import XGBClassifier
from probability_calibration import apply_calibration, log_odds

DATASET_VERSION = "ml-dataset-v1b"
TRAINER_VERSION = "phase2a-1.3"

# Normalised / structurally meaningful inputs only. We deliberately exclude:
# - absolute close_price
# - absolute raw volumes
# - outcome/future fields
# - approximate market breadth for this first production candidate
# Constant features are also removed automatically per objective.
FEATURE_CANDIDATES = [
    "price_vs_sma20_pct",
    "price_vs_sma50_pct",
    "price_vs_sma150_pct",
    "price_vs_sma200_pct",
    "sma200_rising",
    "trend_template_pass",
    "atr_pct",
    "dist_52w_high_pct",
    "dist_52w_low_pct",
    "avg_volume_ratio_20_50",
    "avg_traded_value_20_log10",
    "return_3m_pct",
    "return_6m_pct",
    "return_12m_pct",
    "momentum_score",
    "rs_percentile",
    "rsi14",
    "adx14",
    "plus_di14",
    "minus_di14",
    "bb_width20_pct",
    "relative_volume20",
    "base_depth60_pct",
    "gap_pct",
    "pivot_distance_pct",
    "vcp_score",
    "contraction_count",
    "contraction_1_pct",
    "contraction_2_pct",
    "contraction_3_pct",
    "contraction_4_pct",
    "final_contraction_pct",
    "progressive_contractions",
    "volume_ratio_5_20",
    "volume_dryup",
]


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="Train OmmAlpha Phase 2A models")
    p.add_argument("--csv", required=True)
    p.add_argument("--out", required=True)
    p.add_argument("--min-feature-coverage", type=float, default=0.60)
    p.add_argument("--min-rows", type=int, default=200)
    p.add_argument("--walk-forward-folds", type=int, default=3)
    return p.parse_args()


def safe_metric(fn, a, b):
    try:
        return float(fn(a, b))
    except Exception:
        return None


def top_fraction_rate(y_true: np.ndarray, probs: np.ndarray, frac: float) -> float | None:
    if len(y_true) == 0:
        return None
    k = max(1, int(math.ceil(len(y_true) * frac)))
    order = np.argsort(-probs)[:k]
    return float(np.mean(y_true[order]))


def evaluate(y_true: np.ndarray, probs: np.ndarray) -> Dict[str, Any]:
    probs = np.asarray(probs, dtype=float)
    pred = (probs >= 0.50).astype(int)
    base = float(np.mean(y_true)) if len(y_true) else None
    t10 = top_fraction_rate(y_true, probs, 0.10)
    t20 = top_fraction_rate(y_true, probs, 0.20)
    return {
        "rows": int(len(y_true)),
        "positive_rate": base,
        "reliability_bins": reliability_bins(y_true, probs),
        "roc_auc": safe_metric(roc_auc_score, y_true, probs),
        "pr_auc": safe_metric(average_precision_score, y_true, probs),
        "brier": safe_metric(brier_score_loss, y_true, probs),
        "log_loss": safe_metric(log_loss, y_true, np.clip(probs, 1e-6, 1-1e-6)),
        "accuracy_at_0_5": safe_metric(accuracy_score, y_true, pred),
        "precision_at_0_5": safe_metric(lambda a,b: precision_score(a,b,zero_division=0), y_true, pred),
        "recall_at_0_5": safe_metric(lambda a,b: recall_score(a,b,zero_division=0), y_true, pred),
        "f1_at_0_5": safe_metric(lambda a,b: f1_score(a,b,zero_division=0), y_true, pred),
        "top_10pct_positive_rate": t10,
        "top_10pct_lift": (t10/base) if (t10 is not None and base and base > 0) else None,
        "top_20pct_positive_rate": t20,
        "top_20pct_lift": (t20/base) if (t20 is not None and base and base > 0) else None,
    }


def reliability_bins(y_true, probs):
    bins = []
    indices = np.minimum((np.asarray(probs) * 10).astype(int), 9)
    for i in range(10):
        mask = indices == i
        if mask.any():
            bins.append({"lower": i / 10, "upper": (i + 1) / 10,
                         "rows": int(mask.sum()),
                         "mean_probability": float(np.mean(probs[mask])),
                         "observed_rate": float(np.mean(np.asarray(y_true)[mask]))})
    return bins


def label_resolution_dates(df, objective):
    """Use actual event dates; legacy exports retain conservative fallback."""
    observed = pd.to_datetime(df["evaluated_through"], errors="coerce")
    if "entry_window_end_date" in df:
        expiry = pd.to_datetime(df["entry_window_end_date"], errors="coerce")
        valid_expiry = df.outcome.eq("NOT_TRIGGERED") & expiry.notna() & (expiry > df.signal_date) & (observed.isna() | (expiry <= observed))
        observed = observed.where(~valid_expiry, expiry)
    column = "entry_date" if objective == "trigger" else "exit_date"
    if column not in df:
        return observed
    event = pd.to_datetime(df[column], errors="coerce")
    allowed = df.target.eq(1) if objective == "trigger" else df.outcome.isin(["TARGET", "STOP", "TIMEOUT"])
    valid = allowed & event.notna() & (event > df.signal_date) & (observed.isna() | (event <= observed))
    return observed.where(~valid, event)


def purge_unresolved(df: pd.DataFrame, cutoff) -> pd.DataFrame:
    """Only use labels fully observed before the next evaluation period."""
    if "evaluated_through" not in df:
        raise ValueError("Missing evaluated_through: re-export historical training data")
    observed = pd.to_datetime(df.get("label_resolved_date", df["evaluated_through"]), errors="coerce")
    return df[observed.notna() & (observed >= df.signal_date) & (observed < cutoff)].copy()


def date_split(df: pd.DataFrame) -> Tuple[pd.DataFrame,pd.DataFrame,pd.DataFrame,Dict[str,str]]:
    dates = np.array(sorted(df["signal_date"].dropna().unique()))
    if len(dates) < 20:
        raise ValueError("Need at least 20 distinct signal dates.")
    train_end = max(1, int(len(dates)*0.70))
    val_end = max(train_end+1, int(len(dates)*0.85))
    val_end = min(val_end, len(dates)-1)
    td, vd, xd = dates[:train_end], dates[train_end:val_end], dates[val_end:]
    tr = df[df.signal_date.isin(td)].copy()
    va = df[df.signal_date.isin(vd)].copy()
    te = df[df.signal_date.isin(xd)].copy()
    tr = purge_unresolved(tr, vd[0])
    va = purge_unresolved(va, xd[0])
    ranges = {
        "train": f"{pd.Timestamp(td[0]).date()} to {pd.Timestamp(td[-1]).date()}",
        "validation": f"{pd.Timestamp(vd[0]).date()} to {pd.Timestamp(vd[-1]).date()}",
        "test": f"{pd.Timestamp(xd[0]).date()} to {pd.Timestamp(xd[-1]).date()}",
    }
    return tr,va,te,ranges


def select_features(df: pd.DataFrame, min_cov: float) -> Tuple[List[str],Dict[str,float],List[str],List[str]]:
    if not 0 < min_cov <= 1:
        raise ValueError("min-feature-coverage must be in (0, 1]")
    available = [c for c in FEATURE_CANDIDATES if c in df.columns]
    coverage = {c: float(pd.to_numeric(df[c], errors="coerce").replace([np.inf, -np.inf], np.nan).notna().mean()) for c in available}
    low_cov = [c for c in available if coverage[c] < min_cov]
    candidates = [c for c in available if coverage[c] >= min_cov]
    constants=[]
    features=[]
    for c in candidates:
        s = pd.to_numeric(df[c], errors="coerce").replace([np.inf, -np.inf], np.nan)
        if s.dropna().nunique() <= 1:
            constants.append(c)
        else:
            features.append(c)
    return features,coverage,low_cov,constants


def make_logistic() -> Pipeline:
    return Pipeline([
        ("imputer", SimpleImputer(strategy="median")),
        ("scale", StandardScaler()),
        ("model", LogisticRegression(max_iter=3000, class_weight="balanced", random_state=42)),
    ])


def make_xgb(y: np.ndarray, params=None) -> XGBClassifier:
    pos=max(1,int(np.sum(y==1))); neg=max(1,int(np.sum(y==0)))
    model = XGBClassifier(
        objective="binary:logistic",
        n_estimators=350,
        learning_rate=0.035,
        max_depth=3,
        min_child_weight=5,
        subsample=0.82,
        colsample_bytree=0.82,
        reg_alpha=0.15,
        reg_lambda=2.5,
        gamma=0.0,
        scale_pos_weight=neg/pos,
        eval_metric="logloss",
        tree_method="hist",
        random_state=42,
        n_jobs=1,
    )
    if params:
        model.set_params(**params)
    return model


def tune_xgb(training, min_cov):
    """Small deterministic search confined to purged training history."""
    candidates = [{}, {"max_depth": 2, "n_estimators": 200},
                  {"max_depth": 3, "min_child_weight": 10, "reg_lambda": 5.0}]
    dates = sorted(training.signal_date.unique())
    folds = []
    for start_fraction, end_fraction in [(0.5, 0.67), (0.67, 0.83), (0.83, 1.0)]:
        start, end = int(len(dates)*start_fraction), int(len(dates)*end_fraction)
        if not start or end <= start:
            continue
        fit = purge_unresolved(training[training.signal_date < dates[start]], dates[start])
        check = training[training.signal_date.isin(dates[start:end])]
        features, _, _, _ = select_features(fit, min_cov)
        if len(fit) < 100 or len(check) < 30 or len(features) < 8 or fit.target.nunique() < 2 or check.target.nunique() < 2:
            continue
        folds.append((fit, check, features))
    if len(folds) < 2:
        return {}, {"status": "default", "reason": "fewer than two usable chronological folds", "folds": len(folds)}
    scores = []
    for params in candidates:
        metrics = []
        for fit, check, features in folds:
            model = make_xgb(fit.target.to_numpy(), params)
            model.fit(fit[features], fit.target.to_numpy(), verbose=False)
            p = model.predict_proba(check[features])[:, 1]
            metrics.append(evaluate(check.target.to_numpy(), p))
        scores.append({"params": params, "mean_pr_auc": float(np.mean([m["pr_auc"] for m in metrics])),
                       "mean_brier": float(np.mean([m["brier"] for m in metrics]))})
    best = sorted(scores, key=lambda s: (-s["mean_pr_auc"], s["mean_brier"]))[0]
    return best["params"], {"status": "tuned", "folds": len(folds), "candidates": scores,
                            "selected_params": best["params"], "basis": "training-only chronological folds"}


def deployment_check(metrics, baseline):
    """Validation-only readiness checks; test results never choose deployment."""
    reasons = []
    if metrics.get("rows", 0) < 50:
        reasons.append("fewer than 50 validation rows")
    for key in ("pr_auc", "brier", "top_10pct_lift", "positive_rate"):
        if metrics.get(key) is None or not np.isfinite(metrics[key]):
            reasons.append(f"missing or non-finite {key}")
    if baseline.get("brier") is None or not np.isfinite(baseline["brier"]):
        reasons.append("missing or non-finite baseline Brier")
    if not reasons:
        if metrics["pr_auc"] <= metrics["positive_rate"]:
            reasons.append("PR-AUC does not beat validation prevalence")
        if metrics["brier"] >= baseline["brier"]:
            reasons.append("Brier does not beat training-prevalence baseline")
        if metrics["top_10pct_lift"] <= 1:
            reasons.append("top-decile lift is not above one")
    return {"eligible": not reasons, "reasons": reasons, "basis": "validation only", "baseline": baseline}


def build_objective_df(all_df: pd.DataFrame, objective: str) -> pd.DataFrame:
    d = all_df[(all_df["leakage_audit_pass"]==1) & all_df["signal_date"].notna()].copy()
    if objective == "trigger":
        eligible = pd.to_numeric(d["eligible_entry_model"], errors="coerce") == 1
        d=d[eligible].copy()
        d["target"] = pd.to_numeric(d["label_entry_triggered"], errors="coerce")
    elif objective == "success":
        eligible = pd.to_numeric(d["eligible_success_model"], errors="coerce") == 1
        d=d[eligible].copy()
        d["target"] = pd.to_numeric(d["label_target_before_stop"], errors="coerce")
    elif objective == "direct_target":
        # Today's setup -> eventual TARGET. Exclude unresolved/ambiguous outcomes.
        outcome=d["outcome"].fillna("").astype(str).str.upper()
        keep=outcome.isin(["TARGET","STOP","TIMEOUT","NOT_TRIGGERED"])
        d=d[keep].copy()
        d["target"]=(d["outcome"].astype(str).str.upper()=="TARGET").astype(int)
    else:
        raise ValueError(objective)
    d=d[pd.to_numeric(d["target"],errors="coerce").isin([0,1])].copy()
    d["target"]=pd.to_numeric(d["target"],errors="coerce").astype(int)
    d["outcome"] = d["outcome"].fillna("").astype(str).str.upper().str.strip()
    d["label_resolved_date"] = label_resolution_dates(d, objective)
    return d


def calibration_split(validation):
    dates = sorted(validation.signal_date.unique())
    if len(dates) < 4:
        return validation.iloc[:0].copy(), validation
    boundary = dates[len(dates) // 2]
    cal = purge_unresolved(validation[validation.signal_date < boundary], boundary)
    selection = validation[validation.signal_date >= boundary].copy()
    if len(cal) < 50 or len(selection) < 30 or cal.target.nunique() < 2 or selection.target.nunique() < 2:
        return validation.iloc[:0].copy(), validation
    return cal, selection


def fit_calibration(y, probabilities):
    if len(y) < 50 or len(np.unique(y)) < 2:
        return {"method": "none", "reason": "insufficient calibration data"}
    model = LogisticRegression(C=1.0, max_iter=1000, random_state=42)
    model.fit(log_odds(probabilities).reshape(-1, 1), y)
    slope = float(model.coef_[0, 0])
    if slope <= 0:
        return {"method": "none", "reason": "non-increasing calibration rejected"}
    return {"method": "sigmoid_logit", "slope": slope, "intercept": float(model.intercept_[0])}


def select_calibration(y, raw, candidate):
    calibrated = apply_calibration(raw, candidate)
    raw_metrics, calibrated_metrics = evaluate(y, raw), evaluate(y, calibrated)
    use = candidate["method"] != "none" and calibrated_metrics["brier"] < raw_metrics["brier"] and calibrated_metrics["log_loss"] <= raw_metrics["log_loss"]
    chosen = candidate if use else {"method": "none", "reason": "calibration unavailable or did not improve validation Brier and log loss"}
    return chosen, {"raw": raw_metrics, "candidate": calibrated_metrics, "selected": chosen["method"]}


def choose_champion(log_m: Dict[str,Any], xgb_m: Dict[str,Any]) -> str:
    # Validation PR-AUC is primary. Brier breaks near-ties in favour of calibration.
    lp=log_m.get("pr_auc"); xp=xgb_m.get("pr_auc")
    if lp is None: return "xgboost"
    if xp is None: return "logistic"
    if abs(lp-xp) <= 0.01:
        lb=log_m.get("brier"); xb=xgb_m.get("brier")
        if lb is not None and xb is not None:
            return "logistic" if lb <= xb else "xgboost"
    return "logistic" if lp > xp else "xgboost"


def walk_forward(df: pd.DataFrame, features: List[str], folds: int, min_cov: float = 0.60) -> List[Dict[str,Any]]:
    dates=np.array(sorted(df.signal_date.unique()))
    folds=max(0,min(5,folds))
    if folds==0 or len(dates)<80: return []
    initial=max(20,int(len(dates)*0.45))
    remaining=len(dates)-initial
    block=max(10,remaining//folds)
    rows=[]
    for i in range(folds):
        start=initial+i*block
        end=len(dates) if i==folds-1 else min(len(dates),start+block)
        if start>=len(dates) or end<=start: break
        train_dates=dates[:start]; test_dates=dates[start:end]
        tr=df[df.signal_date.isin(train_dates)]; te=df[df.signal_date.isin(test_dates)]
        tr=purge_unresolved(tr, test_dates[0])
        if len(tr)<100 or len(te)<20 or tr.target.nunique()<2 or te.target.nunique()<2: continue
        fold_features,_,_,_=select_features(tr,min_cov)
        if len(fold_features)<8: continue
        Xtr=tr[fold_features]; ytr=tr.target.to_numpy(); Xte=te[fold_features]; yte=te.target.to_numpy()
        lm=make_logistic(); lm.fit(Xtr,ytr); lp=lm.predict_proba(Xte)[:,1]
        xm=make_xgb(ytr); xm.fit(Xtr,ytr,verbose=False); xp=xm.predict_proba(Xte)[:,1]
        rows.append({
            "fold":len(rows)+1,
            "features":fold_features,
            "train_end":str(pd.Timestamp(train_dates[-1]).date()),
            "test_start":str(pd.Timestamp(test_dates[0]).date()),
            "test_end":str(pd.Timestamp(test_dates[-1]).date()),
            "train_rows":int(len(tr)),"test_rows":int(len(te)),
            "base_rate":float(np.mean(yte)),
            "logistic_pr_auc":safe_metric(average_precision_score,yte,lp),
            "logistic_roc_auc":safe_metric(roc_auc_score,yte,lp),
            "logistic_top10_lift":evaluate(yte,lp)["top_10pct_lift"],
            "xgb_pr_auc":safe_metric(average_precision_score,yte,xp),
            "xgb_roc_auc":safe_metric(roc_auc_score,yte,xp),
            "xgb_top10_lift":evaluate(yte,xp)["top_10pct_lift"],
        })
    return rows


def train_one(all_df: pd.DataFrame, objective: str, root: Path, min_cov: float, min_rows: int, wf_folds:int) -> Dict[str,Any]:
    df=build_objective_df(all_df,objective)
    if len(df)<min_rows: raise RuntimeError(f"{objective}: only {len(df)} eligible rows")
    for c in FEATURE_CANDIDATES:
        if c in df: df[c]=pd.to_numeric(df[c],errors="coerce").replace([np.inf,-np.inf],np.nan)
    tr,va,te,ranges=date_split(df)
    calibration_rows, va = calibration_split(va)
    features,coverage,low_cov,constants=select_features(tr,min_cov)
    if len(features)<8: raise RuntimeError(f"{objective}: only {len(features)} usable features")
    for c in features: df[c]=pd.to_numeric(df[c],errors="coerce")
    Xtr,ytr=tr[features],tr.target.to_numpy(); Xv,yv=va[features],va.target.to_numpy(); Xt,yt=te[features],te.target.to_numpy()
    for name, partition in [("train",tr),("validation",va),("test",te)]:
        if partition.target.nunique()<2:
            raise RuntimeError(f"{objective}: {name} split has fewer than two classes after date purging ({partition.target.value_counts().to_dict()}). Re-export with exit_date and entry_window_end_date; do not disable purging.")

    logistic=make_logistic(); logistic.fit(Xtr,ytr)
    lvp=logistic.predict_proba(Xv)[:,1]; ltp=logistic.predict_proba(Xt)[:,1]
    xgb_params, tuning = tune_xgb(tr, min_cov)
    xgbm=make_xgb(ytr, xgb_params); xgbm.fit(Xtr,ytr,verbose=False)
    xvp=xgbm.predict_proba(Xv)[:,1]; xtp=xgbm.predict_proba(Xt)[:,1]
    raw_ltp, raw_xtp = ltp.copy(), xtp.copy()
    calibrations, calibration_diagnostics = {}, {}
    for name, model, validation_probs in [("logistic", logistic, lvp), ("xgboost", xgbm, xvp)]:
        candidate = {"method": "none", "reason": "insufficient calibration data"}
        if not calibration_rows.empty:
            candidate = fit_calibration(calibration_rows.target.to_numpy(), model.predict_proba(calibration_rows[features])[:, 1])
        calibrations[name], calibration_diagnostics[name] = select_calibration(yv, validation_probs, candidate)
    lvp=apply_calibration(lvp,calibrations["logistic"]); ltp=apply_calibration(ltp,calibrations["logistic"])
    xvp=apply_calibration(xvp,calibrations["xgboost"]); xtp=apply_calibration(xtp,calibrations["xgboost"])
    lm_val=evaluate(yv,lvp); xm_val=evaluate(yv,xvp)
    lm_test=evaluate(yt,ltp); xm_test=evaluate(yt,xtp)
    champion=choose_champion(lm_val,xm_val)

    od=root/objective; od.mkdir(parents=True,exist_ok=True)
    joblib.dump(logistic, od/'evaluated_logistic.joblib')
    joblib.dump(xgbm, od/'evaluated_xgboost.joblib')
    xgbm.save_model(od/'evaluated_xgboost.json')

    # Preserve the exact evaluated base model and its out-of-sample calibrator.
    prod=logistic if champion=="logistic" else xgbm
    joblib.dump(prod,od/'production_champion.joblib')
    if champion=="xgboost": prod.save_model(od/'production_champion_xgboost.json')

    # Diagnostics.
    if hasattr(xgbm,"feature_importances_"):
        pd.DataFrame({"feature":features,"xgb_importance":xgbm.feature_importances_}).sort_values("xgb_importance",ascending=False).to_csv(od/'xgb_feature_importance.csv',index=False)
    log_model=logistic.named_steps["model"]
    pd.DataFrame({"feature":features,"logistic_coefficient":log_model.coef_[0],"abs_coefficient":np.abs(log_model.coef_[0])}).sort_values("abs_coefficient",ascending=False).to_csv(od/'logistic_coefficients.csv',index=False)

    scored=te[["signal_date","symbol","signal_id","outcome","target"]].copy()
    scored["logistic_probability"]=ltp; scored["xgb_probability"]=xtp
    scored["logistic_raw_probability"]=raw_ltp; scored["xgb_raw_probability"]=raw_xtp
    scored.to_csv(od/'test_predictions.csv',index=False)
    wf=walk_forward(df,features,wf_folds,min_cov)
    pd.DataFrame(wf).to_csv(od/'walk_forward.csv',index=False)

    result={
        "objective":objective,
        "tuning":tuning,
        "deployment":deployment_check(lm_val if champion=="logistic" else xm_val,
            evaluate(yv,np.full(len(yv),float(np.mean(ytr))))),
        "rows":int(len(df)),
        "positive_rate":float(df.target.mean()),
        "features":features,
        "feature_schema_hash":hashlib.sha256("\n".join(features).encode()).hexdigest(),
        "feature_coverage":coverage,
        "feature_selection_basis":"purged training partition only",
        "purged_rows":int(len(df)-len(tr)-len(va)-len(te)-len(calibration_rows)),
        "label_boundary_policy":"objective resolution date strictly before next partition; evaluated_through fallback",
        "precise_label_rows":int((df.label_resolved_date != pd.to_datetime(df.evaluated_through, errors="coerce")).sum()),
        "production_fit_policy":"preserve evaluated model and calibration; no post-test refit",
        "calibration":calibrations[champion],
        "calibration_diagnostics":calibration_diagnostics,
        "calibration_range":None if calibration_rows.empty else {"start":str(calibration_rows.signal_date.min().date()), "end":str(calibration_rows.signal_date.max().date())},
        "selection_range":{"start":str(va.signal_date.min().date()), "end":str(va.signal_date.max().date())},
        "dropped_low_coverage":low_cov,
        "dropped_constant":constants,
        "split_ranges":ranges,
        "split_rows":{"train":len(tr),"calibration":len(calibration_rows),"validation":len(va),"test":len(te)},
        "logistic":{"validation":lm_val,"test":lm_test},
        "xgboost":{"validation":xm_val,"test":xm_test},
        "champion":champion,
        "champion_selection_basis":"validation PR-AUC; Brier breaks <=0.01 PR-AUC ties",
        "champion_test": lm_test if champion=="logistic" else xm_test,
        "champion_raw_test": evaluate(yt,raw_ltp if champion=="logistic" else raw_xtp),
        "walk_forward":wf,
    }
    (od/'metrics.json').write_text(json.dumps(result,indent=2,default=str),encoding='utf-8')
    (od/'features.json').write_text(json.dumps(features,indent=2),encoding='utf-8')
    return result


def main() -> None:
    args=parse_args(); csv_path=Path(args.csv); out=Path(args.out); out.mkdir(parents=True,exist_ok=True)
    df=pd.read_csv(csv_path,low_memory=False)
    if df.empty: raise SystemExit("Dataset is empty")
    required={"dataset_version","signal_id","symbol","source","evaluated_through","signal_date","leakage_audit_pass","outcome","eligible_entry_model","label_entry_triggered","eligible_success_model","label_target_before_stop"}
    miss=sorted(required-set(df.columns))
    if miss: raise SystemExit(f"Missing columns: {miss}")
    versions=set(df.dataset_version.dropna().astype(str).unique())
    if versions!={DATASET_VERSION}: raise SystemExit(f"Unexpected dataset versions: {sorted(versions)}")
    df["signal_date"]=pd.to_datetime(df.signal_date,errors='coerce')
    if df.signal_date.isna().any(): raise SystemExit("Invalid or missing signal dates")
    if df.signal_id.isna().any(): raise SystemExit("Missing signal IDs")
    if not df.source.fillna('').astype(str).str.strip().str.upper().eq('BACKTEST').all():
        raise SystemExit("Training requires BACKTEST-only rows; keep LIVE data out of training")
    df["leakage_audit_pass"]=pd.to_numeric(df.leakage_audit_pass,errors='coerce').fillna(0)

    # Global integrity audit.
    leakage_rows=int((df.leakage_audit_pass!=1).sum())
    duplicate_signal_ids=int(df.signal_id.duplicated().sum()) if "signal_id" in df else 0
    if leakage_rows>0: raise SystemExit(f"Leakage audit failed for {leakage_rows} rows")
    if duplicate_signal_ids>0: raise SystemExit(f"Duplicate signal IDs: {duplicate_signal_ids}")

    objectives={}
    for objective in ["trigger","direct_target","success"]:
        objectives[objective]=train_one(df,objective,out,args.min_feature_coverage,args.min_rows,args.walk_forward_folds)

    summary={
        "status":"trained",
        "trainer_version":TRAINER_VERSION,
        "dataset_version":DATASET_VERSION,
        "trained_at_utc":datetime.now(timezone.utc).isoformat(),
        "python":sys.version.split()[0],
        "platform":platform.platform(),
        "libraries":{"numpy":np.__version__,"pandas":pd.__version__,"scikit_learn":sklearn.__version__,"xgboost":xgboost.__version__,"joblib":joblib.__version__},
        "dataset_rows":int(len(df)),
        "dataset_min_date":str(df.signal_date.min().date()),
        "dataset_max_date":str(df.signal_date.max().date()),
        "leakage_fail_rows":leakage_rows,
        "duplicate_signal_ids":duplicate_signal_ids,
        "objectives":objectives,
        "deployment":{"eligible":all(v["deployment"]["eligible"] for v in objectives.values()),
                      "reasons":{k:v["deployment"]["reasons"] for k,v in objectives.items() if not v["deployment"]["eligible"]}},
    }
    (out/'phase2_metrics.json').write_text(json.dumps(summary,indent=2,default=str),encoding='utf-8')
    registry={
        "trainer_version":TRAINER_VERSION,
        "dataset_version":DATASET_VERSION,
        "deployment":summary["deployment"],
        "trained_at_utc":summary["trained_at_utc"],
        "models":{
            k:{"champion":v["champion"],"artifact":f"{k}/production_champion.joblib","features":f"{k}/features.json","metrics":f"{k}/metrics.json","calibration":v["calibration"]}
            for k,v in objectives.items()
        }
    }
    (out/'model_registry.json').write_text(json.dumps(registry,indent=2),encoding='utf-8')
    print(json.dumps({"status":"trained","dataset_rows":len(df),"output":str(out.resolve()),"champions":{k:v['champion'] for k,v in objectives.items()},"test_metrics":{k:v['champion_test'] for k,v in objectives.items()}},indent=2,default=str))

if __name__=='__main__':
    main()
