"""
Agente de Datos de Leonex.

Descarga datos diarios de mercado, los normaliza, calcula metricas basicas y
guarda todo en SQLite para que otros agentes puedan leer una unica fuente.
"""

from __future__ import annotations

import json
import logging
import random
import sqlite3
import time
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 Iterable

import numpy as np
import pandas as pd
import yfinance as yf


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"
SNAPSHOT_PATH = DASHBOARD_DATA_DIR / "market_snapshot.json"
HISTORY_PATH = DASHBOARD_DATA_DIR / "price_history.json"

# Reintentos de descarga diaria. Un vacio de yfinance casi siempre es
# throttling/rate-limit transitorio (no un ticker inexistente); reintentar con
# backoff evita dias a medias como el 25/06 (146/509 tickers en 1d).
_DL_MAX_RETRIES = 4
_DL_BASE_DELAY_S = 1.5
# Espaciado defensivo ENTRE tickers. Sin esto, ~510 descargas seguidas sin
# pausa disparan el rate-limit de Yahoo tras ~146 tickers (el plateau 146/509
# clavado el 25 y 26/06): a partir de ahi el resto falla TODOS los reintentos
# porque la ventana de bloqueo dura mas que el backoff (~12s). El agente
# intradia ya usa 0.4s por ticker y por eso nunca cae a medias. 510*0.4≈3min
# extra, asumible para un cron diario.
_SLEEP_BETWEEN_TICKERS_S = 0.4


DEFAULT_UNIVERSE = {
    "equities": ["SPY", "QQQ", "IWM", "AAPL", "MSFT", "NVDA", "AMZN", "META", "GOOGL"],
    "forex": ["EURUSD=X", "GBPUSD=X", "USDJPY=X"],
    "crypto": ["BTC-USD", "ETH-USD", "SOL-USD"],
}

SP500_CSV_PATH = DATA_DIR / "sp500_components.csv"


def _load_prices_tickers_by_class(db_path: Path) -> dict[str, list[str]] | None:
    """Tickers YA presentes en la tabla `prices`, agrupados por asset_class.

    Es la MISMA fuente que usa el agente intradia (`SELECT DISTINCT ticker FROM
    prices`) y por eso el intradia nunca se queda corto. Sirve para que el
    universo diario sea AUTO-REPARABLE: una vez un ticker entra en prices, lo
    seguimos refrescando aunque el CSV del S&P 500 se haya encogido al fallback
    de 150 (agente_universe lo clobberea cuando Wikipedia bloquea la IP)."""
    if not db_path.exists():
        return None
    try:
        with sqlite3.connect(f"file:{db_path}?mode=ro", uri=True) as conn:
            rows = conn.execute(
                "SELECT DISTINCT ticker, asset_class FROM prices"
            ).fetchall()
    except Exception:
        return None
    if not rows:
        return None
    out: dict[str, list[str]] = {}
    for ticker, asset_class in rows:
        cls = asset_class if asset_class in ("forex", "crypto") else "equities"
        out.setdefault(cls, []).append(str(ticker))
    return out or None


