"""
Agente Drift Monitor de Leonex.

El riesgo silencioso de cualquier sistema de trading no es perder un trade,
es que la estrategia haya dejado de funcionar hace semanas sin que te enteres.
Este agente compara cada vez que se ejecuta:

    - Sharpe rolling de los ultimos N trades  vs  Sharpe IS de la validacion POOLED
    - Win rate rolling                         vs  win rate historico del Triple Barrier
    - Profit factor rolling                    vs  profit factor historico
    - Drawdown actual de los ultimos trades

Y clasifica el estado del sistema:

    GREEN   →  metricas dentro del rango esperado, sistema sano
    YELLOW  →  metricas degradadas pero todavia en rango aceptable
    RED     →  metricas fuera de rango: activa flag PAUSE_NEW_TRADES

Cuando se activa PAUSE_NEW_TRADES, el executor deja de abrir posiciones
nuevas pero el close_monitor sigue operando hasta liquidar todo lo abierto.
Es el equivalente a "el sistema se autoprotege" sin que tengas que vigilarlo
tu manualmente.

Uso:
    python agents/agente_drift_monitor.py
    python agents/agente_drift_monitor.py --window 60      # ventana de trades para rolling
    python agents/agente_drift_monitor.py --reset-pause    # desactiva la pausa manual
"""

from __future__ import annotations

import argparse
import json
import logging
import math
import sqlite3
import sys
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 / "drift_report.json"
VALIDATION_PATH = DASHBOARD_DATA_DIR / "validation_report.json"
TRADES_PATH = DASHBOARD_DATA_DIR / "trades_report.json"


# ── Umbrales (calibrados conservadoramente) ──────────────────────────────
# Sharpe rolling vs IS:
#   ratio >= 0.50  → GREEN
#   ratio 0.20-0.50 → YELLOW
#   ratio < 0.20    → RED + PAUSE_NEW_TRADES
SHARPE_RATIO_RED = 0.20
SHARPE_RATIO_YELLOW = 0.50

# Win rate absoluto: si cae bajo este umbral, RED inmediato
WIN_RATE_RED_FLOOR = 0.30
WIN_RATE_YELLOW_FLOOR = 0.38

# Profit Factor minimo aceptable
PROFIT_FACTOR_RED = 0.90      # bajo 0.90 = sistema perdedor
PROFIT_FACTOR_YELLOW = 1.05

# Drawdown sobre los trades rolling
DD_RED = 0.20                  # 20% drawdown rolling = RED
DD_YELLOW = 0.10

# ── Lógica STICKY (anti-flicker) ─────────────────────────────────────────
# Cuando el sistema esta PAUSADO por drift, no salimos de la pausa al primer
# GREEN — exigimos N evaluaciones GREEN consecutivas. Esto evita el caso
# tipico: el rolling de 60 trades sufre un par de operaciones buenas y se
# va a +3 Sharpe; al siguiente run vuelve a -2 y re-RED. Mejor estabilidad
# que reactividad. RED y YELLOW siguen siendo inmediatos.
STICKY_GREEN_CONFIRMATIONS = 3

# Numero minimo de trades rolling para hacer evaluacion
MIN_TRADES_FOR_EVAL = 20


@dataclass
class DriftMetric:
    name: str
    value_recent: float
    value_baseline: float
    threshold_yellow: float
    threshold_red: float
    status: str                # "GREEN" | "YELLOW" | "RED"
    note: str = ""


@dataclass
class DriftReport:
    generated_at: str
    window_trades: int
    n_recent_trades: int
    overall_status: str         # "GREEN" | "YELLOW" | "RED" | "NO_DATA"
    pause_new_trades: bool
    metrics: list[DriftMetric] = field(default_factory=list)
    summary_note: str = ""
    recent_events: list[dict] = field(default_factory=list)


def ensure_schema(db_path: Path) -> None:
    with sqlite3.connect(db_path) as conn:
        conn.execute(
            """
            CREATE TABLE IF NOT EXISTS drift_events (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                timestamp TEXT NOT NULL,
                window_trades INTEGER,
                n_recent_trades INTEGER,
                overall_status TEXT NOT NULL,
                pause_new_trades INTEGER NOT NULL,
                sharpe_recent REAL,
                sharpe_baseline REAL,
                win_rate_recent REAL,
                win_rate_baseline REAL,
                profit_factor_recent REAL,
                profit_factor_baseline REAL,
                drawdown_recent REAL,
                summary_note TEXT
            )
            """
        )
        conn.execute(
            """
            CREATE TABLE IF NOT EXISTS system_state (
                key TEXT PRIMARY KEY,
                value TEXT NOT NULL,
                updated_at TEXT NOT NULL
            )
            """
        )
        conn.commit()


