"""
Agente de Señales de Leonex — Régimen + Validación combinados.

Lee precios y régimen histórico desde SQLite, aplica una señal distinta
según el régimen de cada día y guarda las señales resultantes en la tabla
``signals`` para que el Agente de Validación pueda evaluarlas con Walk-Forward,
DSR y Monte Carlo.

Lógica adaptativa (long-only por defecto):
    TENDENCIA_ALCISTA  →  long si momentum_20d > 0
    TENDENCIA_BAJISTA  →  flat  (en modo long-short: short si momentum_20d < 0)
    LATERAL            →  mean reversion: long si RSI(14) < 30
                          (en long-short: short si RSI > 70)
    CRISIS             →  flat  (volatilidad anormal: fuera del mercado)
    DESCONOCIDO        →  flat

Salidas:
    SQLite: tabla ``signals`` con (ticker, date, signal, regime, strategy)
    JSON:   dashboard/data/signals_report.json con resumen por activo

Uso:
    python agents/agente_senales.py
    python agents/agente_senales.py --long-short
    python agents/agente_senales.py --ticker SPY
"""

from __future__ import annotations

import argparse
import json
import logging
import sqlite3
from dataclasses import asdict, dataclass
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

# Reutilizamos los indicadores del Agente de Régimen.
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent))
from agente_regimen import calculate_indicators, classify_regime  # noqa: E402

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

STRATEGY_NAME = "regime_adaptive_v1"

# Parámetros de las sub-estrategias
MOMENTUM_LOOKBACK = 20
RSI_PERIOD = 14
RSI_OVERSOLD = 30
RSI_OVERBOUGHT = 70


@dataclass
class TickerSummary:
    ticker: str
    n_days: int
    n_long: int
    n_short: int
    n_flat: int
    regime_breakdown: dict
    last_date: str
    last_regime: str
    last_signal: int
    last_rsi: float = 0.0
    last_momentum_20: float = 0.0
    last_close: float = 0.0


def compute_rsi(close: pd.Series, period: int = RSI_PERIOD) -> pd.Series:
    """RSI clásico de Wilder."""
    delta = close.diff()
    gain = delta.clip(lower=0).ewm(alpha=1 / period, adjust=False).mean()
    loss = (-delta.clip(upper=0)).ewm(alpha=1 / period, adjust=False).mean()
    rs = gain / loss.replace(0, np.nan)
    rsi = 100 - (100 / (1 + rs))
    return rsi.fillna(50.0)


def build_regime_signal(prices: pd.DataFrame, long_short: bool = False) -> pd.DataFrame:
    """
    Construye una señal condicional al régimen para cada día.

    Devuelve un DataFrame indexado por fecha con columnas:
        regime, momentum_20, rsi_14, raw_signal, signal

    raw_signal = señal del día (potencial lookahead)
    signal     = raw_signal.shift(1)  → la usada para el retorno del día siguiente
    """
    df = calculate_indicators(prices.copy())
    df["regime"] = df.apply(classify_regime, axis=1)
    df["momentum_20"] = df["close"].pct_change(MOMENTUM_LOOKBACK)
    df["rsi_14"] = compute_rsi(df["close"], RSI_PERIOD)

    def decide(row: pd.Series) -> int:
        regime = row["regime"]
        mom = row["momentum_20"]
        rsi = row["rsi_14"]

        if regime == "TENDENCIA_ALCISTA":
            if not pd.isna(mom) and mom > 0:
                return 1
            return 0
        if regime == "TENDENCIA_BAJISTA":
            if long_short and not pd.isna(mom) and mom < 0:
                return -1
            return 0
        if regime == "LATERAL":
            if not pd.isna(rsi):
                if rsi < RSI_OVERSOLD:
                    return 1
                if long_short and rsi > RSI_OVERBOUGHT:
                    return -1
            return 0
        # CRISIS y DESCONOCIDO → flat
        return 0

    df["raw_signal"] = df.apply(decide, axis=1).astype(int)
    df["signal"] = df["raw_signal"].shift(1).fillna(0).astype(int)
    return df[["regime", "momentum_20", "rsi_14", "raw_signal", "signal"]]


