"""
Agente Meta-Labeling de Leonex.

Implementa el Meta-Labeling descrito por Marcos Lopez de Prado en Advances in
Financial Machine Learning (Wiley 2018), capitulo 3.6 (concepto), capitulo 4
(sample weights por uniqueness) y capitulo 7 (purged k-fold + walk-forward).

REFERENCIAS BIBLIOGRAFICAS
- Lopez de Prado, M. (2018). Advances in Financial Machine Learning. Wiley.
  Capitulo 3 (Triple Barrier + Meta-Labeling), capitulo 4 (sample uniqueness),
  capitulo 7 (Purged CV + walk-forward), capitulo 8 (feature importance).
- Lopez de Prado, M. (2023). Causal Factor Investing. Cambridge Elements.
  Capitulo 5 (PC algorithm para distinguir features causales de spurious),
  base del modo --use-causal-features de este agente.
- Hastie, Tibshirani, Friedman (2009). The Elements of Statistical Learning,
  2nd ed. Stanford. Capitulos 4-10 sobre clasificacion binaria + gradient
  boosting. Disponible libre en hastie.su.domains/ElemStatLearn.
- Aronson, D. (2007). Evidence-Based Technical Analysis. Wiley. Capitulo 6:
  multiple testing problem — motivo por el que el AUC OOF se mide contra el
  baseline 0.50 y un AUC 0.516 (caso actual de Leonex) NO es edge real.

El sistema funciona en dos capas:

    Modelo PRIMARIO  →  agente_senales.py produce señales direccionales
    Modelo SECUNDARIO →  agente_triple_barrier.py simula trades cerrados
                          con TP, SL y timeout
    Modelo META      →  ESTE agente. Aprende a filtrar los trades del Triple
                          Barrier prediciendo si seran ganadores o no.

Sin Meta:
    1377 trades → win rate 42.6%, profit factor 1.27, MC p-value 0.99

Con Meta (esperado tras filtrar):
    ~600-700 trades → win rate 55-65%, PF 1.6-2.0, MC p-value < 0.05

El modelo es un clasificador binario (LightGBM con fallback a
GradientBoostingClassifier de sklearn). Entrena con Walk-Forward de 3 folds.

Uso:
    python agents/agente_meta.py
    python agents/agente_meta.py --strategy regime_adaptive_v1_lo
    python agents/agente_meta.py --thresholds 0.50,0.55,0.60,0.65
"""

from __future__ import annotations

import argparse
import json
import logging
import math
import sqlite3
from dataclasses import asdict, dataclass, field
from datetime import datetime
try:
    from datetime import UTC  # Python 3.11+
except ImportError:  # Python 3.10
    from datetime import timezone
    UTC = timezone.utc
from pathlib import Path
from typing import Optional

import numpy as np
import pandas as pd


PROJECT_ROOT = Path(__file__).resolve().parents[1]
DATA_DIR = PROJECT_ROOT / "data"
LOGS_DIR = PROJECT_ROOT / "logs"
DASHBOARD_DATA_DIR = PROJECT_ROOT / "dashboard" / "data"
DB_PATH = DATA_DIR / "Leonex.sqlite"
REPORT_PATH = DASHBOARD_DATA_DIR / "meta_report.json"

DEFAULT_STRATEGY = "regime_adaptive_v1_lo"
DEFAULT_N_FOLDS = 3
DEFAULT_THRESHOLDS = (0.50, 0.55, 0.60, 0.65)
CAUSAL_REPORT_PATH = DASHBOARD_DATA_DIR / "causal_report.json"


def load_causal_features() -> tuple[list[str], list[str]]:
    """Lee causal_report.json y devuelve (causal_features, spurious_features).
    Si no existe o esta vacio, devuelve ([], [])."""
    if not CAUSAL_REPORT_PATH.exists():
        return [], []
    try:
        data = json.loads(CAUSAL_REPORT_PATH.read_text(encoding="utf-8"))
        return (
            list(data.get("causal_features") or []),
            list(data.get("spurious_features") or []),
        )
    except Exception:
        return [], []


