#!/usr/bin/env python3
"""
OmmAlpha ML Phase 1B
-------------------
Train a leakage-conscious breakout-success classifier from the CSV exported by
admin/export_ml_dataset_v1.php.

Target definition:
    TARGET  -> 1
    STOP/TIMEOUT -> 0
    PENDING / AMBIGUOUS / NOT_TRIGGERED -> excluded

The split is chronological by signal DATE. Rows are never randomly shuffled
between train/validation/test periods.
"""

from __future__ import annotations

import argparse
import hashlib
import json
import math
from pathlib import Path
from typing import Dict, List, Tuple

import joblib
import numpy as np
import pandas as pd
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, DMatrix

DATASET_VERSION = "ml-dataset-v1b"

# Intentionally excludes future/outcome columns.
# Breadth/market context is exported for later versions, but v1 does not use
# approximate historical breadth until its membership quality is trustworthy.
FEATURE_CANDIDATES = [
    "close_price",
    "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_20",
    "avg_volume_50",
    "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 XGBoost v1")
    p.add_argument("--csv", required=True, help="ML Dataset v1 CSV")
    p.add_argument("--out", default="./model_output", help="Output directory")
    p.add_argument(
        "--min-feature-coverage",
        type=float,
        default=0.60,
        help="Drop feature if non-null coverage is below this fraction (default 0.60)",
    )
    p.add_argument(
        "--min-rows",
        type=int,
        default=200,
        help="Minimum eligible rows required to train (default 200)",
    )
    return p.parse_args()


def safe_metric(fn, y_true, y_prob_or_pred):
    try:
        return float(fn(y_true, y_prob_or_pred))
    except Exception:
        return None


def precision_top_fraction(y_true: np.ndarray, probs: np.ndarray, frac: float = 0.10) -> 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, float | None]:
    pred = (probs >= 0.50).astype(int)
    base_rate = float(np.mean(y_true)) if len(y_true) else None
    top10 = precision_top_fraction(y_true, probs, 0.10)

    metrics = {
        "rows": int(len(y_true)),
        "positive_rate": base_rate,
        "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),
        "precision_top_10pct": top10,
        "lift_top_10pct": (top10 / base_rate) if (top10 is not None and base_rate and base_rate > 0) else None,
    }
    return metrics


def chronological_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) < 12:
        raise ValueError("Need at least 12 distinct signal dates for a chronological split.")

    train_end_i = max(1, int(len(dates) * 0.70))
    val_end_i = max(train_end_i + 1, int(len(dates) * 0.85))
    val_end_i = min(val_end_i, len(dates) - 1)

    train_dates = dates[:train_end_i]
    val_dates = dates[train_end_i:val_end_i]
    test_dates = dates[val_end_i:]

    train = df[df["signal_date"].isin(train_dates)].copy()
    val = df[df["signal_date"].isin(val_dates)].copy()
    test = df[df["signal_date"].isin(test_dates)].copy()

    ranges = {
        "train": f"{train_dates[0]} to {train_dates[-1]}",
        "validation": f"{val_dates[0]} to {val_dates[-1]}",
        "test": f"{test_dates[0]} to {test_dates[-1]}",
    }
    return train, val, test, ranges


