"""
Agente Fractional Differentiation de Leonex (AFML cap. 5).

Por cada ticker del universo activo, busca el MINIMO orden fraccional d ∈ [0,1]
que vuelve la serie log-precio estacionaria (test ADF p_value <= 0.05).

Por que importa:
    Si en el futuro queremos usar precios (o log-precios) DIRECTAMENTE como
    feature del Meta o de otro modelo, deben ser estacionarios para no
    sesgar el aprendizaje. Lopez de Prado demuestra que con d~0.4 se logra
    estacionariedad preservando >80% de la memoria de la serie original.

Estado actual de Leonex:
    El Meta entrena con features ya estacionarias (momentum_20, RSI,
    sigma_entry, regime dummies). Asi que este agente es DIAGNOSTICO /
    PREPARATORIO: te dice cual seria el d optimo si en el futuro anadimos
    log_price como feature.

Uso:
    python agents/agente_frac_diff.py
    python agents/agente_frac_diff.py --d-grid 0.0,0.1,0.2,0.3,0.4,0.5
"""

from __future__ import annotations

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

import numpy as np
import pandas as pd

sys.path.insert(0, str(Path(__file__).resolve().parent))
from fractional_diff import (  # noqa: E402
    find_min_d_for_stationarity,
    adf_pvalue,
    DEFAULT_D_GRID,
)


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

DEFAULT_MIN_HISTORY = 250    # dias minimo para evaluar
DEFAULT_USE_LOG = True       # log-precio en vez de precio crudo (recomendado)


@dataclass
class TickerResult:
    ticker: str
    n_obs: int
    log_used: bool
    adf_pvalue_at_d0: Optional[float]    # p-value de la serie original
    d_min: Optional[float]
    adf_pvalue_at_d_min: Optional[float]
    correlation_with_original: Optional[float]
    verdict: str                          # PASS | FAIL | INSUFFICIENT_DATA
    note: str = ""


@dataclass
class FracDiffReport:
    generated_at: str
    n_tickers_analyzed: int
    n_pass: int
    n_fail: int
    n_insufficient: int
    d_recommended_global: Optional[float]   # mediana de d_min en los PASS
    d_grid: list[float] = field(default_factory=list)
    adf_threshold: float = 0.05
    use_log_price: bool = True
    tickers: list[TickerResult] = field(default_factory=list)
    summary_note: str = ""


def load_active_universe(db_path: Path = DB_PATH) -> list[str]:
    """Lee la tabla active_universe (o fallback a tickers con datos)."""
    if not db_path.exists():
        return []
    with sqlite3.connect(db_path) as conn:
        try:
            df = pd.read_sql("SELECT DISTINCT ticker FROM active_universe", conn)
            if not df.empty:
                return df["ticker"].tolist()
        except Exception:
            pass
        try:
            df = pd.read_sql("SELECT DISTINCT ticker FROM prices", conn)
            return df["ticker"].tolist()
        except Exception:
            return []


def load_prices_for_ticker(ticker: str, db_path: Path = DB_PATH) -> pd.Series:
    """Lee close-price para un ticker, indexado por fecha."""
    with sqlite3.connect(db_path) as conn:
        df = pd.read_sql(
            "SELECT date, close FROM prices WHERE ticker = ? ORDER BY date ASC",
            conn, params=(ticker,), parse_dates=["date"],
        )
    if df.empty:
        return pd.Series(dtype=float)
    df = df.set_index("date")
    return df["close"]


def analyze_ticker(
    ticker: str, db_path: Path,
    d_grid: list[float],
    use_log: bool = DEFAULT_USE_LOG,
    min_history: int = DEFAULT_MIN_HISTORY,
    logger: Optional[logging.Logger] = None,
) -> TickerResult:
    prices = load_prices_for_ticker(ticker, db_path)
    if len(prices) < min_history:
        return TickerResult(
            ticker=ticker, n_obs=int(len(prices)), log_used=use_log,
            adf_pvalue_at_d0=None, d_min=None,
            adf_pvalue_at_d_min=None, correlation_with_original=None,
            verdict="INSUFFICIENT_DATA",
            note=f"Only {len(prices)} obs, minimum {min_history}",
        )
    # Usar log para suavizar (recomendacion AFML)
    series = np.log(prices) if use_log else prices
    series = series.dropna()
    if len(series) < min_history:
        return TickerResult(
            ticker=ticker, n_obs=int(len(series)), log_used=use_log,
            adf_pvalue_at_d0=None, d_min=None,
            adf_pvalue_at_d_min=None, correlation_with_original=None,
            verdict="INSUFFICIENT_DATA",
            note=f"After dropna only {len(series)} obs",
        )

    # p-value de la serie original (d=0)
    p0 = adf_pvalue(series)

    # Busqueda
    res = find_min_d_for_stationarity(series, d_grid=d_grid)
    d_min = res["d_min"]
    if d_min is None:
        return TickerResult(
            ticker=ticker, n_obs=int(len(series)), log_used=use_log,
            adf_pvalue_at_d0=p0, d_min=None,
            adf_pvalue_at_d_min=None, correlation_with_original=None,
            verdict="FAIL", note=res["note"],
        )
    # Tomamos el row del d_min
    row = next(r for r in res["all_results"] if r["d"] == d_min)
    return TickerResult(
        ticker=ticker, n_obs=int(len(series)), log_used=use_log,
        adf_pvalue_at_d0=p0,
        d_min=float(d_min),
        adf_pvalue_at_d_min=row.get("adf_pvalue"),
        correlation_with_original=row.get("correlation_with_original"),
        verdict="PASS", note=res["note"],
    )