@dataclass
class ThresholdMetrics:
    threshold: float
    n_trades_keep: int
    n_trades_skip: int
    keep_rate: float          # % de trades aceptados
    win_rate_keep: float      # % de ganadores entre aceptados
    avg_return_keep: float    # retorno medio entre aceptados (%)
    total_return_keep: float  # suma de retornos aceptados (%)
    profit_factor_keep: float
    expectancy_keep: float


@dataclass
class MetaReport:
    generated_at: str
    strategy: str
    model: str
    n_trades_total: int
    n_features: int
    feature_names: list[str] = field(default_factory=list)
    auc_oof: float = 0.0       # AUC out-of-fold
    accuracy_oof: float = 0.0
    baseline_win_rate: float = 0.0
    metrics_by_threshold: list[ThresholdMetrics] = field(default_factory=list)
    n_folds: int = DEFAULT_N_FOLDS
    feature_importance: list[tuple[str, float]] = field(default_factory=list)
    # Modo causal: si esta activo, solo se entrena con features que el
    # PC algorithm marca como causales (no spurious).
    causal_mode: str = "all_features"        # "all_features" | "causal_only"
    causal_features_available: list[str] = field(default_factory=list)
    spurious_features_excluded: list[str] = field(default_factory=list)
    features_filter_note: str = ""
    # Sampling uniqueness (López de Prado AFML cap.4): corrige overlapping
    # events del Triple Barrier ponderando trades correlacionados.
    sample_weighting: str = "none"            # "none" | "uniqueness"
    avg_uniqueness: float = 1.0               # 1.0 = sin solape; <1 = trades correlados
    n_trades_overlapped: int = 0              # trades con uniqueness < 0.95
    # ── Comparison pass: SIEMPRE corremos el modo opuesto sobre los mismos
    # folds para que el usuario vea si el filtro causal-only le esta costando
    # edge predictivo (AUC). El primary es el modo activo en produccion;
    # comparison es el contrafactico. delta_auc > 0 ⇒ primary mejor.
    auc_oof_primary: float = 0.0              # = auc_oof (mantiene compat)
    auc_oof_comparison: Optional[float] = None
    comparison_mode: str = "n/a"              # "causal_only" | "extended_all" | "n/a"
    comparison_n_features: int = 0
    delta_auc: Optional[float] = None         # primary - comparison
    comparison_note: str = ""


def _try_lightgbm():
    try:
        import lightgbm as lgb
        return lgb
    except ImportError:
        return None


def _fit_predict(X_train, y_train, X_test, sample_weight=None):
    """Entrena modelo binario con sample_weight opcional (uniqueness).

    sample_weight: None o array shape (n_train,) con peso por muestra.
    AFML cap.4: usar uniqueness para corregir overlapping events del
    Triple Barrier. Sin pesos, el modelo sobreestima informacion cuando
    dos trades del mismo ticker se solapan en el tiempo.
    """
    lgb = _try_lightgbm()
    if lgb is not None:
        # LightGBM: rapido, robusto, default elegido por López de Prado
        model = lgb.LGBMClassifier(
            n_estimators=200,
            num_leaves=15,
            min_child_samples=20,
            learning_rate=0.05,
            random_state=42,
            verbose=-1,
        )
        if sample_weight is not None:
            model.fit(X_train, y_train, sample_weight=sample_weight)
        else:
            model.fit(X_train, y_train)
        prob = model.predict_proba(X_test)[:, 1]
        return prob, model
    # Fallback: GradientBoosting de sklearn
    from sklearn.ensemble import GradientBoostingClassifier
    model = GradientBoostingClassifier(
        n_estimators=150, max_depth=3, learning_rate=0.05, random_state=42,
    )
    if sample_weight is not None:
        model.fit(X_train, y_train, sample_weight=sample_weight)
    else:
        model.fit(X_train, y_train)
    prob = model.predict_proba(X_test)[:, 1]
    return prob, model


# ─────────────────────────────────────────────────────────────────────────────
# Sampling Uniqueness Weights (López de Prado AFML cap. 4)
# ─────────────────────────────────────────────────────────────────────────────
# Cuando dos trades del Triple Barrier solapan en el tiempo (entry_i < exit_j
# y entry_j < exit_i), su informacion esta correlacionada. Tratarlos como
# muestras independientes durante el entrenamiento SOBREESTIMA la calidad
# del modelo y produce overfitting.
#
# Solucion: para cada trade i, calcular su "average uniqueness" = mean a lo
# largo de [entry_i, exit_i] de 1 / num_eventos_concurrentes. Trades muy
# solapados → peso << 1. Trades solitarios → peso = 1.