def main() -> None:
    args = parse_args()
    csv_path = Path(args.csv)
    out_dir = Path(args.out)
    out_dir.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_date",
        "leakage_audit_pass",
        "eligible_success_model",
        "label_target_before_stop",
    }
    missing_required = sorted(required - set(df.columns))
    if missing_required:
        raise SystemExit(f"Missing required columns: {missing_required}")

    versions = set(df["dataset_version"].dropna().astype(str).unique())
    if versions and versions != {DATASET_VERSION}:
        raise SystemExit(f"Unexpected dataset version(s): {sorted(versions)}")

    df["signal_date"] = pd.to_datetime(df["signal_date"], errors="coerce")

    # Only leakage-audited, resolved samples for this first objective.
    df = df[
        (pd.to_numeric(df["leakage_audit_pass"], errors="coerce") == 1)
        & (pd.to_numeric(df["eligible_success_model"], errors="coerce") == 1)
    ].copy()
    df["target"] = pd.to_numeric(df["label_target_before_stop"], errors="coerce")
    df = df[df["target"].isin([0, 1]) & df["signal_date"].notna()].copy()
    df["target"] = df["target"].astype(int)

    if len(df) < args.min_rows:
        raise SystemExit(
            f"Only {len(df)} eligible rows. Need at least {args.min_rows}. "
            "Run larger historical backtests first."
        )

    available = [c for c in FEATURE_CANDIDATES if c in df.columns]
    coverage = {
        c: float(pd.to_numeric(df[c], errors="coerce").notna().mean())
        for c in available
    }
    features = [c for c in available if coverage[c] >= args.min_feature_coverage]

    if len(features) < 8:
        raise SystemExit(
            f"Only {len(features)} features meet coverage threshold. "
            "Dataset is not ready for meaningful training."
        )

    # Coerce model features to numeric. XGBoost naturally handles NaN.
    for c in features:
        df[c] = pd.to_numeric(df[c], errors="coerce")

    train, val, test, ranges = chronological_split(df)
    if min(len(train), len(val), len(test)) == 0:
        raise SystemExit("Chronological split produced an empty partition.")

    X_train, y_train = train[features], train["target"].to_numpy()
    X_val, y_val = val[features], val["target"].to_numpy()
    X_test, y_test = test[features], test["target"].to_numpy()

    # ------------------------- Baseline -------------------------
    baseline = Pipeline([
        ("imputer", SimpleImputer(strategy="median")),
        ("scale", StandardScaler()),
        ("model", LogisticRegression(max_iter=3000, class_weight="balanced", random_state=42)),
    ])
    baseline.fit(X_train, y_train)
    baseline_val_prob = baseline.predict_proba(X_val)[:, 1]
    baseline_test_prob = baseline.predict_proba(X_test)[:, 1]

    # ------------------------- XGBoost --------------------------
    positives = max(1, int(np.sum(y_train == 1)))
    negatives = max(1, int(np.sum(y_train == 0)))
    scale_pos_weight = negatives / positives

    model = XGBClassifier(
        objective="binary:logistic",
        n_estimators=650,
        learning_rate=0.03,
        max_depth=4,
        min_child_weight=5,
        subsample=0.80,
        colsample_bytree=0.80,
        reg_alpha=0.10,
        reg_lambda=2.0,
        gamma=0.0,
        scale_pos_weight=scale_pos_weight,
        eval_metric="logloss",
        tree_method="hist",
        random_state=42,
        n_jobs=-1,
    )
    model.fit(X_train, y_train, eval_set=[(X_val, y_val)], verbose=False)

    xgb_val_prob = model.predict_proba(X_val)[:, 1]
    xgb_test_prob = model.predict_proba(X_test)[:, 1]

    metrics = {
        "dataset_version": DATASET_VERSION,
        "target": "label_target_before_stop",
        "target_definition": "TARGET=1; STOP/TIMEOUT=0; other outcomes excluded",
        "rows_total_eligible": int(len(df)),
        "features_used": features,
        "feature_coverage": coverage,
        "dropped_for_low_coverage": [c for c in available if c not in features],
        "split_ranges": ranges,
        "split_rows": {"train": len(train), "validation": len(val), "test": len(test)},
        "baseline": {
            "validation": evaluate(y_val, baseline_val_prob),
            "test": evaluate(y_test, baseline_test_prob),
        },
        "xgboost": {
            "validation": evaluate(y_val, xgb_val_prob),
            "test": evaluate(y_test, xgb_test_prob),
        },
    }

    # Explicit comparison. We do not assume XGBoost is better.
    base_pr = metrics["baseline"]["test"]["pr_auc"]
    xgb_pr = metrics["xgboost"]["test"]["pr_auc"]
    metrics["xgboost_beats_baseline_on_test_pr_auc"] = (
        None if base_pr is None or xgb_pr is None else bool(xgb_pr > base_pr)
    )

    feature_hash = hashlib.sha256("\n".join(features).encode("utf-8")).hexdigest()
    metrics["feature_schema_hash"] = feature_hash

    # Save artifacts.
    joblib.dump(baseline, out_dir / "baseline_logistic.joblib")
    model.save_model(out_dir / "xgb_breakout_success_v1.json")

    importance = pd.DataFrame({
        "feature": features,
        "gain_importance": model.feature_importances_,
    }).sort_values("gain_importance", ascending=False)
    importance.to_csv(out_dir / "feature_importance.csv", index=False)

    # XGBoost contribution values for explainability on the held-out test set.
    booster = model.get_booster()
    dtest = DMatrix(X_test, feature_names=features)
    contrib = booster.predict(dtest, pred_contribs=True)
    mean_abs = np.mean(np.abs(contrib[:, :-1]), axis=0)
    shap_global = pd.DataFrame({
        "feature": features,
        "mean_abs_contribution": mean_abs,
    }).sort_values("mean_abs_contribution", ascending=False)
    shap_global.to_csv(out_dir / "global_contributions.csv", index=False)

    scored = test[["signal_date", "symbol", "signal_id", "outcome", "target"]].copy()
    scored["baseline_probability"] = baseline_test_prob
    scored["xgb_probability"] = xgb_test_prob
    scored = scored.sort_values(["signal_date", "xgb_probability"], ascending=[True, False])
    scored.to_csv(out_dir / "test_predictions.csv", index=False)

    with open(out_dir / "metrics.json", "w", encoding="utf-8") as f:
        json.dump(metrics, f, indent=2, default=str)

    model_card = f"""OmmAlpha XGBoost v1\n\nObjective: breakout success after trigger\nDataset: {DATASET_VERSION}\nRows: {len(df)}\nFeatures: {len(features)}\nTrain: {ranges['train']}\nValidation: {ranges['validation']}\nTest: {ranges['test']}\n\nBaseline test PR-AUC: {base_pr}\nXGBoost test PR-AUC: {xgb_pr}\nXGBoost beats baseline: {metrics['xgboost_beats_baseline_on_test_pr_auc']}\n\nDo not deploy merely because training completed. Deploy only after out-of-sample\nmetrics, calibration and top-ranked expectancy are acceptable.\n"""
    (out_dir / "MODEL_CARD.txt").write_text(model_card, encoding="utf-8")

    print(json.dumps({
        "status": "trained",
        "rows": len(df),
        "features": len(features),
        "split_ranges": ranges,
        "baseline_test": metrics["baseline"]["test"],
        "xgboost_test": metrics["xgboost"]["test"],
        "xgboost_beats_baseline": metrics["xgboost_beats_baseline_on_test_pr_auc"],
        "output": str(out_dir.resolve()),
    }, indent=2, default=str))


if __name__ == "__main__":
    main()