class AgenteSenales:
    def __init__(
        self,
        db_path: Path = DB_PATH,
        report_path: Path = REPORT_PATH,
        long_short: bool = False,
    ) -> None:
        self.db_path = db_path
        self.report_path = report_path
        self.long_short = long_short
        self.strategy = STRATEGY_NAME + ("_ls" if long_short else "_lo")
        self.logger = self._build_logger()
        DASHBOARD_DATA_DIR.mkdir(parents=True, exist_ok=True)
        self._ensure_schema()

    def run(self, tickers: Optional[list[str]] = None) -> list[TickerSummary]:
        self.logger.info(
            "Iniciando Agente de Señales (estrategia=%s)", self.strategy
        )
        summaries: list[TickerSummary] = []
        with sqlite3.connect(self.db_path) as conn:
            if tickers is None:
                tickers_df = pd.read_sql(
                    "SELECT DISTINCT ticker FROM prices ORDER BY ticker",
                    conn,
                )
                tickers = tickers_df["ticker"].tolist()

            # Limpiamos señales previas de esta estrategia para los tickers procesados.
            placeholders = ",".join("?" * len(tickers))
            conn.execute(
                f"DELETE FROM signals WHERE strategy = ? AND ticker IN ({placeholders})",
                [self.strategy, *tickers],
            )
            conn.commit()

            for ticker in tickers:
                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) < 60:
                    self.logger.warning(
                        "%s: datos insuficientes (%d filas)", ticker, len(prices)
                    )
                    continue

                signals_df = build_regime_signal(prices, long_short=self.long_short)
                self._save_signals(conn, ticker, signals_df)

                summary = self._summarize(ticker, signals_df, prices)
                summaries.append(summary)
                self.logger.info(
                    "%s → long=%d short=%d flat=%d | último: %s/%s/%d",
                    ticker,
                    summary.n_long,
                    summary.n_short,
                    summary.n_flat,
                    summary.last_date,
                    summary.last_regime,
                    summary.last_signal,
                )
            conn.commit()

        self._export(summaries)
        self.logger.info(
            "Agente de Señales terminado. Activos procesados=%d", len(summaries)
        )
        return summaries

    # ── Persistencia ─────────────────────────────────────────────────────────
    def _save_signals(
        self,
        conn: sqlite3.Connection,
        ticker: str,
        signals_df: pd.DataFrame,
    ) -> None:
        rows = signals_df.reset_index(names="date").copy()
        rows["date"] = rows["date"].dt.strftime("%Y-%m-%d")
        rows["ticker"] = ticker
        rows["strategy"] = self.strategy
        rows["momentum_20"] = rows["momentum_20"].astype(float).fillna(0.0)
        rows["rsi_14"] = rows["rsi_14"].astype(float).fillna(50.0)
        rows["signal"] = rows["signal"].astype(int)
        # Importante: usar tuplas Python nativas — to_records() empaqueta
        # numpy int64 que SQLite guarda como BLOB.
        records = [
            (
                str(r.ticker),
                str(r.date),
                str(r.strategy),
                int(r.signal),
                str(r.regime),
                float(r.momentum_20),
                float(r.rsi_14),
            )
            for r in rows[
                ["ticker", "date", "strategy", "signal", "regime", "momentum_20", "rsi_14"]
            ].itertuples(index=False)
        ]
        conn.executemany(
            """
            INSERT OR REPLACE INTO signals
                (ticker, date, strategy, signal, regime, momentum_20, rsi_14)
            VALUES (?, ?, ?, ?, ?, ?, ?)
            """,
            records,
        )

    def _summarize(self, ticker: str, signals_df: pd.DataFrame, prices: pd.DataFrame) -> TickerSummary:
        signal = signals_df["signal"]
        regime = signals_df["regime"]
        n_long = int((signal == 1).sum())
        n_short = int((signal == -1).sum())
        n_flat = int((signal == 0).sum())
        regime_counts = regime.value_counts().to_dict()
        last_rsi = float(signals_df["rsi_14"].iloc[-1]) if "rsi_14" in signals_df else 0.0
        last_mom = float(signals_df["momentum_20"].iloc[-1]) if "momentum_20" in signals_df else 0.0
        if pd.isna(last_mom):
            last_mom = 0.0
        last_close = float(prices["close"].iloc[-1])
        return TickerSummary(
            ticker=ticker,
            n_days=int(len(signals_df)),
            n_long=n_long,
            n_short=n_short,
            n_flat=n_flat,
            regime_breakdown={str(k): int(v) for k, v in regime_counts.items()},
            last_date=signals_df.index[-1].strftime("%Y-%m-%d"),
            last_regime=str(regime.iloc[-1]),
            last_signal=int(signal.iloc[-1]),
            last_rsi=round(last_rsi, 2),
            last_momentum_20=round(last_mom, 4),
            last_close=round(last_close, 4),
        )

    def _export(self, summaries: list[TickerSummary]) -> None:
        payload = {
            "generated_at": datetime.now(UTC).isoformat(),
            "strategy": self.strategy,
            "long_short": self.long_short,
            "config": {
                "momentum_lookback": MOMENTUM_LOOKBACK,
                "rsi_period": RSI_PERIOD,
                "rsi_oversold": RSI_OVERSOLD,
                "rsi_overbought": RSI_OVERBOUGHT,
            },
            "summary": {
                "total": len(summaries),
                "con_senal_hoy": sum(1 for s in summaries if s.last_signal != 0),
            },
            "tickers": [asdict(s) for s in summaries],
        }
        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)

    # ── Esquema y logging ────────────────────────────────────────────────────
    def _ensure_schema(self) -> None:
        with sqlite3.connect(self.db_path) as conn:
            conn.execute(
                """
                CREATE TABLE IF NOT EXISTS signals (
                    ticker TEXT NOT NULL,
                    date TEXT NOT NULL,
                    strategy TEXT NOT NULL,
                    signal INTEGER NOT NULL,
                    regime TEXT,
                    momentum_20 REAL,
                    rsi_14 REAL,
                    PRIMARY KEY (ticker, date, strategy)
                )
                """
            )
            conn.commit()

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


def main() -> None:
    parser = argparse.ArgumentParser(description="Agente de Señales adaptativas al régimen")
    parser.add_argument("--long-short", action="store_true",
                        help="Habilitar señales short además de long")
    parser.add_argument("--ticker", type=str, default=None,
                        help="Procesar solo un ticker concreto")
    args = parser.parse_args()

    # UTF-8 en Windows
    import sys
    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 = AgenteSenales(long_short=args.long_short)
    summaries = agente.run(tickers=[args.ticker] if args.ticker else None)

    print(f"\n{'=' * 55}")
    print(f"Leonex -- Agente de Señales completado")
    print(f"{'=' * 55}")
    print(f"Estrategia        : {agente.strategy}")
    print(f"Activos procesados: {len(summaries)}")
    print(f"Con señal hoy     : {sum(1 for s in summaries if s.last_signal != 0)}")
    print(f"{'-' * 55}")
    for s in summaries:
        flag = "LONG" if s.last_signal == 1 else ("SHORT" if s.last_signal == -1 else "FLAT")
        print(f"  {s.ticker:10s} {s.last_regime:18s} → {flag}")


if __name__ == "__main__":
    main()