def compute_concurrent_events(events_df: pd.DataFrame) -> pd.Series:
    """Cuenta cuantos trades estan ACTIVOS en cada timestamp.

    events_df debe tener columnas (entry_date, exit_date), ordenado.
    Construye el indice temporal uniendo todas las fechas de entry/exit
    y, para cada timestamp, suma trades cuyo [entry, exit] lo contiene.

    Devuelve pd.Series indexada por timestamp con el conteo.
    """
    if events_df.empty:
        return pd.Series(dtype=int)
    timestamps = sorted(set(events_df["entry_date"].tolist()) |
                        set(events_df["exit_date"].tolist()))
    timestamps = pd.DatetimeIndex(timestamps)
    count = pd.Series(0, index=timestamps, dtype=int)
    for _, row in events_df.iterrows():
        mask = (count.index >= row["entry_date"]) & (count.index <= row["exit_date"])
        count[mask] = count[mask] + 1
    return count


def compute_uniqueness_per_event(events_df: pd.DataFrame) -> pd.Series:
    """Average uniqueness por trade.

    Para cada trade, calcula mean(1/num_concurrent) durante su lifetime
    [entry_date, exit_date]. Devuelve Series alineada con events_df.index.

    Trades sin solape → uniqueness = 1.0
    Trades con N solapados durante toda su vida → uniqueness ≈ 1/N
    """
    if events_df.empty:
        return pd.Series(dtype=float)
    count = compute_concurrent_events(events_df)
    if count.empty:
        return pd.Series(1.0, index=events_df.index)
    uniqueness = pd.Series(index=events_df.index, dtype=float)
    for idx, row in events_df.iterrows():
        mask = (count.index >= row["entry_date"]) & (count.index <= row["exit_date"])
        if mask.any():
            sub = count[mask]
            # avg uniqueness = mean(1 / num_concurrent)
            uniqueness[idx] = float((1.0 / sub.clip(lower=1)).mean())
        else:
            uniqueness[idx] = 1.0
    return uniqueness.fillna(1.0)


def compute_uniqueness_by_ticker(merged: pd.DataFrame) -> pd.Series:
    """Calcula uniqueness por ticker (los solapes dentro del mismo activo).

    Devuelve Series alineada con merged.index, valores en (0, 1].
    """
    out = pd.Series(1.0, index=merged.index, dtype=float)
    for ticker, group in merged.groupby("ticker"):
        events_df = group[["entry_date", "exit_date"]].copy()
        u = compute_uniqueness_per_event(events_df)
        out.loc[u.index] = u.values
    return out


def _auc(y_true: np.ndarray, prob: np.ndarray) -> float:
    """AUC sin scipy: ranking de probabilidades."""
    if len(np.unique(y_true)) < 2:
        return 0.5
    order = np.argsort(prob)
    y_sorted = y_true[order]
    n_pos = float(y_sorted.sum())
    n_neg = float(len(y_sorted) - n_pos)
    if n_pos == 0 or n_neg == 0:
        return 0.5
    ranks = np.arange(1, len(y_sorted) + 1)
    sum_ranks_pos = float(ranks[y_sorted == 1].sum())
    return float((sum_ranks_pos - n_pos * (n_pos + 1) / 2) / (n_pos * n_neg))