def main() -> int:
    parser = argparse.ArgumentParser(description="Agente Fractional Differentiation de Leonex")
    parser.add_argument("--d-grid", type=str,
                        default=",".join(f"{d:.1f}" for d in DEFAULT_D_GRID),
                        help="Grid de d a probar, separado por comas. Default 0.0..1.0 en pasos 0.1")
    parser.add_argument("--use-log", action="store_true", default=True,
                        help="Usar log(precio) en vez de precio crudo (recomendado).")
    parser.add_argument("--min-history", type=int, default=DEFAULT_MIN_HISTORY)
    parser.add_argument("--limit-tickers", type=int, default=0,
                        help="Si > 0, limita el numero de tickers analizados (smoke test).")
    args = parser.parse_args()

    LOGS_DIR.mkdir(parents=True, exist_ok=True)
    DASHBOARD_DATA_DIR.mkdir(parents=True, exist_ok=True)
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
        handlers=[
            logging.FileHandler(LOGS_DIR / "agente_frac_diff.log", encoding="utf-8"),
            logging.StreamHandler(),
        ],
    )
    log = logging.getLogger("agente_frac_diff")

    d_grid = [float(x) for x in args.d_grid.split(",") if x.strip()]
    universe = load_active_universe(DB_PATH)
    if args.limit_tickers and len(universe) > args.limit_tickers:
        universe = universe[: args.limit_tickers]
    log.info("Analizando %d tickers con d_grid=%s", len(universe), d_grid)

    results: list[TickerResult] = []
    for i, ticker in enumerate(universe, 1):
        try:
            r = analyze_ticker(ticker, DB_PATH, d_grid=d_grid,
                               use_log=args.use_log, min_history=args.min_history,
                               logger=log)
            results.append(r)
            tag = {"PASS": "[OK]", "FAIL": "[FAIL]", "INSUFFICIENT_DATA": "[--]"}[r.verdict]
            d_str = f"d_min={r.d_min:.1f}" if r.d_min is not None else "d_min=-"
            log.info("  %3d/%d %s %s %s n_obs=%d", i, len(universe), tag,
                     ticker, d_str, r.n_obs)
        except Exception as exc:
            log.warning("Error analizando %s: %s", ticker, exc)
            results.append(TickerResult(
                ticker=ticker, n_obs=0, log_used=args.use_log,
                adf_pvalue_at_d0=None, d_min=None,
                adf_pvalue_at_d_min=None, correlation_with_original=None,
                verdict="FAIL", note=f"exception: {exc!r}",
            ))

    n_pass = sum(1 for r in results if r.verdict == "PASS")
    n_fail = sum(1 for r in results if r.verdict == "FAIL")
    n_insuf = sum(1 for r in results if r.verdict == "INSUFFICIENT_DATA")

    # Recomendacion global: mediana de d_min entre los PASS
    d_pass = [r.d_min for r in results if r.verdict == "PASS" and r.d_min is not None]
    d_global = float(np.median(d_pass)) if d_pass else None

    summary = (
        f"Analyzed {len(results)} tickers: {n_pass} PASS, {n_fail} FAIL, "
        f"{n_insuf} insufficient data. "
        + (f"Global recommended d (median of d_min over PASS) = {d_global:.2f}."
           if d_global is not None else
           "Could not derive a global d for lack of PASSes.")
    )
    log.info(summary)

    report = FracDiffReport(
        generated_at=datetime.now(UTC).isoformat(),
        n_tickers_analyzed=len(results),
        n_pass=n_pass, n_fail=n_fail, n_insufficient=n_insuf,
        d_recommended_global=d_global,
        d_grid=d_grid, adf_threshold=0.05,
        use_log_price=args.use_log,
        tickers=results,
        summary_note=summary,
    )
    payload = {k: v for k, v in asdict(report).items() if k != "tickers"}
    payload["tickers"] = [asdict(t) for t in report.tickers]
    REPORT_PATH.write_text(
        json.dumps(payload, indent=2, ensure_ascii=False, default=str),
        encoding="utf-8",
    )
    log.info("Reporte exportado → %s", REPORT_PATH)

    # Resumen consola
    sep = "=" * 64
    print(f"\n{sep}")
    print("Leonex -- Fractional Differentiation (AFML cap.5) — diagnostico")
    print(sep)
    print(f"Tickers analizados : {len(results)}")
    print(f"  PASS             : {n_pass}")
    print(f"  FAIL             : {n_fail}")
    print(f"  INSUFFICIENT     : {n_insuf}")
    if d_global is not None:
        print(f"  d recomendado    : {d_global:.2f}  (mediana de d_min en PASSes)")
    print()
    print("Top 10 tickers por orden de d_min:")
    sorted_pass = sorted([r for r in results if r.verdict == "PASS"],
                         key=lambda r: (r.d_min if r.d_min is not None else 99))
    for r in sorted_pass[:10]:
        print(f"  {r.ticker:<10} d_min={r.d_min:.1f}  ADF_p_d0={r.adf_pvalue_at_d0 or 0:.4f}  "
              f"ADF_p_dmin={r.adf_pvalue_at_d_min or 0:.4f}  "
              f"corr={r.correlation_with_original or 0:.3f}")
    print(sep)
    return 0


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