"""
Agente de Validación de Leonex — Anti-overfitting.

Implementa tres capas de validación basadas en López de Prado
(Advances in Financial Machine Learning):

1. Walk-Forward con Purge & Embargo
   - Divide la serie en ventanas rodantes de entrenamiento/test
   - Elimina (purge) las filas de entrenamiento solapadas con el test
   - Añade un embargo de N días tras cada test para evitar leakage

2. Deflated Sharpe Ratio (DSR)
   - Penaliza el Sharpe observado según el número de estrategias
     probadas (trials) y las propiedades estadísticas de la curva
   - DSR < 0.95 → la estrategia no sobrevive al p-hacking

3. Monte Carlo Permutation Test
   - Baraja los retornos diarios 5000 veces
   - Calcula el percentil del Sharpe real respecto a los permutados
   - p-value < 0.05 → el resultado es estadísticamente robusto

Lee señales de SQLite (tabla 'signals' si existe, o genera una
señal momentum de demostración sobre los precios reales) y exporta
el informe a dashboard/data/validation_report.json.
"""

from __future__ import annotations

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 / "validation_report.json"

# ── Configuración ────────────────────────────────────────────────────────────
TRAIN_DAYS = 180          # ~9 meses de entrenamiento por fold
TEST_DAYS = 60            # ~3 meses de test por fold
EMBARGO_DAYS = 5          # días de buffer entre train y test
N_FOLDS = 3               # número de folds walk-forward
N_MC_PERMUTATIONS = 5_000 # permutaciones Monte Carlo
RISK_FREE_DAILY = 0.0     # tasa libre de riesgo diaria (simplificado)
TARGET_DSR = 0.95         # umbral mínimo de DSR para aprobar
TARGET_PVALUE = 0.05      # umbral máximo de p-value MC
# n_trials=1 desactiva la deflacion (e_max=0) -> el DSR satura a 1.000 con
# muestras grandes. El DSR EXISTE para corregir el multiple testing, asi que el
# default debe ser > 1. 10 = numero razonable de variantes probadas del carrier.
DEFAULT_N_TRIALS = 10


# ── Dataclasses de resultados ─────────────────────────────────────────────────
@dataclass
class FoldResult:
    fold: int
    train_start: str
    train_end: str
    test_start: str
    test_end: str
    n_train: int
    n_test: int
    sharpe_test: float
    total_return_test: float
    max_drawdown_test: float
    n_trades: int


@dataclass
class ValidationReport:
    ticker: str
    strategy: str
    generated_at: str
    # Walk-Forward
    folds: list[FoldResult] = field(default_factory=list)
    mean_sharpe_wf: float = 0.0
    mean_return_wf: float = 0.0
    mean_drawdown_wf: float = 0.0
    consistency_ratio: float = 0.0   # % de folds con Sharpe > 0
    # DSR
    sharpe_is: float = 0.0           # Sharpe in-sample completo
    dsr: float = 0.0
    dsr_passes: bool = False
    n_trials: int = 1
    # Monte Carlo
    mc_pvalue: float = 1.0
    mc_passes: bool = False
    mc_sharpe_percentile: float = 0.0
    # Veredicto final
    verdict: str = "RECHAZADO"
    verdict_detail: str = ""


# ── Utilidades estadísticas ───────────────────────────────────────────────────

def sharpe_ratio(returns: np.ndarray, rfr: float = RISK_FREE_DAILY) -> float:
    """Sharpe anualizado (√252)."""
    excess = returns - rfr
    std = excess.std(ddof=1)
    if std == 0 or np.isnan(std):
        return 0.0
    return float(excess.mean() / std * math.sqrt(252))


def max_drawdown(returns: np.ndarray) -> float:
    """Max drawdown sobre la curva de equity (valor negativo)."""
    equity = np.cumprod(1 + returns)
    peak = np.maximum.accumulate(equity)
    dd = (equity - peak) / peak
    return float(dd.min())


def total_return(returns: np.ndarray) -> float:
    return float(np.prod(1 + returns) - 1)


def deflated_sharpe_ratio(
    sharpe_obs: float,
    n_returns: int,
    n_trials: int,
    skewness: float = 0.0,
    kurtosis: float = 3.0,
) -> float:
    """
    DSR según López de Prado (2018).
    Calcula la probabilidad de que el Sharpe observado sea genuino
    dado que se probaron N estrategias.
    """
    from scipy.stats import norm

    if n_returns < 5 or sharpe_obs <= 0:
        return 0.0

    # Sharpe máximo esperado bajo H0 (ruido puro)
    gamma = 0.5772156649  # constante de Euler-Mascheroni
    e_max = (
        (1 - gamma) * norm.ppf(1 - 1 / n_trials)
        + gamma * norm.ppf(1 - 1 / (n_trials * math.e))
        if n_trials > 1
        else 0.0
    )

    # Varianza del estimador del Sharpe
    sr_std = math.sqrt(
        (1 + (0.5 * sharpe_obs**2) - skewness * sharpe_obs
         + ((kurtosis - 3) / 4) * sharpe_obs**2)
        / (n_returns - 1)
    )

    if sr_std == 0:
        return 0.0

    z = (sharpe_obs - e_max) / sr_std
    return float(norm.cdf(z))