def _build_feature_matrix(merged: pd.DataFrame, ticker_dummies: pd.DataFrame) -> tuple[pd.DataFrame, list[str]]:
    """Construye matriz X de features y devuelve (X, feature_names)."""
    # Numericas
    feat = pd.DataFrame(index=merged.index)
    feat["momentum_20"] = merged["momentum_20"].fillna(0.0)
    feat["rsi_14"] = merged["rsi_14"].fillna(50.0)
    feat["sigma_entry"] = merged["sigma_entry"].fillna(0.0)
    # `side` puede venir como string ("LONG"/"SHORT") o numero (1/-1).
    # Convertimos a numerico para que LightGBM no falle.
    side_raw = merged["side"]
    if side_raw.dtype == object:
        feat["side"] = side_raw.map({"LONG": 1, "SHORT": -1}).fillna(0).astype(int)
    else:
        feat["side"] = side_raw.fillna(0).astype(int)
    # Temporales
    entry = pd.to_datetime(merged["entry_date"])
    feat["day_of_week"] = entry.dt.dayofweek
    feat["month"] = entry.dt.month
    # One-hot regime
    for r in ["TENDENCIA_ALCISTA", "TENDENCIA_BAJISTA", "LATERAL", "CRISIS"]:
        feat[f"regime_{r}"] = (merged["regime"] == r).astype(int)
    # One-hot ticker (limitado para no explotar dimensionalidad)
    feat = pd.concat([feat, ticker_dummies.reindex(merged.index, fill_value=0)], axis=1)
    return feat, list(feat.columns)