def get_system_state(key: str, default: Optional[str] = None,
                     db_path: Path = DB_PATH) -> Optional[str]:
    """Lee un flag global del sistema. Devuelve `default` si no existe."""
    try:
        with sqlite3.connect(db_path) as conn:
            row = conn.execute(
                "SELECT value FROM system_state WHERE key = ?", (key,)
            ).fetchone()
            return row[0] if row else default
    except sqlite3.OperationalError:
        return default


def set_system_state(key: str, value: str, db_path: Path = DB_PATH) -> None:
    """Guarda o actualiza un flag global."""
    ensure_schema(db_path)
    with sqlite3.connect(db_path) as conn:
        conn.execute(
            """
            INSERT INTO system_state (key, value, updated_at)
            VALUES (?, ?, ?)
            ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at
            """,
            (key, str(value), datetime.now(UTC).isoformat()),
        )
        conn.commit()


def get_pause_flag(db_path: Path = DB_PATH) -> bool:
    """Conveniencia: True si el sistema esta pausado para nuevas entradas."""
    v = get_system_state("pause_new_trades", default="false", db_path=db_path)
    return str(v).lower() in ("true", "1", "yes")


def _read_validation_baseline() -> dict:
    """Lee el report POOLED para obtener Sharpe IS de referencia."""
    if not VALIDATION_PATH.exists():
        return {}
    try:
        data = json.loads(VALIDATION_PATH.read_text(encoding="utf-8"))
        for r in data.get("reports", []):
            if r.get("ticker") == "POOLED":
                return r
    except Exception:
        pass
    return {}


def _read_trades_baseline() -> dict:
    """Lee summary del trades_report.json (historico Triple Barrier completo)."""
    if not TRADES_PATH.exists():
        return {}
    try:
        return json.loads(TRADES_PATH.read_text(encoding="utf-8")).get("summary", {})
    except Exception:
        return {}


def _fetch_recent_trades(window: int, db_path: Path) -> pd.DataFrame:
    """Devuelve los `window` trades mas recientes (por exit_date) de la estrategia activa."""
    if not db_path.exists():
        return pd.DataFrame()
    strategy = "regime_adaptive_v1_lo"  # principal por ahora
    with sqlite3.connect(db_path) as conn:
        df = pd.read_sql(
            """
            SELECT ticker, entry_date, exit_date, return_pct, days_held, label, exit_reason
            FROM trades
            WHERE strategy = ?
            ORDER BY exit_date DESC LIMIT ?
            """,
            conn, params=(strategy, int(window)),
            parse_dates=["entry_date", "exit_date"],
        )
    # los queremos en orden cronologico ascendente para metricas como drawdown
    return df.sort_values("exit_date").reset_index(drop=True)


def _annual_sharpe_trades(rets_pct: np.ndarray, avg_days_held: float) -> float:
    if len(rets_pct) < 2:
        return 0.0
    rets = rets_pct / 100.0
    std = float(rets.std(ddof=1))
    if std <= 0:
        return 0.0
    trades_per_year = 252.0 / max(avg_days_held, 1.0)
    return float(rets.mean() / std * math.sqrt(trades_per_year))


def _profit_factor(rets_pct: np.ndarray) -> float:
    rets = rets_pct
    wins = rets[rets > 0].sum()
    losses = -rets[rets < 0].sum()
    if losses <= 0:
        return 999.0 if wins > 0 else 0.0
    return float(wins / losses)


def _drawdown(rets_pct: np.ndarray) -> float:
    """Max drawdown sobre la curva de equity de los retornos rolling."""
    if len(rets_pct) == 0:
        return 0.0
    eq = np.cumprod(1 + rets_pct / 100.0)
    peak = np.maximum.accumulate(eq)
    return float(((eq - peak) / peak).min())