def monte_carlo_pvalue(
    returns: np.ndarray,
    observed_sharpe: float,
    n_permutations: int = N_MC_PERMUTATIONS,
    seed: int = 42,
) -> tuple[float, float]:
    """
    Permutation test: baraja los retornos N veces y calcula qué
    fracción de los Sharpes permutados supera el observado.
    Devuelve (p_value, percentile).
    """
    rng = np.random.default_rng(seed)
    perm_sharpes = np.empty(n_permutations)
    r = returns.copy()
    for i in range(n_permutations):
        rng.shuffle(r)
        perm_sharpes[i] = sharpe_ratio(r)

    p_value = float((perm_sharpes >= observed_sharpe).mean())
    percentile = float(100 * (perm_sharpes < observed_sharpe).mean())
    return p_value, percentile


def bootstrap_pvalue_under_h0(  # ← Bootstrap H0 para validar trades
    returns: np.ndarray,
    observed_sharpe: float,
    annual_factor: float = 252.0,
    n_bootstrap: int = N_MC_PERMUTATIONS,
    seed: int = 42,
) -> tuple[float, float]:
    """
    Bootstrap del null para retornos por trade.

    El permutation test simple sobre Sharpe NO funciona: permutar el orden
    de los retornos no cambia ni media ni desviacion tipica, asi que el
    Sharpe permutado coincide con el observado. Por eso pasamos al bootstrap
    bajo la hipotesis nula H0: μ = 0 (no hay edge en la media):

    1. Centramos los retornos restandoles su media → distribucion con μ=0
    2. Generamos n_bootstrap muestras con reemplazo del mismo tamaño
    3. Calculamos Sharpe anualizado de cada muestra
    4. p_value = % de Sharpes nulos que superan el observado

    Esto SI detecta si la media es estadisticamente > 0.
    """
    n = len(returns)
    if n < 5:
        return 1.0, 0.0
    centered = returns - returns.mean()
    rng = np.random.default_rng(seed)
    boot_sharpes = np.empty(n_bootstrap)
    for i in range(n_bootstrap):
        sample = rng.choice(centered, size=n, replace=True)
        std_s = sample.std(ddof=1)
        boot_sharpes[i] = (
            sample.mean() / std_s * math.sqrt(annual_factor)
            if std_s > 0 else 0.0
        )
    p_value = float((boot_sharpes >= observed_sharpe).mean())
    percentile = float(100 * (boot_sharpes < observed_sharpe).mean())
    return p_value, percentile


# ── Señal de demostración (momentum simple) ───────────────────────────────────

def build_momentum_signal(prices: pd.DataFrame, lookback: int = 20) -> pd.Series:
    """
    Señal momentum: long si el retorno de los últimos `lookback` días > 0.
    Retorna una serie de posiciones (+1 / -1 / 0) por fecha.
    """
    mom = prices["close"].pct_change(lookback)
    signal = mom.apply(lambda x: 1 if x > 0 else (-1 if x < 0 else 0))
    return signal.shift(1)  # evitar lookahead: señal de hoy → posición de mañana


def compute_strategy_returns(
    prices: pd.DataFrame,
    signal: pd.Series,
) -> pd.Series:
    """Aplica la señal sobre los retornos diarios."""
    daily_ret = prices["close"].pct_change()
    strat_ret = (daily_ret * signal).dropna()
    return strat_ret


# ── Walk-Forward con Purge & Embargo ─────────────────────────────────────────