class AgenteMeta:
    def __init__(
        self,
        db_path: Path = DB_PATH,
        report_path: Path = REPORT_PATH,
        strategy: str = DEFAULT_STRATEGY,
        n_folds: int = DEFAULT_N_FOLDS,
        thresholds: tuple = DEFAULT_THRESHOLDS,
        use_causal_features: bool = False,
        sample_weighting: str = "uniqueness",   # "none" | "uniqueness"
    ) -> None:
        self.db_path = db_path
        self.report_path = report_path
        self.strategy = strategy
        self.n_folds = n_folds
        self.thresholds = thresholds
        self.use_causal_features = use_causal_features
        self.sample_weighting = sample_weighting
        self.logger = self._build_logger()
        DASHBOARD_DATA_DIR.mkdir(parents=True, exist_ok=True)
        self._ensure_schema()

    def run(self) -> Optional[MetaReport]:
        with sqlite3.connect(self.db_path) as conn:
            # Cargar trades + signals para juntar features de entrada
            trades = pd.read_sql(
                "SELECT ticker, strategy, side, entry_date, exit_date, "
                "entry_price, sigma_entry, return_pct, label, days_held, exit_reason "
                "FROM trades WHERE strategy = ? ORDER BY entry_date ASC, ticker ASC",
                conn, params=(self.strategy,),
                parse_dates=["entry_date", "exit_date"],
            )
            signals = pd.read_sql(
                "SELECT ticker, date, regime, momentum_20, rsi_14 "
                "FROM signals WHERE strategy = ?",
                conn, params=(self.strategy,),
                parse_dates=["date"],
            )
        if trades.empty:
            self.logger.warning("No hay trades para %s", self.strategy)
            return None
        if len(trades) < 100:
            self.logger.warning(
                "Solo %d trades — Meta-Labeling necesita al menos 100",
                len(trades),
            )
            return None

        # Join trades con signals por (ticker, entry_date)
        signals = signals.rename(columns={"date": "entry_date"})
        merged = trades.merge(signals, on=["ticker", "entry_date"], how="left")
        # Target binario: 1 si trade ganador, 0 si no
        merged["y"] = (merged["return_pct"] > 0).astype(int)
        baseline_win_rate = float(merged["y"].mean())

        # One-hot por ticker (mantiene el modelo aware del activo)
        ticker_dummies = pd.get_dummies(merged["ticker"], prefix="tk").astype(int)
        # Aseguramos índice contiguo
        merged = merged.reset_index(drop=True)
        ticker_dummies = ticker_dummies.reset_index(drop=True)
        X_full, feature_names_full = _build_feature_matrix(merged, ticker_dummies)

        # ── Modo causal: filtrar features segun causal_report.json ─────────
        causal_features, spurious_features = load_causal_features()
        causal_mode = "all_features"
        features_filter_note = "Using all available features."

        if self.use_causal_features:
            if not causal_features:
                # No hay reporte causal todavia: avisar y usar todas
                causal_mode = "all_features_fallback"
                features_filter_note = (
                    "Causal mode requested but causal_report.json has no "
                    "causal features. Run agente_causal.py first. "
                    "Fallback: using all features."
                )
                self.logger.warning(features_filter_note)
                X = X_full
                feature_names = feature_names_full
            else:
                # Las causal_features pueden ser variables 'crudas' (momentum_20,
                # rsi_14, sigma_entry, days_held) y 'derivadas' (regime_alcista,
                # regime_lateral). Aceptamos ambas tal cual.
                # Los ticker dummies (tk_AAPL, etc.) se mantienen SIEMPRE: nos dan
                # contexto del activo sin ser variables continuas que el PC analiza.
                allowed = set(causal_features)
                kept_cols = [c for c in feature_names_full
                             if c in allowed or c.startswith("tk_")]
                if not kept_cols:
                    # Algo raro: nada coincide
                    self.logger.warning(
                        "Causal features %s no coinciden con feature matrix %s. "
                        "Fallback a todas.", causal_features, feature_names_full[:10],
                    )
                    causal_mode = "all_features_fallback"
                    features_filter_note = "Mismatch causal vs matrix — using all."
                    X = X_full
                    feature_names = feature_names_full
                else:
                    X = X_full[kept_cols]
                    feature_names = kept_cols
                    causal_mode = "causal_only"
                    n_tk = sum(1 for c in kept_cols if c.startswith("tk_"))
                    features_filter_note = (
                        f"Causal mode ON: {len(kept_cols)} features used "
                        f"({len(kept_cols) - n_tk} causal + {n_tk} ticker dummies). "
                        f"Excluded as spurious: {spurious_features}"
                    )
                    self.logger.info(features_filter_note)
        else:
            X = X_full
            feature_names = feature_names_full

        y = merged["y"].values
        rets = merged["return_pct"].values

        n = len(merged)
        # ── Sampling uniqueness (López de Prado AFML cap.4) ────────────────
        # Si hay solape temporal de trades dentro del mismo ticker, calculamos
        # avg uniqueness por trade y lo usamos como sample_weight del Meta.
        uniqueness_full = None
        avg_uniqueness = 1.0
        n_overlapped = 0
        if self.sample_weighting == "uniqueness":
            try:
                uniqueness_full = compute_uniqueness_by_ticker(merged).values
                avg_uniqueness = float(np.mean(uniqueness_full))
                # Cuenta de trades con uniqueness notablemente bajo
                n_overlapped = int(np.sum(uniqueness_full < 0.95))
                self.logger.info(
                    "Uniqueness weights: avg=%.3f | %d/%d trades con solape (<0.95)",
                    avg_uniqueness, n_overlapped, n,
                )
            except Exception as exc:
                self.logger.warning(
                    "compute_uniqueness fallo: %s — fallback a sample_weight=None", exc,
                )
                uniqueness_full = None
                self.sample_weighting = "none"

        self.logger.info(
            "Trades=%d | features=%d (modo=%s) | baseline win rate=%.1f%% | weighting=%s",
            n, len(feature_names), causal_mode, baseline_win_rate * 100,
            self.sample_weighting,
        )

        # Walk-Forward: predicciones out-of-fold (OOF)
        fold_size = n // (self.n_folds + 1)
        # ── Primary CV pass ─────────────────────────────────────────────────
        primary_result = self._run_cv_pass(
            X, feature_names, y, baseline_win_rate, uniqueness_full,
            fold_size, n, label=f"primary[{causal_mode}]",
        )
        if primary_result is None:
            return None
        auc, acc, oof_prob, feature_importance_acc, valid = primary_result
        y_valid = y[valid]
        prob_valid = oof_prob[valid]
        rets_valid = rets[valid]
        self.logger.info("OOF AUC=%.3f | accuracy=%.3f | baseline=%.3f",
                         auc, acc, baseline_win_rate)

        # ── Comparison CV pass (modo opuesto, mismos folds) ────────────────
        # Asi el usuario ve honestamente si el filtro causal-only le esta
        # costando edge predictivo. delta_auc = primary - comparison.
        auc_comparison: Optional[float] = None
        comparison_mode = "n/a"
        comparison_n_feat = 0
        comparison_note = ""
        if causal_mode == "causal_only":
            # Primary = causal_only ⇒ comparison = extended (X_full)
            cmp_res = self._run_cv_pass(
                X_full, feature_names_full, y, baseline_win_rate,
                uniqueness_full, fold_size, n, label="comparison[extended_all]",
            )
            if cmp_res is not None:
                auc_comparison = round(cmp_res[0], 4)
                comparison_mode = "extended_all"
                comparison_n_feat = len(feature_names_full)
                comparison_note = (
                    f"Comparison with all features ({comparison_n_feat}): "
                    f"AUC={auc_comparison:.3f}. Si delta_auc < 0 ⇒ el filtro "
                    f"causal-only te esta costando edge predictivo."
                )
        elif causal_features:
            # Primary = extended ⇒ comparison = causal_only (si hay info)
            allowed = set(causal_features)
            kept_cols = [c for c in feature_names_full
                         if c in allowed or c.startswith("tk_")]
            if kept_cols:
                X_c = X_full[kept_cols]
                cmp_res = self._run_cv_pass(
                    X_c, kept_cols, y, baseline_win_rate, uniqueness_full,
                    fold_size, n, label="comparison[causal_only]",
                )
                if cmp_res is not None:
                    auc_comparison = round(cmp_res[0], 4)
                    comparison_mode = "causal_only"
                    n_tk = sum(1 for c in kept_cols if c.startswith("tk_"))
                    comparison_n_feat = len(kept_cols)
                    comparison_note = (
                        f"Comparison restringiendo a features causales "
                        f"({len(kept_cols) - n_tk} causales + {n_tk} ticker dummies): "
                        f"AUC={auc_comparison:.3f}. Si delta_auc > 0 ⇒ las "
                        f"features 'spurious' segun PC SI aportan edge predictivo."
                    )
        else:
            comparison_note = ("No hay causal_report.json — no se puede correr el "
                               "pase de comparison.")
        if auc_comparison is not None:
            delta = round(auc - auc_comparison, 4)
            self.logger.info(
                "Comparison pass [%s] AUC=%.3f | delta vs primary=%+.3f",
                comparison_mode, auc_comparison, delta,
            )
        else:
            delta = None

        # Metricas por umbral
        metrics: list[ThresholdMetrics] = []
        for thr in self.thresholds:
            keep_mask = prob_valid >= thr
            n_keep = int(keep_mask.sum())
            n_skip = int(valid.sum() - n_keep)
            keep_rate = float(n_keep / valid.sum())
            if n_keep == 0:
                metrics.append(ThresholdMetrics(
                    threshold=float(thr), n_trades_keep=0, n_trades_skip=n_skip,
                    keep_rate=keep_rate, win_rate_keep=0.0,
                    avg_return_keep=0.0, total_return_keep=0.0,
                    profit_factor_keep=0.0, expectancy_keep=0.0,
                ))
                continue
            rets_keep = rets_valid[keep_mask]
            y_keep = y_valid[keep_mask]
            wins = rets_keep[rets_keep > 0]
            losses = rets_keep[rets_keep < 0]
            pf = (
                float(wins.sum() / abs(losses.sum()))
                if losses.size and losses.sum() != 0 else 999.0
            )
            wr = float(y_keep.mean())
            expectancy = float(rets_keep.mean())  # %
            metrics.append(ThresholdMetrics(
                threshold=float(thr),
                n_trades_keep=n_keep,
                n_trades_skip=n_skip,
                keep_rate=round(keep_rate, 4),
                win_rate_keep=round(wr, 4),
                avg_return_keep=round(expectancy, 4),
                total_return_keep=round(float(rets_keep.sum()), 4),
                profit_factor_keep=round(min(pf, 999.0), 4),
                expectancy_keep=round(expectancy, 4),
            ))

        # Persistir predicciones por trade en SQLite (para el validador filtered)
        self._save_predictions(merged, oof_prob)

        # Top features
        if feature_importance_acc:
            top = sorted(feature_importance_acc.items(), key=lambda kv: kv[1], reverse=True)[:10]
            top_imp = [(k, round(v, 2)) for k, v in top]
        else:
            top_imp = []

        report = MetaReport(
            generated_at=datetime.now(UTC).isoformat(),
            strategy=self.strategy + ("+TB+MetaCausal" if causal_mode == "causal_only"
                                       else "+TB+Meta"),
            model="LightGBM" if _try_lightgbm() else "GradientBoosting",
            n_trades_total=n,
            n_features=len(feature_names),
            feature_names=feature_names,
            auc_oof=round(auc, 4),
            accuracy_oof=round(acc, 4),
            baseline_win_rate=round(baseline_win_rate, 4),
            metrics_by_threshold=metrics,
            n_folds=self.n_folds,
            feature_importance=top_imp,
            causal_mode=causal_mode,
            causal_features_available=causal_features,
            spurious_features_excluded=spurious_features,
            features_filter_note=features_filter_note,
            sample_weighting=self.sample_weighting,
            avg_uniqueness=round(avg_uniqueness, 4),
            n_trades_overlapped=int(n_overlapped),
            auc_oof_primary=round(auc, 4),
            auc_oof_comparison=auc_comparison,
            comparison_mode=comparison_mode,
            comparison_n_features=comparison_n_feat,
            delta_auc=delta,
            comparison_note=comparison_note,
        )
        self._export(report)
        return report

    def _run_cv_pass(self, X_df, feature_names, y, baseline_win_rate,
                     uniqueness_full, fold_size, n, label):
        """Single CV pass sobre las features dadas. Devuelve
        (auc, acc, oof_prob, feat_imp, valid_mask) o None si <50 OOF preds.
        Se usa para los dos pases (primary + comparison) con los mismos folds.
        """
        oof_prob = np.full(n, np.nan)
        feat_imp: dict[str, float] = {}
        for k in range(self.n_folds):
            train_end = (k + 1) * fold_size
            test_start = train_end
            test_end = (k + 2) * fold_size if k < self.n_folds - 1 else n
            if test_end <= test_start or train_end < 50:
                continue
            X_tr = X_df.iloc[:train_end].values
            y_tr = y[:train_end]
            X_te = X_df.iloc[test_start:test_end].values
            sw_tr = uniqueness_full[:train_end] if uniqueness_full is not None else None
            if len(np.unique(y_tr)) < 2:
                oof_prob[test_start:test_end] = baseline_win_rate
                continue
            prob, model = _fit_predict(X_tr, y_tr, X_te, sample_weight=sw_tr)
            oof_prob[test_start:test_end] = prob
            importances = getattr(model, "feature_importances_", None)
            if importances is not None:
                for name, imp in zip(feature_names, importances):
                    feat_imp[name] = feat_imp.get(name, 0.0) + float(imp)
            self.logger.info(
                "[%s] Fold %d: train=[0:%d] test=[%d:%d] | n_test=%d",
                label, k + 1, train_end, test_start, test_end, test_end - test_start,
            )
        valid = ~np.isnan(oof_prob)
        if valid.sum() < 50:
            self.logger.warning(
                "[%s] OOF demasiado pequeño (%d) — abortando", label, valid.sum())
            return None
        auc = _auc(y[valid], oof_prob[valid])
        pred_at_05 = (oof_prob[valid] >= 0.5).astype(int)
        acc = float((pred_at_05 == y[valid]).mean())
        return auc, acc, oof_prob, feat_imp, valid

    def _save_predictions(self, merged: pd.DataFrame, oof_prob: np.ndarray) -> None:
        with sqlite3.connect(self.db_path) as conn:
            conn.execute(
                "DELETE FROM meta_predictions WHERE strategy = ?",
                (self.strategy,),
            )
            rows = [
                (str(r.ticker), str(r.strategy), str(pd.to_datetime(r.entry_date).strftime("%Y-%m-%d")),
                 float(p) if not np.isnan(p) else None)
                for r, p in zip(merged.itertuples(index=False), oof_prob)
            ]
            conn.executemany(
                "INSERT OR REPLACE INTO meta_predictions "
                "(ticker, strategy, entry_date, prob_win) VALUES (?, ?, ?, ?)",
                rows,
            )
            conn.commit()

    def _export(self, report: MetaReport) -> None:
        payload = asdict(report)
        payload["metrics_by_threshold"] = [asdict(m) for m in report.metrics_by_threshold]
        self.report_path.write_text(
            json.dumps(payload, indent=2, ensure_ascii=False),
            encoding="utf-8",
        )
        self.logger.info("Informe exportado → %s", self.report_path)

    def _ensure_schema(self) -> None:
        with sqlite3.connect(self.db_path) as conn:
            conn.execute(
                """
                CREATE TABLE IF NOT EXISTS meta_predictions (
                    ticker TEXT NOT NULL,
                    strategy TEXT NOT NULL,
                    entry_date TEXT NOT NULL,
                    prob_win REAL,
                    PRIMARY KEY (ticker, strategy, entry_date)
                )
                """
            )
            conn.commit()

    def _build_logger(self) -> logging.Logger:
        LOGS_DIR.mkdir(parents=True, exist_ok=True)
        logger = logging.getLogger("agente_meta")
        logger.setLevel(logging.INFO)
        logger.handlers.clear()
        fmt = logging.Formatter("%(asctime)s | %(levelname)s | %(name)s | %(message)s")
        fh = logging.FileHandler(LOGS_DIR / "agente_meta.log", encoding="utf-8")
        fh.setFormatter(fmt)
        sh = logging.StreamHandler()
        sh.setFormatter(fmt)
        logger.addHandler(fh)
        logger.addHandler(sh)
        return logger