def evaluate_drift(window: int, db_path: Path) -> DriftReport:
    """Genera el report comparando rolling vs baseline."""
    ensure_schema(db_path)
    baseline_val = _read_validation_baseline()
    baseline_trades = _read_trades_baseline()
    df = _fetch_recent_trades(window, db_path)
    now_iso = datetime.now(UTC).isoformat()

    if df.empty:
        return DriftReport(
            generated_at=now_iso, window_trades=window, n_recent_trades=0,
            overall_status="NO_DATA", pause_new_trades=get_pause_flag(db_path),
            summary_note="No trades in the table to evaluate drift.",
        )
    if len(df) < MIN_TRADES_FOR_EVAL:
        return DriftReport(
            generated_at=now_iso, window_trades=window, n_recent_trades=len(df),
            overall_status="NO_DATA", pause_new_trades=get_pause_flag(db_path),
            summary_note=f"Only {len(df)} trades — minimum {MIN_TRADES_FOR_EVAL} to evaluate.",
        )

    rets = df["return_pct"].values
    avg_days = float(df["days_held"].mean()) if "days_held" in df else 3.0
    n = len(rets)

    # Metricas rolling
    sharpe_recent = _annual_sharpe_trades(rets, avg_days)
    win_rate_recent = float((rets > 0).mean())
    pf_recent = _profit_factor(rets)
    dd_recent = _drawdown(rets)

    # Baselines
    sharpe_baseline = float(baseline_val.get("sharpe_is", 0.0))
    win_rate_baseline = float(baseline_trades.get("win_rate", 0.0))
    pf_baseline = float(baseline_trades.get("profit_factor", 1.0))

    # ── Evaluacion por metrica ───────────────────────────────────────────
    metrics: list[DriftMetric] = []

    # Sharpe ratio: comparado como ratio relativo (resiste cambio de escala)
    if sharpe_baseline > 0:
        ratio = sharpe_recent / sharpe_baseline
    else:
        ratio = 0.0 if sharpe_recent <= 0 else 1.0
    if ratio < SHARPE_RATIO_RED or sharpe_recent < 0:
        s_status = "RED"
    elif ratio < SHARPE_RATIO_YELLOW:
        s_status = "YELLOW"
    else:
        s_status = "GREEN"
    metrics.append(DriftMetric(
        name="Sharpe Annual (rolling vs IS)",
        value_recent=round(sharpe_recent, 4),
        value_baseline=round(sharpe_baseline, 4),
        threshold_yellow=SHARPE_RATIO_YELLOW,
        threshold_red=SHARPE_RATIO_RED,
        status=s_status,
        note=f"ratio={ratio:.2f}",
    ))

    # Win rate absoluto
    if win_rate_recent < WIN_RATE_RED_FLOOR:
        w_status = "RED"
    elif win_rate_recent < WIN_RATE_YELLOW_FLOOR:
        w_status = "YELLOW"
    else:
        w_status = "GREEN"
    metrics.append(DriftMetric(
        name="Win Rate (rolling)",
        value_recent=round(win_rate_recent, 4),
        value_baseline=round(win_rate_baseline, 4),
        threshold_yellow=WIN_RATE_YELLOW_FLOOR,
        threshold_red=WIN_RATE_RED_FLOOR,
        status=w_status,
    ))

    # Profit factor absoluto
    if pf_recent < PROFIT_FACTOR_RED:
        p_status = "RED"
    elif pf_recent < PROFIT_FACTOR_YELLOW:
        p_status = "YELLOW"
    else:
        p_status = "GREEN"
    metrics.append(DriftMetric(
        name="Profit Factor (rolling)",
        value_recent=round(pf_recent, 4),
        value_baseline=round(pf_baseline, 4),
        threshold_yellow=PROFIT_FACTOR_YELLOW,
        threshold_red=PROFIT_FACTOR_RED,
        status=p_status,
    ))

    # Drawdown rolling
    if abs(dd_recent) > DD_RED:
        d_status = "RED"
    elif abs(dd_recent) > DD_YELLOW:
        d_status = "YELLOW"
    else:
        d_status = "GREEN"
    metrics.append(DriftMetric(
        name="Drawdown (rolling)",
        value_recent=round(dd_recent, 4),
        value_baseline=0.0,
        threshold_yellow=DD_YELLOW,
        threshold_red=DD_RED,
        status=d_status,
    ))

    # ── Status global y decision (con lógica STICKY anti-flicker) ────────
    statuses = [m.status for m in metrics]
    red_metrics = [m.name for m in metrics if m.status == "RED"]

    # 1) Calcular el status "crudo" de esta evaluacion (sin sticky aun)
    if "RED" in statuses:
        raw_status = "RED"
    elif "YELLOW" in statuses:
        raw_status = "YELLOW"
    else:
        raw_status = "GREEN"

    # 2) Sticky: ¿estamos pausados ahora mismo POR DRIFT? Si si, y el crudo
    # da GREEN, exigimos N evaluaciones GREEN consecutivas antes de salir.
    was_paused = get_pause_flag(db_path)
    prev_reason = get_system_state("pause_reason", "", db_path) or ""
    drift_was_pausing = bool(was_paused) and prev_reason.startswith("drift_monitor:")

    sticky_pending = 0      # cuantos GREEN consecutivos llevamos contando el actual
    sticky_needed = STICKY_GREEN_CONFIRMATIONS

    if raw_status == "RED":
        overall = "RED"
        pause = True
        note = ("System in RED — PAUSE_NEW_TRADES activated. The executor will "
                "not open new positions. The close_monitor keeps managing the "
                "open ones until they are liquidated.")
    elif raw_status == "YELLOW":
        # YELLOW no flipa la pausa — preserva el estado anterior.
        # Asi evitamos un YELLOW intermedio des-pausando o pausando spurious.
        overall = "YELLOW"
        pause = drift_was_pausing      # mantiene la pausa de drift si la habia
        if drift_was_pausing:
            note = ("System in YELLOW — degraded metrics. PAUSE_NEW_TRADES "
                    "se MANTIENE (estaba activo por drift). Espero confirmacion "
                    "GREEN para liberar.")
        else:
            note = ("System in YELLOW — degraded metrics but still operational. "
                    "Monitor.")
    else:  # raw_status == "GREEN"
        if drift_was_pausing:
            # Necesitamos N consecutivos antes de des-pausar.
            # Contamos los GREEN de los ultimos (N-1) eventos del drift, e
            # incluimos el actual como el Nth.
            try:
                with sqlite3.connect(db_path) as _conn:
                    prev_statuses = [
                        r[0] for r in _conn.execute(
                            "SELECT overall_status FROM drift_events "
                            "ORDER BY id DESC LIMIT ?", (sticky_needed - 1,)
                        ).fetchall()
                    ]
            except Exception:
                prev_statuses = []
            consecutive_greens = 1   # el actual cuenta
            for s in prev_statuses:
                if s == "GREEN":
                    consecutive_greens += 1
                else:
                    break
            sticky_pending = consecutive_greens
            if consecutive_greens >= sticky_needed:
                # Suficientes confirmaciones: liberamos
                overall = "GREEN"
                pause = False
                note = (f"System back in GREEN after {consecutive_greens} "
                        f"consecutive GREEN evaluations (sticky exit). "
                        f"PAUSE_NEW_TRADES released.")
            else:
                # Aun no llegamos al umbral — seguimos pausados
                overall = "GREEN"
                pause = True
                note = (f"Metrics now GREEN but PAUSE_NEW_TRADES kept active "
                        f"(sticky exit: {consecutive_greens}/{sticky_needed} "
                        f"consecutive GREEN). Releasing only after "
                        f"{sticky_needed - consecutive_greens} more GREEN run(s).")
        else:
            overall = "GREEN"
            pause = False
            note = "System in GREEN — metrics within range."

    # Aplicar la decision al system_state
    set_system_state("pause_new_trades", "true" if pause else "false", db_path)
    # pause_reason legible: si pausamos, dejamos prefijo y metricas RED;
    # si NO pausamos (recuperacion), limpiamos la razon solo si el motivo
    # previo provenia del drift_monitor (para no pisar pausas de risk_monitor).
    if pause:
        reason = f"drift_monitor: RED ({', '.join(red_metrics) or 'unknown'})"
        set_system_state("pause_reason", reason, db_path)
    else:
        prev = get_system_state("pause_reason", "", db_path) or ""
        if prev.startswith("drift_monitor:"):
            set_system_state("pause_reason", "", db_path)

    # Persistir el evento
    with sqlite3.connect(db_path) as conn:
        conn.execute(
            """
            INSERT INTO drift_events (
                timestamp, window_trades, n_recent_trades, overall_status, pause_new_trades,
                sharpe_recent, sharpe_baseline, win_rate_recent, win_rate_baseline,
                profit_factor_recent, profit_factor_baseline, drawdown_recent, summary_note
            ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
            """,
            (now_iso, window, n, overall, int(pause),
             sharpe_recent, sharpe_baseline, win_rate_recent, win_rate_baseline,
             pf_recent, pf_baseline, dd_recent, note),
        )
        # Recuperar eventos recientes para el report
        recent = conn.execute(
            """
            SELECT timestamp, overall_status, pause_new_trades, n_recent_trades,
                   sharpe_recent, win_rate_recent, profit_factor_recent, drawdown_recent, summary_note
            FROM drift_events ORDER BY id DESC LIMIT 15
            """
        ).fetchall()
        conn.commit()

    cols = ["timestamp", "overall_status", "pause_new_trades", "n_recent_trades",
            "sharpe_recent", "win_rate_recent", "profit_factor_recent",
            "drawdown_recent", "summary_note"]
    recent_events = [dict(zip(cols, r)) for r in recent]

    return DriftReport(
        generated_at=now_iso,
        window_trades=window,
        n_recent_trades=n,
        overall_status=overall,
        pause_new_trades=pause,
        metrics=metrics,
        summary_note=note,
        recent_events=recent_events,
    )