def walk_forward_validation(
    prices: pd.DataFrame,
    signal: pd.Series,
    n_folds: int = N_FOLDS,
    train_days: int = TRAIN_DAYS,
    test_days: int = TEST_DAYS,
    embargo_days: int = EMBARGO_DAYS,
) -> list[FoldResult]:
    """
    Walk-Forward con purge y embargo.
    - Purge: elimina del train las últimas `embargo_days` filas antes del test
    - Embargo: salta `embargo_days` días entre el final del test y el inicio del próximo train
    """
    results: list[FoldResult] = []
    dates = prices.index
    n = len(dates)
    required = train_days + embargo_days + test_days

    if n < required:
        return results

    # Punto de inicio: dejamos espacio para todos los folds
    fold_size = test_days + embargo_days
    start_offset = n - (n_folds * fold_size + train_days)
    if start_offset < 0:
        start_offset = 0

    for fold_idx in range(n_folds):
        train_start_i = start_offset + fold_idx * fold_size
        train_end_i = train_start_i + train_days - embargo_days - 1  # purge
        test_start_i = train_start_i + train_days
        test_end_i = min(test_start_i + test_days - 1, n - 1)

        if test_end_i >= n or train_end_i < train_start_i:
            continue

        # Índices de fechas
        train_idx = dates[train_start_i: train_end_i + 1]
        test_idx = dates[test_start_i: test_end_i + 1]

        # Retornos del test usando la señal entrenada en el train
        test_signal = signal.reindex(test_idx)
        test_prices = prices.reindex(test_idx)
        test_returns = compute_strategy_returns(test_prices, test_signal).dropna()

        if len(test_returns) < 10:
            continue

        r = test_returns.values
        n_trades = int((signal.reindex(test_idx).diff().fillna(0) != 0).sum())

        fold = FoldResult(
            fold=fold_idx + 1,
            train_start=str(train_idx[0].date()),
            train_end=str(train_idx[-1].date()),
            test_start=str(test_idx[0].date()),
            test_end=str(test_idx[-1].date()),
            n_train=len(train_idx),
            n_test=len(test_idx),
            sharpe_test=round(sharpe_ratio(r), 4),
            total_return_test=round(total_return(r), 4),
            max_drawdown_test=round(max_drawdown(r), 4),
            n_trades=n_trades,
        )
        results.append(fold)

    return results


# ── Agente principal ───────────────────────────────────────────────────────────