def _load_sp500_data_pool(db_path: Path) -> dict[str, list[str]] | None:
    """Data pool = UNION de tres fuentes, para que el universo diario NUNCA se
    encoja silenciosamente:
      1. sp500_components.csv (S&P 500 completo, si esta disponible).
      2. Todos los tickers YA presentes en la tabla `prices` (auto-reparable:
         mismo set que sigue el intradia; impide que un fallo de scrape deje de
         refrescar ~360 tickers ya trackeados, como paso el 25-27/06).
      3. crypto/forex legacy de active_universe.
    Es el conjunto que se DESCARGA y que estudia el Strategy Lab — distinto del
    universo OPERATIVO (top-N) del executor, que sigue siendo active_universe.
    Devuelve None solo si NO hay ninguna fuente (entonces el caller hace fallback)."""
    pool: dict[str, set[str]] = {}

    # 1) CSV del S&P 500 (puede faltar o estar encogido al fallback de 150)
    if SP500_CSV_PATH.exists():
        try:
            symbols = (pd.read_csv(SP500_CSV_PATH)["Symbol"]
                       .astype(str).str.strip().str.replace(".", "-", regex=False))
            equities = {s for s in symbols if s and s.lower() != "nan"}
            if equities:
                pool.setdefault("equities", set()).update(equities)
        except Exception:
            pass

    # 2) Tickers ya presentes en prices (la clave del auto-reparado)
    prices_pool = _load_prices_tickers_by_class(db_path)
    if prices_pool:
        for cls, tks in prices_pool.items():
            pool.setdefault(cls, set()).update(tks)

    # 3) crypto/forex legacy de active_universe
    db = _load_universe_from_db(db_path)
    if db:
        for cls in ("crypto", "forex"):
            if db.get(cls):
                pool.setdefault(cls, set()).update(db[cls])

    if not pool or not any(pool.values()):
        return None
    return {cls: sorted(tks) for cls, tks in pool.items() if tks}


def _load_universe_from_db(db_path: Path) -> dict[str, list[str]] | None:
    """Lee active_universe si existe y tiene filas. Devuelve None si no hay datos."""
    if not db_path.exists():
        return None
    try:
        import sqlite3 as _sql
        with _sql.connect(db_path) as conn:
            rows = conn.execute(
                "SELECT ticker, asset_class FROM active_universe"
            ).fetchall()
    except Exception:
        return None
    if not rows:
        return None
    universe: dict[str, list[str]] = {}
    for ticker, asset_class in rows:
        cls = asset_class if asset_class in ("equity", "equities", "forex", "crypto") else "equities"
        # Normalizamos 'equity' (singular) → 'equities' para mantener consistencia
        if cls == "equity":
            cls = "equities"
        universe.setdefault(cls, []).append(ticker)
    return universe


@dataclass(frozen=True)
class AssetSnapshot:
    ticker: str
    asset_class: str
    last_close: float
    last_date: str
    daily_return: float
    volatility_20d: float
    volatility_60d: float
    momentum_20d: float
    momentum_60d: float
    zscore_20d: float
    data_points: int