class AgenteDriftMonitor:
    def __init__(self, db_path: Path = DB_PATH, report_path: Path = REPORT_PATH,
                 window: int = 60) -> None:
        self.db_path = db_path
        self.report_path = report_path
        self.window = window
        self.logger = self._build_logger()
        DASHBOARD_DATA_DIR.mkdir(parents=True, exist_ok=True)

    def run(self) -> DriftReport:
        self.logger.info("Iniciando Drift Monitor (ventana=%d trades)", self.window)
        report = evaluate_drift(self.window, self.db_path)
        self.logger.info(
            "Drift: status=%s | trades_recientes=%d | pause=%s",
            report.overall_status, report.n_recent_trades, report.pause_new_trades,
        )
        self._export(report)
        return report

    def _export(self, report: DriftReport) -> None:
        payload = {
            **{k: v for k, v in asdict(report).items() if k != "metrics"},
            "metrics": [asdict(m) for m in report.metrics],
        }
        self.report_path.write_text(
            json.dumps(payload, indent=2, ensure_ascii=False, default=str),
            encoding="utf-8",
        )
        self.logger.info("Informe exportado → %s", self.report_path)

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


def main() -> int:
    parser = argparse.ArgumentParser(description="Drift Monitor de Leonex")
    parser.add_argument("--window", type=int, default=60,
                        help="Numero de trades recientes para rolling (default 60)")
    parser.add_argument("--reset-pause", action="store_true",
                        help="Desactiva manualmente el flag PAUSE_NEW_TRADES.")
    args = parser.parse_args()

    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")

    if args.reset_pause:
        set_system_state("pause_new_trades", "false")
        set_system_state("pause_reason", "")
        print("[OK] PAUSE_NEW_TRADES desactivado manualmente y pause_reason limpiada.")
        return 0

    agente = AgenteDriftMonitor(window=args.window)
    report = agente.run()

    sep = "=" * 64
    print(f"\n{sep}")
    print("Leonex -- Drift Monitor")
    print(sep)
    print(f"Estado global    : {report.overall_status}")
    print(f"Ventana          : {report.window_trades} trades  |  usados: {report.n_recent_trades}")
    print(f"PAUSE_NEW_TRADES : {report.pause_new_trades}")
    print(f"Nota             : {report.summary_note}")
    print("-" * 64)
    for m in report.metrics:
        tag = {"GREEN": "[OK]", "YELLOW": "[!!]", "RED": "[XX]"}.get(m.status, "[??]")
        print(f"  {tag} {m.name:<32} reciente={m.value_recent:>8.3f}  "
              f"base={m.value_baseline:>8.3f}  {m.note}")
    print(sep)
    if report.overall_status == "RED":
        print("Sistema PAUSADO. El executor no abrira nuevas posiciones.")
        print("Para reanudar manualmente: python agents/agente_drift_monitor.py --reset-pause")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