class AgenteValidacion:
    def __init__(
        self,
        db_path: Path = DB_PATH,
        report_path: Path = REPORT_PATH,
        n_trials: int = 1,
        strategy: Optional[str] = None,
        triple_barrier: bool = False,
        meta_filter: Optional[float] = None,
    ) -> None:
        """
        strategy:
            None  → genera momentum_20d internamente (modo legacy)
            "regime_adaptive_v1_lo" o "regime_adaptive_v1_ls" → lee la
            tabla signals producida por agente_senales.py
        triple_barrier:
            True → lee la tabla `trades` (retornos por trade) producida por
            agente_triple_barrier.py. Cambia el universo del validador: en vez
            de retornos diarios sobre una señal continua, opera sobre los
            retornos discretos de cada trade. Esto reduce el ruido y eleva
            la calidad estadística de DSR y Monte Carlo.
        """
        self.db_path = db_path
        self.report_path = report_path
        self.n_trials = n_trials
        self.strategy_override = strategy
        self.triple_barrier = triple_barrier
        self.meta_filter = meta_filter  # None = no filtrar; float = filtrar prob_win >= filtro
        self.logger = self._build_logger()
        DASHBOARD_DATA_DIR.mkdir(parents=True, exist_ok=True)

    def run(self, tickers: Optional[list[str]] = None) -> list[ValidationReport]:
        mode = "trades (Triple Barrier)" if self.triple_barrier else "diario"
        self.logger.info("Iniciando Agente de Validación — modo %s", mode)
        reports: list[ValidationReport] = []

        with sqlite3.connect(self.db_path) as conn:
            if tickers is None:
                table = "trades" if self.triple_barrier else "prices"
                where = " WHERE strategy = ?" if self.triple_barrier else ""
                sql = f"SELECT DISTINCT ticker FROM {table}{where}"
                params = (self.strategy_override,) if self.triple_barrier else ()
                tickers_df = pd.read_sql(sql, conn, params=params)
                tickers = tickers_df["ticker"].tolist()

            # En modo Triple Barrier añadimos primero el report POOLED (de cartera).
            if self.triple_barrier:
                pooled = self._validate_pooled_trades(conn)
                if pooled is not None:
                    reports.append(pooled)
                    self.logger.info(
                        "POOLED → veredicto=%s | DSR=%.3f | p-value=%.3f | WF Sharpe=%.2f",
                        pooled.verdict, pooled.dsr, pooled.mc_pvalue, pooled.mean_sharpe_wf,
                    )

            for ticker in tickers:
                if self.triple_barrier:
                    report = self._validate_ticker_trades(ticker, conn)
                    if report is None:
                        continue
                    reports.append(report)
                    self.logger.info(
                        "%s → veredicto=%s | DSR=%.3f | p-value=%.3f | trades=%d",
                        ticker, report.verdict, report.dsr, report.mc_pvalue,
                        len(report.folds) and sum(f.n_trades for f in report.folds),
                    )
                    continue

                self.logger.info("Validando %s", ticker)
                prices = pd.read_sql(
                    "SELECT date, open, high, low, close, volume "
                    "FROM prices WHERE ticker = ? ORDER BY date ASC",
                    conn,
                    params=(ticker,),
                    index_col="date",
                    parse_dates=["date"],
                )

                if len(prices) < N_FOLDS * (TRAIN_DAYS + TEST_DAYS + EMBARGO_DAYS) + 30:
                    self.logger.warning(
                        "Datos insuficientes para %s (%d filas)", ticker, len(prices)
                    )
                    continue

                report = self._validate_ticker(ticker, prices)
                reports.append(report)
                self.logger.info(
                    "%s → veredicto=%s | DSR=%.3f | p-value=%.3f | WF Sharpe=%.2f",
                    ticker,
                    report.verdict,
                    report.dsr,
                    report.mc_pvalue,
                    report.mean_sharpe_wf,
                )

        self._export(reports)
        self.logger.info(
            "Validación completada. Activos: %d — Aprobados: %d",
            len(reports),
            sum(1 for r in reports if r.verdict == "APROBADO"),
        )
        return reports

    def _validate_ticker(self, ticker: str, prices: pd.DataFrame) -> ValidationReport:
        if self.strategy_override:
            strategy = self.strategy_override
            signal = self._load_signal_from_db(ticker, prices.index, strategy)
            if signal is None:
                self.logger.warning(
                    "%s: sin señales en tabla signals para estrategia=%s; saltando",
                    ticker,
                    strategy,
                )
                # Devolvemos un report vacío con veredicto SIN_DATOS para no romper el flujo.
                return ValidationReport(
                    ticker=ticker,
                    strategy=strategy,
                    generated_at=datetime.now(UTC).isoformat(),
                    verdict="SIN_DATOS",
                    verdict_detail="No signals recorded for this strategy.",
                )
        else:
            strategy = "momentum_20d"
            signal = build_momentum_signal(prices, lookback=20)
        all_returns = compute_strategy_returns(prices, signal).dropna()

        # ── Walk-Forward ──────────────────────────────────────────────────
        folds = walk_forward_validation(prices, signal)

        if folds:
            sharpes = [f.sharpe_test for f in folds]
            returns_ = [f.total_return_test for f in folds]
            dds = [f.max_drawdown_test for f in folds]
            mean_sharpe_wf = float(np.mean(sharpes))
            mean_return_wf = float(np.mean(returns_))
            mean_drawdown_wf = float(np.mean(dds))
            consistency = sum(1 for s in sharpes if s > 0) / len(sharpes)
        else:
            mean_sharpe_wf = 0.0
            mean_return_wf = 0.0
            mean_drawdown_wf = 0.0
            consistency = 0.0

        # ── Deflated Sharpe Ratio ─────────────────────────────────────────
        r_arr = all_returns.values
        sharpe_is = sharpe_ratio(r_arr)
        skew = float(pd.Series(r_arr).skew())
        kurt = float(pd.Series(r_arr).kurtosis() + 3)  # scipy usa exceso
        dsr = deflated_sharpe_ratio(
            sharpe_obs=sharpe_is,
            n_returns=len(r_arr),
            n_trials=self.n_trials,
            skewness=skew,
            kurtosis=kurt,
        )
        dsr_passes = dsr >= TARGET_DSR

        # ── Monte Carlo ───────────────────────────────────────────────────
        mc_pvalue, mc_percentile = monte_carlo_pvalue(r_arr, sharpe_is)
        mc_passes = mc_pvalue < TARGET_PVALUE

        # ── Veredicto ─────────────────────────────────────────────────────
        passed = dsr_passes and mc_passes and mean_sharpe_wf > 0 and consistency >= 0.5
        if passed:
            verdict = "APROBADO"
            verdict_detail = (
                f"Walk-Forward mean Sharpe={mean_sharpe_wf:.2f}, "
                f"DSR={dsr:.3f}, MC p-value={mc_pvalue:.3f}. "
                "The signal survives the three validation layers."
            )
        else:
            fails = []
            if not dsr_passes:
                fails.append(f"DSR={dsr:.3f} < {TARGET_DSR}")
            if not mc_passes:
                fails.append(f"MC p-value={mc_pvalue:.3f} ≥ {TARGET_PVALUE}")
            if mean_sharpe_wf <= 0:
                fails.append(f"WF mean Sharpe={mean_sharpe_wf:.2f} <= 0")
            if consistency < 0.5:
                fails.append(f"Consistency={consistency:.0%} < 50%")
            verdict = "RECHAZADO"
            verdict_detail = "Failures: " + " | ".join(fails)

        return ValidationReport(
            ticker=ticker,
            strategy=strategy,
            generated_at=datetime.now(UTC).isoformat(),
            folds=folds,
            mean_sharpe_wf=round(mean_sharpe_wf, 4),
            mean_return_wf=round(mean_return_wf, 4),
            mean_drawdown_wf=round(mean_drawdown_wf, 4),
            consistency_ratio=round(consistency, 4),
            sharpe_is=round(sharpe_is, 4),
            dsr=round(dsr, 4),
            dsr_passes=dsr_passes,
            n_trials=self.n_trials,
            mc_pvalue=round(mc_pvalue, 4),
            mc_passes=mc_passes,
            mc_sharpe_percentile=round(mc_percentile, 2),
            verdict=verdict,
            verdict_detail=verdict_detail,
        )

    def _export(self, reports: list[ValidationReport]) -> None:
        payload = {
            "generated_at": datetime.now(UTC).isoformat(),
            "config": {
                "train_days": TRAIN_DAYS,
                "test_days": TEST_DAYS,
                "embargo_days": EMBARGO_DAYS,
                "n_folds": N_FOLDS,
                "n_mc_permutations": N_MC_PERMUTATIONS,
                "target_dsr": TARGET_DSR,
                "target_pvalue": TARGET_PVALUE,
            },
            "summary": {
                "total": len(reports),
                "aprobados": sum(1 for r in reports if r.verdict == "APROBADO"),
                "rechazados": sum(1 for r in reports if r.verdict == "RECHAZADO"),
            },
            "reports": [
                {**asdict(r), "folds": [asdict(f) for f in r.folds]}
                for r in reports
            ],
        }
        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)

    # ── Modo Triple Barrier: validacion sobre retornos por trade ─────────────
    def _validate_ticker_trades(
        self, ticker: str, conn: sqlite3.Connection,
    ) -> Optional[ValidationReport]:
        strategy = self.strategy_override or "regime_adaptive_v1_lo"
        if self.meta_filter is not None:
            # JOIN con meta_predictions y filtramos por probabilidad
            trades = pd.read_sql(
                "SELECT t.entry_date, t.exit_date, t.return_pct, t.label, t.days_held, t.exit_reason, "
                "       m.prob_win "
                "FROM trades t LEFT JOIN meta_predictions m "
                "  ON t.ticker = m.ticker AND t.strategy = m.strategy AND t.entry_date = m.entry_date "
                "WHERE t.ticker = ? AND t.strategy = ? AND COALESCE(m.prob_win, 0) >= ? "
                "ORDER BY t.entry_date ASC",
                conn, params=(ticker, strategy, float(self.meta_filter)),
                parse_dates=["entry_date", "exit_date"],
            )
        else:
            trades = pd.read_sql(
                "SELECT entry_date, exit_date, return_pct, label, days_held, exit_reason "
                "FROM trades WHERE ticker = ? AND strategy = ? ORDER BY entry_date ASC",
                conn, params=(ticker, strategy),
                parse_dates=["entry_date", "exit_date"],
            )
        if trades.empty:
            self.logger.warning("Sin trades para %s/%s", ticker, strategy)
            return None

        # Retorno por trade en escala decimal (return_pct estaba en porcentaje)
        rets = trades["return_pct"].values / 100.0
        if len(rets) < 30:
            self.logger.warning(
                "Solo %d trades para %s — saltando validacion",
                len(rets), ticker,
            )
            return None

        # Walk-Forward sobre trades: 3 folds cronologicos
        folds = self._wf_trades(rets, trades, n_folds=3, train_frac=0.6, test_frac=0.2)

        if folds:
            sharpes = [f.sharpe_test for f in folds]
            returns_ = [f.total_return_test for f in folds]
            dds = [f.max_drawdown_test for f in folds]
            mean_sharpe_wf = float(np.mean(sharpes))
            mean_return_wf = float(np.mean(returns_))
            mean_drawdown_wf = float(np.mean(dds))
            consistency = sum(1 for s in sharpes if s > 0) / len(sharpes)
        else:
            mean_sharpe_wf = mean_return_wf = mean_drawdown_wf = consistency = 0.0

        # Sharpe in-sample sobre trades (anualizado por frecuencia real)
        avg_days = float(trades["days_held"].mean())
        trades_per_year = 252.0 / max(avg_days, 1.0)
        std = rets.std(ddof=1)
        if std > 0 and len(rets) > 1:
            sharpe_is = float(rets.mean() / std * math.sqrt(trades_per_year))
        else:
            sharpe_is = 0.0
        skew = float(pd.Series(rets).skew()) if len(rets) > 2 else 0.0
        kurt = float(pd.Series(rets).kurtosis() + 3) if len(rets) > 3 else 3.0
        dsr = deflated_sharpe_ratio(
            sharpe_obs=sharpe_is,
            n_returns=len(rets),
            n_trials=self.n_trials,
            skewness=skew,
            kurtosis=kurt,
        )
        dsr_passes = dsr >= TARGET_DSR

        # Bootstrap bajo H0 sobre retornos por trade:
        # centramos retornos en cero (asumimos μ=0 bajo H0 "no hay edge") y
        # generamos muestras con reemplazo del mismo tamaño que rets. El
        # Sharpe observado se compara contra esa distribución nula. Este test
        # SI detecta edge en la media, a diferencia del permutation simple.
        mc_pvalue, mc_percentile = bootstrap_pvalue_under_h0(
            rets, sharpe_is, trades_per_year,
        )
        mc_passes = mc_pvalue < TARGET_PVALUE

        passed = dsr_passes and mc_passes and mean_sharpe_wf > 0 and consistency >= 0.5
        if passed:
            verdict = "APROBADO"
            verdict_detail = (
                f"Walk-Forward Sharpe={mean_sharpe_wf:.2f}, DSR={dsr:.3f}, "
                f"MC p-value={mc_pvalue:.3f}. The strategy survives the four "
                f"layers over Triple Barrier trades."
            )
        else:
            fails = []
            if not dsr_passes:
                fails.append(f"DSR={dsr:.3f} < {TARGET_DSR}")
            if not mc_passes:
                fails.append(f"MC p-value={mc_pvalue:.3f} ≥ {TARGET_PVALUE}")
            if mean_sharpe_wf <= 0:
                fails.append(f"WF mean Sharpe={mean_sharpe_wf:.2f} <= 0")
            if consistency < 0.5:
                fails.append(f"Consistency={consistency:.0%} < 50%")
            verdict = "RECHAZADO"
            verdict_detail = "Failures: " + " | ".join(fails)

        return ValidationReport(
            ticker=ticker,
            strategy=strategy + "+TB",
            generated_at=datetime.now(UTC).isoformat(),
            folds=folds,
            mean_sharpe_wf=round(mean_sharpe_wf, 4),
            mean_return_wf=round(mean_return_wf, 4),
            mean_drawdown_wf=round(mean_drawdown_wf, 4),
            consistency_ratio=round(consistency, 4),
            sharpe_is=round(sharpe_is, 4),
            dsr=round(dsr, 4),
            dsr_passes=dsr_passes,
            n_trials=self.n_trials,
            mc_pvalue=round(mc_pvalue, 4),
            mc_passes=mc_passes,
            mc_sharpe_percentile=round(mc_percentile, 2),
            verdict=verdict,
            verdict_detail=verdict_detail,
        )

    def _wf_trades(
        self,
        rets: np.ndarray,
        trades: pd.DataFrame,
        n_folds: int = 3,
        train_frac: float = 0.6,
        test_frac: float = 0.2,
    ) -> list[FoldResult]:
        """Walk-forward sobre trades cronologicos. Sin overlapping ni purge
        porque cada trade ya es un evento discreto independiente del siguiente."""
        n = len(rets)
        results: list[FoldResult] = []
        train_size = max(int(n * train_frac), 30)
        test_size = max(int(n * test_frac), 10)
        if train_size + test_size > n:
            return results

        # Folds desplazados: avanzamos en pasos test_size desde el final
        avg_days = float(trades["days_held"].mean())
        trades_per_year = 252.0 / max(avg_days, 1.0)

        for k in range(n_folds):
            test_end = n - k * test_size
            test_start = test_end - test_size
            if test_start <= train_size:
                break
            r = rets[test_start:test_end]
            std = r.std(ddof=1)
            sharpe = float(r.mean() / std * math.sqrt(trades_per_year)) if std > 0 else 0.0
            total = float(np.prod(1 + r) - 1)
            mdd = max_drawdown(r)
            fold = FoldResult(
                fold=n_folds - k,
                train_start=str(trades["entry_date"].iloc[test_start - train_size]),
                train_end=str(trades["entry_date"].iloc[test_start - 1]),
                test_start=str(trades["entry_date"].iloc[test_start]),
                test_end=str(trades["entry_date"].iloc[test_end - 1]),
                n_train=train_size,
                n_test=test_size,
                sharpe_test=round(sharpe, 4),
                total_return_test=round(total, 4),
                max_drawdown_test=round(mdd, 4),
                n_trades=test_size,
            )
            results.append(fold)
        return list(reversed(results))

    # ── Validacion POOLED de cartera: agrega los trades de todos los tickers ─
    def _validate_pooled_trades(self, conn: sqlite3.Connection) -> Optional[ValidationReport]:
        strategy = self.strategy_override or "regime_adaptive_v1_lo"
        if self.meta_filter is not None:
            trades = pd.read_sql(
                "SELECT t.ticker, t.entry_date, t.exit_date, t.return_pct, t.days_held, t.label, t.exit_reason, "
                "       m.prob_win "
                "FROM trades t LEFT JOIN meta_predictions m "
                "  ON t.ticker = m.ticker AND t.strategy = m.strategy AND t.entry_date = m.entry_date "
                "WHERE t.strategy = ? AND COALESCE(m.prob_win, 0) >= ? "
                "ORDER BY t.entry_date ASC, t.ticker ASC",
                conn, params=(strategy, float(self.meta_filter)),
                parse_dates=["entry_date", "exit_date"],
            )
        else:
            trades = pd.read_sql(
                "SELECT ticker, entry_date, exit_date, return_pct, days_held, label, exit_reason "
                "FROM trades WHERE strategy = ? ORDER BY entry_date ASC, ticker ASC",
                conn, params=(strategy,),
                parse_dates=["entry_date", "exit_date"],
            )
        if trades.empty:
            self.logger.warning("Sin trades para validacion pooled de %s", strategy)
            return None

        rets = trades["return_pct"].values / 100.0
        n = len(rets)
        if n < 100:
            self.logger.warning("Solo %d trades pooled — saltando", n)
            return None

        # Walk-Forward sobre el pool cronologico
        # train 60% / test 20% / 3 folds en cascada
        folds = self._wf_trades(rets, trades, n_folds=3, train_frac=0.6, test_frac=0.2)
        if folds:
            sharpes = [f.sharpe_test for f in folds]
            returns_ = [f.total_return_test for f in folds]
            dds = [f.max_drawdown_test for f in folds]
            mean_sharpe_wf = float(np.mean(sharpes))
            mean_return_wf = float(np.mean(returns_))
            mean_drawdown_wf = float(np.mean(dds))
            consistency = sum(1 for s in sharpes if s > 0) / len(sharpes)
        else:
            mean_sharpe_wf = mean_return_wf = mean_drawdown_wf = consistency = 0.0

        # Sharpe in-sample sobre pool de trades
        avg_days = float(trades["days_held"].mean())
        trades_per_year = 252.0 / max(avg_days, 1.0)
        std = rets.std(ddof=1)
        sharpe_is = (
            float(rets.mean() / std * math.sqrt(trades_per_year))
            if std > 0 and n > 1 else 0.0
        )
        skew = float(pd.Series(rets).skew()) if n > 2 else 0.0
        kurt = float(pd.Series(rets).kurtosis() + 3) if n > 3 else 3.0
        # N EFECTIVO: 86k trades solapados entre 510 activos NO son 86k
        # observaciones independientes. Los trades del MISMO dia estan
        # correlacionados transversalmente, asi que colapsan a ~1 observacion.
        # Usar el nº de fechas de entrada distintas evita que la varianza del
        # estimador sea absurdamente pequena y el DSR sature a 1.000.
        try:
            n_eff = int(trades["entry_date"].dt.normalize().nunique())
        except Exception:
            n_eff = n
        n_eff = max(5, min(n, n_eff))
        dsr = deflated_sharpe_ratio(
            sharpe_obs=sharpe_is, n_returns=n_eff, n_trials=self.n_trials,
            skewness=skew, kurtosis=kurt,
        )
        dsr_passes = dsr >= TARGET_DSR

        # Bootstrap bajo H0 (mu=0): el sharpe observado contra la distribucion
        # nula generada por muestreo con reemplazo de retornos centrados.
        mc_pvalue, mc_percentile = bootstrap_pvalue_under_h0(
            rets, sharpe_is, trades_per_year,
        )
        mc_passes = mc_pvalue < TARGET_PVALUE

        passed = dsr_passes and mc_passes and mean_sharpe_wf > 0 and consistency >= 0.5
        if passed:
            verdict = "APROBADO"
            verdict_detail = (
                f"Pooled portfolio passes the 4 filters over {n} trades: "
                f"WF Sharpe={mean_sharpe_wf:.2f}, DSR={dsr:.3f}, MC p={mc_pvalue:.3f}, "
                f"consistency={consistency:.0%}."
            )
        else:
            fails = []
            if not dsr_passes: fails.append(f"DSR={dsr:.3f} < {TARGET_DSR}")
            if not mc_passes: fails.append(f"MC p-value={mc_pvalue:.3f} >= {TARGET_PVALUE}")
            if mean_sharpe_wf <= 0: fails.append(f"WF Sharpe={mean_sharpe_wf:.2f} <= 0")
            if consistency < 0.5: fails.append(f"Consistency={consistency:.0%} < 50%")
            verdict = "RECHAZADO"
            verdict_detail = f"Portfolio ({n} trades) failures: " + " | ".join(fails)

        return ValidationReport(
            ticker="POOLED",
            strategy=strategy + "+TB",
            generated_at=datetime.now(UTC).isoformat(),
            folds=folds,
            mean_sharpe_wf=round(mean_sharpe_wf, 4),
            mean_return_wf=round(mean_return_wf, 4),
            mean_drawdown_wf=round(mean_drawdown_wf, 4),
            consistency_ratio=round(consistency, 4),
            sharpe_is=round(sharpe_is, 4),
            dsr=round(dsr, 4),
            dsr_passes=dsr_passes,
            n_trials=self.n_trials,
            mc_pvalue=round(mc_pvalue, 4),
            mc_passes=mc_passes,
            mc_sharpe_percentile=round(mc_percentile, 2),
            verdict=verdict,
            verdict_detail=verdict_detail,
        )

    def _load_signal_from_db(
        self,
        ticker: str,
        index: pd.DatetimeIndex,
        strategy: str,
    ) -> Optional[pd.Series]:
        """Lee la tabla signals y devuelve una serie de posiciones alineada al índice de precios."""
        with sqlite3.connect(self.db_path) as conn:
            df = pd.read_sql(
                "SELECT date, signal FROM signals "
                "WHERE ticker = ? AND strategy = ? ORDER BY date ASC",
                conn,
                params=(ticker, strategy),
                index_col="date",
                parse_dates=["date"],
            )
        if df.empty:
            return None
        signal = df["signal"].astype(float)
        # Alinear con el índice de prices; los días sin señal → 0
        signal = signal.reindex(index).fillna(0.0)
        return signal

    def _build_logger(self) -> logging.Logger:
        LOGS_DIR.mkdir(parents=True, exist_ok=True)
        logger = logging.getLogger("agente_validacion")
        logger.setLevel(logging.INFO)
        logger.handlers.clear()
        fmt = logging.Formatter("%(asctime)s | %(levelname)s | %(name)s | %(message)s")
        fh = logging.FileHandler(LOGS_DIR / "agente_validacion.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
    import argparse

    parser = argparse.ArgumentParser(description="Agente de Validacion Leonex")
    parser.add_argument(
        "--strategy",
        type=str,
        default=None,
        help="Si se indica, lee senales de la tabla signals (p.ej. regime_adaptive_v1_lo). "
             "Si se omite usa momentum_20d generado internamente.",
    )
    parser.add_argument("--n-trials", type=int, default=DEFAULT_N_TRIALS,
                        help="Numero de estrategias probadas (para corregir el DSR via deflacion). "
                             "1 = sin correccion (NO recomendado: el DSR satura a 1.000).")
    parser.add_argument("--triple-barrier", action="store_true",
                        help="Validar sobre retornos por trade (tabla 'trades') en vez de retornos diarios")
    parser.add_argument("--meta-filter", type=float, default=None,
                        help="Solo validar trades con prob_win >= umbral (requiere haber corrido agente_meta.py)")
    args = parser.parse_args()

    # Fijar stdout a UTF-8 en Windows para evitar UnicodeEncodeError
    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 = AgenteValidacion(
        n_trials=args.n_trials,
        strategy=args.strategy,
        triple_barrier=args.triple_barrier,
        meta_filter=args.meta_filter,
    )
    reports = agente.run()

    sep = "=" * 60
    dash = "-" * 60
    print(f"\n{sep}")
    print("Leonex -- Informe de Validacion Anti-Overfitting")
    print(sep)
    print(f"Estrategia         : {args.strategy or 'momentum_20d (legacy)'}")
    print(f"Activos analizados : {len(reports)}")
    print(f"Aprobados          : {sum(1 for r in reports if r.verdict == 'APROBADO')}")
    print(f"Rechazados         : {sum(1 for r in reports if r.verdict == 'RECHAZADO')}")
    print(f"Sin datos          : {sum(1 for r in reports if r.verdict == 'SIN_DATOS')}")
    print(dash)
    for r in sorted(reports, key=lambda x: x.mean_sharpe_wf, reverse=True):
        badge = {
            "APROBADO": "[OK]",
            "RECHAZADO": "[--]",
            "SIN_DATOS": "[??]",
        }.get(r.verdict, "[--]")
        print(
            f"{badge} {r.ticker:<12} | WF Sharpe={r.mean_sharpe_wf:+.2f} "
            f"| DSR={r.dsr:.3f} | MC p={r.mc_pvalue:.3f} "
            f"| {r.verdict}"
        )
    print(sep)
    print(f"Informe completo -> {REPORT_PATH}")
    # eof


if __name__ == "__main__":
    main()