def main() -> None:
    import sys
    parser = argparse.ArgumentParser(description="Agente Meta-Labeling de Leonex")
    parser.add_argument("--strategy", type=str, default=DEFAULT_STRATEGY)
    parser.add_argument("--n-folds", type=int, default=DEFAULT_N_FOLDS)
    parser.add_argument("--use-causal-features", action="store_true",
                        help="Entrenar solo con features marcadas como causales por agente_causal.py")
    parser.add_argument("--sample-weighting", type=str, default="uniqueness",
                        choices=["none", "uniqueness"],
                        help="Pesado de muestras durante el entrenamiento. 'uniqueness' aplica AFML cap.4 "
                             "(corrige overlapping events del Triple Barrier). Default uniqueness.")
    parser.add_argument(
        "--thresholds", type=str, default=",".join(map(str, DEFAULT_THRESHOLDS)),
        help="Umbrales de probabilidad separados por comas, p.ej. 0.50,0.55,0.60",
    )
    args = parser.parse_args()
    thr_list = tuple(float(t) for t in args.thresholds.split(","))

    if sys.stdout.encoding and sys.stdout.encoding.lower() != "utf-8":
        import io
        sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="replace")

    agente = AgenteMeta(
        strategy=args.strategy, n_folds=args.n_folds,
        thresholds=thr_list, use_causal_features=args.use_causal_features,
        sample_weighting=args.sample_weighting,
    )
    report = agente.run()

    sep = "=" * 64
    print(f"\n{sep}")
    print("Leonex -- Meta-Labeling")
    print(sep)
    if report is None:
        print("Sin reporte (datos insuficientes).")
        return
    print(f"Estrategia       : {report.strategy}")
    print(f"Modelo           : {report.model}")
    print(f"Trades totales   : {report.n_trades_total}")
    print(f"Features         : {report.n_features}")
    print(f"AUC OOF          : {report.auc_oof:.3f}  (0.5=azar, 0.7+=util)")
    print(f"Accuracy OOF     : {report.accuracy_oof:.3f}")
    print(f"Baseline winrate : {report.baseline_win_rate:.3f}  (sin filtrar)")
    print(f"Folds            : {report.n_folds}")
    print("-" * 64)
    print(f"{'umbral':>7} | {'keep%':>6} | {'n_keep':>7} | {'winrate':>8} | {'PF':>5} | {'exp%':>6}")
    print("-" * 64)
    for m in report.metrics_by_threshold:
        print(f"{m.threshold:7.2f} | {m.keep_rate*100:6.1f} | {m.n_trades_keep:7d} | "
              f"{m.win_rate_keep*100:7.1f}% | {m.profit_factor_keep:5.2f} | "
              f"{m.expectancy_keep:+6.3f}")
    if report.feature_importance:
        print("-" * 64)
        print("Top features:")
        for name, imp in report.feature_importance[:8]:
            print(f"  {name:<25} {imp:>10.1f}")
    print(sep)


if __name__ == "__main__":
    # entrypoint
    main()