class AgenteDatos:
    def __init__(
        self,
        db_path: Path = DB_PATH,
        snapshot_path: Path = SNAPSHOT_PATH,
        history_path: Path = HISTORY_PATH,
        universe: dict[str, list[str]] | None = None,
    ) -> None:
        self.db_path = db_path
        self.snapshot_path = snapshot_path
        self.history_path = history_path
        # Si no se pasa un universo explicito, intentar leer active_universe
        # (poblada por agente_universe.py). Si no existe o esta vacia,
        # fallback al DEFAULT_UNIVERSE (los 15 legacy).
        if universe is None:
            pool = _load_sp500_data_pool(db_path)
            if pool:
                self.universe = pool
                self.universe_source = "sp500_data_pool"
            else:
                dynamic = _load_universe_from_db(db_path)
                self.universe = dynamic if dynamic else DEFAULT_UNIVERSE
                self.universe_source = "active_universe" if dynamic else "default_legacy"
        else:
            self.universe = universe
            self.universe_source = "explicit"
        self.logger = self._build_logger()
        self._ensure_directories()
        self._ensure_schema()
        total = sum(len(v) for v in self.universe.values())
        self.logger.info("Universo cargado: %d tickers (fuente=%s)", total, self.universe_source)

    def run(self, period: str = "5y", incremental: bool = True) -> list[AssetSnapshot]:
        """Ejecuta el ciclo completo del agente.

        Args:
            period: rango maximo a descargar (e.g. "5y", "10y", "max").
            incremental: si True (default), descarga solo lo nuevo desde el
                ultimo timestamp en SQLite + 30 dias de overlap. Si la tabla
                esta vacia para un ticker, descarga `period` completo. Asi el
                historico se acumula con el tiempo sin descargas pesadas.
        """
        mode = "incremental" if incremental else "full"
        self.logger.info("Iniciando Agente de Datos (mode=%s, period=%s)", mode, period)
        snapshots: list[AssetSnapshot] = []
        histories: dict[str, pd.DataFrame] = {}

        for asset_class, tickers in self.universe.items():
            for ticker in tickers:
                # Throttle ANTES de cada descarga (incluida la 1a, coste ~0.4s)
                # para mantener el ritmo por debajo del rate-limit de Yahoo.
                # Va al principio del bucle para ejecutarse SIEMPRE, aunque la
                # iteracion previa hiciera `continue` o lanzara excepcion.
                time.sleep(_SLEEP_BETWEEN_TICKERS_S)
                try:
                    # Si incremental: detecta ultimo timestamp en SQLite
                    effective_period = period
                    # repair=True solo en la descarga COMPLETA (primera vez por
                    # ticker). En el refresh diario incremental se desactiva para
                    # no disparar peticiones extra a Yahoo y agotar el cupo.
                    use_repair = True
                    if incremental:
                        last_date = self._get_last_price_date(ticker)
                        if last_date is not None:
                            # Si tenemos algo, descargamos solo los ultimos 30 dias
                            # (overlap de seguridad por si el cron fallo varios dias)
                            effective_period = "30d"
                            use_repair = False
                            self.logger.info(
                                "%s incremental: ultimo dato=%s → descargo %s",
                                ticker, last_date, effective_period,
                            )
                        else:
                            self.logger.info(
                                "%s sin historico previo → descarga completa %s",
                                ticker, effective_period,
                            )

                    prices = self.download_asset(ticker=ticker, period=effective_period,
                                                 use_repair=use_repair)
                    if prices.empty:
                        self.logger.warning("%s no devolvio datos", ticker)
                        continue

                    self.save_prices(ticker=ticker, asset_class=asset_class, prices=prices)
                    # Para el snapshot necesitamos el historico COMPLETO desde SQLite
                    # (no solo los 30 dias del incremental), si no las metricas vol/momentum
                    # 60d se calculan mal.
                    full_history = self._load_full_history(ticker)
                    if full_history.empty:
                        full_history = prices
                    histories[ticker] = full_history
                    snapshot = self.build_snapshot(
                        ticker=ticker,
                        asset_class=asset_class,
                        prices=full_history,
                    )
                    snapshots.append(snapshot)
                    self.logger.info(
                        "%s actualizado: +%d nuevas, historico total=%d filas",
                        ticker, len(prices), len(full_history),
                    )
                except Exception:
                    self.logger.exception("Error actualizando %s", ticker)

        self.export_snapshot(snapshots)
        self.export_history(histories)
        self.logger.info(
            "Agente de Datos terminado (mode=%s). Activos OK=%d",
            mode, len(snapshots),
        )
        return snapshots

    def _get_last_price_date(self, ticker: str) -> str | None:
        """Devuelve la fecha del ultimo close registrado para el ticker, o None."""
        try:
            with sqlite3.connect(self.db_path) as conn:
                row = conn.execute(
                    "SELECT MAX(date) FROM prices WHERE ticker = ?", (ticker,),
                ).fetchone()
                return row[0] if row and row[0] else None
        except sqlite3.OperationalError:
            return None

    def _load_full_history(self, ticker: str) -> pd.DataFrame:
        """Carga TODO el historico del ticker desde SQLite para snapshot/metricas."""
        try:
            with sqlite3.connect(self.db_path) as conn:
                df = pd.read_sql(
                    "SELECT date, open, high, low, close, volume, return_log "
                    "FROM prices WHERE ticker = ? ORDER BY date ASC",
                    conn, params=(ticker,), parse_dates=["date"],
                )
            if df.empty:
                return df
            return df.set_index("date")
        except Exception:
            return pd.DataFrame()

    @staticmethod
    def _repair_close_spikes(prices: pd.DataFrame, ticker: str = "",
                             factor: float = 2.5) -> pd.DataFrame:
        """Repara ticks transitorios: una barra cuyo close salta >factor x
        respecto a AMBOS vecinos (sube y luego revierte) es casi siempre un
        precio erroneo de origen, no un movimiento real. Lo sustituye por la
        media geometrica de los vecinos. NO toca movimientos sostenidos (p.ej.
        MU +14% en un dia que NO revierte queda intacto).
        """
        if len(prices) < 3:
            return prices
        c = prices["close"].to_numpy(dtype="float64")
        prev = c[:-2]
        cur = c[1:-1]
        nxt = c[2:]
        with np.errstate(divide="ignore", invalid="ignore"):
            up = cur / prev
            down = cur / nxt
        # spike al alza: cur >> prev y cur >> next  | spike a la baja: <<
        bad = ((up > factor) & (down > factor)) | ((up < 1 / factor) & (down < 1 / factor))
        idx = np.where(bad)[0] + 1  # offset por el slice [1:-1]
        if len(idx) > 0:
            # media geometrica de los vecinos sanos
            c[idx] = np.sqrt(c[idx - 1] * c[idx + 1])
            prices = prices.copy()
            prices["close"] = c
        return prices

    def _download_once(self, ticker: str, period: str,
                       repair: bool = True) -> pd.DataFrame:
        """Una sola llamada a yfinance (1d). repair=True corrige errores de
        origen — splits mal aplicados, precios x100 (mezcla de divisa) y ticks
        absurdos (origen de artefactos tipo DD +178% o KLAC). Cae a la firma
        antigua si la version instalada no soporta `repair`.

        IMPORTANTE: repair=True hace PETICIONES EXTRA a Yahoo por ticker (re-
        descarga para detectar splits/100x). Con ~503 tickers eso multiplica el
        volumen de requests y agota el cupo de Yahoo a mitad (~147 cargados, el
        resto vacio). Por eso el refresh diario incremental lo llama con
        repair=False (1 req/ticker, como el agente intradia que SI completa los
        ~502); repair solo se usa en la descarga COMPLETA inicial, que es rara."""
        try:
            return yf.download(
                ticker,
                period=period,
                interval="1d",
                auto_adjust=True,
                repair=repair,
                progress=False,
                threads=False,
            )
        except TypeError:
            # yfinance antiguo sin parametro repair.
            return yf.download(
                ticker,
                period=period,
                interval="1d",
                auto_adjust=True,
                progress=False,
                threads=False,
            )

    def download_asset(self, ticker: str, period: str = "3y",
                       use_repair: bool = True) -> pd.DataFrame:
        """Descarga OHLCV diario (con reintentos) y normaliza columnas.

        Un resultado vacio de yfinance casi siempre es throttling/rate-limit
        transitorio, no que el ticker no exista. Por eso reintentamos con
        backoff exponencial + jitter en vez de saltar al primer vacio: sin
        esto, una ventana de rate limit deja el dia incompleto (p.ej. 146/509
        tickers en 1d el 25/06). Si tras agotar los reintentos sigue vacio, se
        registra un ERROR — ya no es un skip silencioso."""
        raw = pd.DataFrame()
        last_exc: Exception | None = None
        for attempt in range(1, _DL_MAX_RETRIES + 1):
            try:
                raw = self._download_once(ticker, period, repair=use_repair)
            except Exception as exc:  # red / transitorio
                last_exc = exc
                raw = pd.DataFrame()
            if not raw.empty:
                break
            if attempt < _DL_MAX_RETRIES:
                delay = _DL_BASE_DELAY_S * (2 ** (attempt - 1))
                delay += random.uniform(0, _DL_BASE_DELAY_S)  # jitter
                self.logger.warning(
                    "%s: yfinance vacio/error (intento %d/%d%s) — reintento "
                    "en %.1fs", ticker, attempt, _DL_MAX_RETRIES,
                    f", {last_exc!r}" if last_exc else "", delay,
                )
                time.sleep(delay)
        if raw.empty:
            self.logger.error(
                "%s: DESCARGA FALLIDA tras %d intentos%s — el dia puede quedar "
                "incompleto", ticker, _DL_MAX_RETRIES,
                f" (ultimo error: {last_exc!r})" if last_exc else "",
            )
            return pd.DataFrame()

        if isinstance(raw.columns, pd.MultiIndex):
            raw.columns = raw.columns.get_level_values(0)

        prices = raw.rename(
            columns={
                "Open": "open",
                "High": "high",
                "Low": "low",
                "Close": "close",
                "Volume": "volume",
            }
        )
        prices = prices[["open", "high", "low", "close", "volume"]].copy()
        prices.index = pd.to_datetime(prices.index).tz_localize(None)
        prices = prices.dropna(subset=["close"])
        prices = self._repair_close_spikes(prices, ticker)
        prices["return_log"] = np.log(prices["close"] / prices["close"].shift(1))
        prices["downloaded_at"] = datetime.now(UTC).isoformat()
        return prices.dropna(subset=["return_log"])

    def save_prices(self, ticker: str, asset_class: str, prices: pd.DataFrame) -> None:
        rows = prices.reset_index(names="date").copy()
        rows["ticker"] = ticker
        rows["asset_class"] = asset_class
        rows["date"] = rows["date"].dt.strftime("%Y-%m-%d")

        # IMPORTANTE: NO usar to_records() porque convierte numpy int64 a
        # binario y SQLite los guarda como BLOB en vez de REAL/INTEGER, lo
        # que rompe lecturas posteriores. Usamos itertuples con casts
        # explicitos a tipos nativos Python.
        records = [
            (
                str(ticker),
                str(asset_class),
                str(row.date),
                None if pd.isna(row.open) else float(row.open),
                None if pd.isna(row.high) else float(row.high),
                None if pd.isna(row.low) else float(row.low),
                None if pd.isna(row.close) else float(row.close),
                None if pd.isna(row.volume) else float(row.volume),
                None if pd.isna(row.return_log) else float(row.return_log),
                str(row.downloaded_at),
            )
            for row in rows.itertuples(index=False)
        ]

        with sqlite3.connect(self.db_path) as conn:
            conn.executemany(
                """
                INSERT OR REPLACE INTO prices (
                    ticker, asset_class, date, open, high, low, close, volume,
                    return_log, downloaded_at
                )
                VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
                """,
                records,
            )
            conn.commit()

    def build_snapshot(
        self,
        ticker: str,
        asset_class: str,
        prices: pd.DataFrame,
    ) -> AssetSnapshot:
        close = prices["close"]
        returns = prices["return_log"]
        last_return = float(returns.iloc[-1])
        vol_20d = float(returns.tail(20).std() * np.sqrt(252))
        vol_60d = float(returns.tail(60).std() * np.sqrt(252))
        momentum_20d = float(close.iloc[-1] / close.shift(20).iloc[-1] - 1)
        momentum_60d = float(close.iloc[-1] / close.shift(60).iloc[-1] - 1)

        rolling_mean = returns.tail(20).mean()
        rolling_std = returns.tail(20).std()
        zscore_20d = 0.0 if rolling_std == 0 else float((last_return - rolling_mean) / rolling_std)

        return AssetSnapshot(
            ticker=ticker,
            asset_class=asset_class,
            last_close=float(close.iloc[-1]),
            last_date=prices.index[-1].strftime("%Y-%m-%d"),
            daily_return=last_return,
            volatility_20d=vol_20d,
            volatility_60d=vol_60d,
            momentum_20d=momentum_20d,
            momentum_60d=momentum_60d,
            zscore_20d=zscore_20d,
            data_points=int(len(prices)),
        )

    def export_snapshot(self, snapshots: Iterable[AssetSnapshot]) -> None:
        payload = {
            "generated_at": datetime.now(UTC).isoformat(),
            "source": "yfinance",
            "assets": [asdict(snapshot) for snapshot in snapshots],
        }
        self.snapshot_path.write_text(
            json.dumps(payload, indent=2, ensure_ascii=False),
            encoding="utf-8",
        )

    def export_history(self, histories: dict[str, pd.DataFrame], lookback: int = 260) -> None:
        payload = {
            "generated_at": datetime.now(UTC).isoformat(),
            "source": "yfinance",
            "lookback": lookback,
            "series": {},
        }

        def _safe_float(v) -> float:
            """Convierte v a float, blindando contra BLOB legacy (bytes)."""
            if v is None:
                return 0.0
            if isinstance(v, (bytes, bytearray)):
                # SQLite BLOB legacy del bug to_records(); ignoramos
                return 0.0
            try:
                if pd.isna(v):
                    return 0.0
            except (TypeError, ValueError):
                pass
            try:
                return float(v)
            except (TypeError, ValueError):
                return 0.0

        def _safe_date(v) -> str:
            if hasattr(v, "strftime"):
                return v.strftime("%Y-%m-%d")
            return str(v)[:10]

        for ticker, prices in histories.items():
            rows = prices.tail(lookback).reset_index(names="date")
            payload["series"][ticker] = [
                {
                    "date": _safe_date(row.date),
                    "close": _safe_float(row.close),
                    "volume": _safe_float(row.volume),
                }
                for row in rows.itertuples(index=False)
            ]

        self.history_path.write_text(
            json.dumps(payload, indent=2, ensure_ascii=False),
            encoding="utf-8",
        )

    def _ensure_directories(self) -> None:
        for directory in (DATA_DIR, LOGS_DIR, DASHBOARD_DATA_DIR):
            directory.mkdir(parents=True, exist_ok=True)

    def _ensure_schema(self) -> None:
        with sqlite3.connect(self.db_path) as conn:
            conn.execute(
                """
                CREATE TABLE IF NOT EXISTS prices (
                    ticker TEXT NOT NULL,
                    asset_class TEXT NOT NULL,
                    date TEXT NOT NULL,
                    open REAL,
                    high REAL,
                    low REAL,
                    close REAL,
                    volume REAL,
                    return_log REAL,
                    downloaded_at TEXT NOT NULL,
                    PRIMARY KEY (ticker, date)
                )
                """
            )
            conn.execute(
                "CREATE INDEX IF NOT EXISTS idx_prices_date ON prices(date)"
            )
            conn.execute(
                "CREATE INDEX IF NOT EXISTS idx_prices_asset_class ON prices(asset_class)"
            )
            conn.commit()

    def _build_logger(self) -> logging.Logger:
        LOGS_DIR.mkdir(parents=True, exist_ok=True)
        logger = logging.getLogger("agente_datos")
        logger.setLevel(logging.INFO)
        logger.handlers.clear()

        formatter = logging.Formatter(
            "%(asctime)s | %(levelname)s | %(name)s | %(message)s"
        )
        file_handler = logging.FileHandler(LOGS_DIR / "agente_datos.log", encoding="utf-8")
        file_handler.setFormatter(formatter)
        stream_handler = logging.StreamHandler()
        stream_handler.setFormatter(formatter)

        logger.addHandler(file_handler)
        logger.addHandler(stream_handler)
        return logger


def main() -> None:
    import argparse
    parser = argparse.ArgumentParser(description="Agente de Datos de Leonex")
    parser.add_argument("--period", type=str, default="10y",
                        help="Rango maximo a descargar la PRIMERA vez por ticker. "
                             "Default 10y. Se ignora en modo incremental si ya hay datos.")
    parser.add_argument("--full", action="store_true",
                        help="Forzar descarga completa (no incremental). Util tras "
                             "ampliar el universo o tras un corrupt del SQLite.")
    args = parser.parse_args()
    snapshots = AgenteDatos().run(period=args.period, incremental=not args.full)
    print(f"Activos actualizados: {len(snapshots)}")
    print(f"Base de datos Leonex: {DB_PATH}")
    print(f"Snapshot dashboard: {SNAPSHOT_PATH}")
    print(f"Historial dashboard: {HISTORY_PATH}")


if __name__ == "__main__":
    main()

