"""
Agente Timeframe Selector de Leonex.

Para cada ticker del universo, evalua un backtest minimalista de momentum
en tres timeframes:
    daily (tabla prices)
    1h    (tabla prices_intraday)
    4h    (tabla prices_intraday)

Calcula Sharpe IS de la estrategia simple "long si momentum_N > 0, short si
< 0", anualizado correctamente segun la frecuencia. Elige el timeframe
ganador SOLO si supera al actual por un margen > 0.20 Sharpe (evita
flip-flopping de mes a mes).

Output:
    Tabla asset_timeframes(ticker, chosen_timeframe, sharpe_daily,
                           sharpe_1h, sharpe_4h, updated_at)
    JSON dashboard/data/timeframe_selector_report.json
"""

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

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

DEFAULT_MARGIN_SHARPE = 0.20      # diferencia minima Sharpe para cambiar de TF
DEFAULT_MOMENTUM_LOOKBACK = {     # barras de lookback por timeframe
    "daily": 20,
    "1h":    24,                  # ~1 dia trading
    "4h":    18,                  # ~3 dias trading equity (1.6 barras/dia)
}
PERIODS_PER_YEAR = {              # para anualizar Sharpe
    "daily": 252,
    "1h":    252 * 6.5,           # sesion RTH equity
    "4h":    252 * 1.6,           # idem agregado a 4h
}
MIN_BARS = {                      # minimo de barras para considerar fiable
    "daily": 120,
    "1h":    500,
    "4h":    200,
}


@dataclass
class TickerSelection:
    ticker: str
    chosen_timeframe: str
    chosen_sharpe: float
    sharpe_daily: Optional[float]
    sharpe_1h: Optional[float]
    sharpe_4h: Optional[float]
    n_bars_daily: int
    n_bars_1h: int
    n_bars_4h: int
    margin_vs_2nd: float
    note: str = ""


@dataclass
class SelectorReport:
    generated_at: str
    n_tickers: int
    n_daily: int = 0
    n_1h: int = 0
    n_4h: int = 0
    margin_threshold: float = DEFAULT_MARGIN_SHARPE
    selections: list[TickerSelection] = field(default_factory=list)
    summary_note: str = ""


def ensure_schema(db_path: Path = DB_PATH) -> None:
    with sqlite3.connect(db_path) as conn:
        conn.execute(
            """
            CREATE TABLE IF NOT EXISTS asset_timeframes (
                ticker TEXT PRIMARY KEY,
                chosen_timeframe TEXT NOT NULL,
                chosen_sharpe REAL,
                sharpe_daily REAL,
                sharpe_1h REAL,
                sharpe_4h REAL,
                margin_vs_2nd REAL,
                updated_at TEXT
            )
            """
        )
        conn.commit()


# ─────────────────────────────────────────────────────────────────────────────
# Backtest minimalista de momentum
# ─────────────────────────────────────────────────────────────────────────────

def load_series(ticker: str, timeframe: str, db_path: Path = DB_PATH) -> pd.Series:
    """Devuelve close indexado por fecha/timestamp."""
    if timeframe == "daily":
        with sqlite3.connect(db_path) as conn:
            df = pd.read_sql(
                "SELECT date AS ts, close FROM prices "
                "WHERE ticker = ? ORDER BY date ASC",
                conn, params=(ticker,), parse_dates=["ts"],
            )
    else:
        with sqlite3.connect(db_path) as conn:
            df = pd.read_sql(
                "SELECT ts, close FROM prices_intraday "
                "WHERE ticker = ? AND timeframe = ? ORDER BY ts ASC",
                conn, params=(ticker, timeframe), parse_dates=["ts"],
            )
    if df.empty:
        return pd.Series(dtype=float)
    return df.set_index("ts")["close"].dropna()


def momentum_backtest_sharpe(close: pd.Series, lookback: int,
                             periods_per_year: float) -> Optional[float]:
    """Backtest minimalista: signal = sign(returns acumulados de lookback periodos).
    Position = signal lag-1. PnL = position * returns_t.
    Sharpe anualizado sin costes.
    """
    if len(close) < lookback + 30:
        return None
    rets = close.pct_change().dropna()
    if rets.empty:
        return None
    # Momentum signal: positivo si retorno acumulado lookback > 0
    momentum = (close / close.shift(lookback) - 1.0).dropna()
    signal = np.sign(momentum).reindex(rets.index, method="ffill").fillna(0.0)
    # Aplicamos signal con lag 1 (no podemos operar el dato de hoy con info de hoy)
    position = signal.shift(1).fillna(0.0)
    pnl = (position * rets).dropna()
    if len(pnl) < 30 or pnl.std() <= 0:
        return None
    sharpe = float(pnl.mean() / pnl.std() * math.sqrt(periods_per_year))
    if not math.isfinite(sharpe):
        return None
    return sharpe


# ─────────────────────────────────────────────────────────────────────────────
# Selector
# ─────────────────────────────────────────────────────────────────────────────

def select_for_ticker(ticker: str, db_path: Path,
                      margin: float = DEFAULT_MARGIN_SHARPE
                      ) -> TickerSelection:
    sharpes: dict[str, Optional[float]] = {}
    n_bars: dict[str, int] = {}
    for tf in ("daily", "1h", "4h"):
        close = load_series(ticker, tf, db_path)
        n_bars[tf] = int(len(close))
        if n_bars[tf] < MIN_BARS[tf]:
            sharpes[tf] = None
            continue
        sharpes[tf] = momentum_backtest_sharpe(
            close, DEFAULT_MOMENTUM_LOOKBACK[tf], PERIODS_PER_YEAR[tf],
        )

    # Default: daily si no hay nada mejor
    candidates = [(tf, s) for tf, s in sharpes.items() if s is not None]
    if not candidates:
        return TickerSelection(
            ticker=ticker, chosen_timeframe="daily", chosen_sharpe=0.0,
            sharpe_daily=sharpes.get("daily"),
            sharpe_1h=sharpes.get("1h"),
            sharpe_4h=sharpes.get("4h"),
            n_bars_daily=n_bars["daily"], n_bars_1h=n_bars["1h"],
            n_bars_4h=n_bars["4h"], margin_vs_2nd=0.0,
            note="insufficient_data_all_timeframes",
        )

    candidates.sort(key=lambda x: x[1], reverse=True)
    best_tf, best_sh = candidates[0]
    margin_vs_2nd = best_sh - (candidates[1][1] if len(candidates) > 1 else best_sh)

    # Anti-flip-flop: si daily existe y la mejora no supera el margen, nos
    # quedamos con daily. Esto da estabilidad ante ruido.
    sh_daily = sharpes.get("daily")
    chosen_tf = best_tf
    chosen_sh = best_sh
    note = f"best_sharpe={best_sh:.2f} on {best_tf}"
    if sh_daily is not None and best_tf != "daily":
        if best_sh - sh_daily < margin:
            chosen_tf = "daily"
            chosen_sh = sh_daily
            note = (f"daily kept: best TF {best_tf} ({best_sh:.2f}) does not beat "
                    f"daily ({sh_daily:.2f}) by margin {margin:.2f}")

    return TickerSelection(
        ticker=ticker, chosen_timeframe=chosen_tf,
        chosen_sharpe=round(chosen_sh, 4),
        sharpe_daily=round(sharpes.get("daily"), 4) if sharpes.get("daily") is not None else None,
        sharpe_1h=round(sharpes.get("1h"), 4) if sharpes.get("1h") is not None else None,
        sharpe_4h=round(sharpes.get("4h"), 4) if sharpes.get("4h") is not None else None,
        n_bars_daily=n_bars["daily"], n_bars_1h=n_bars["1h"],
        n_bars_4h=n_bars["4h"],
        margin_vs_2nd=round(margin_vs_2nd, 4),
        note=note,
    )


def persist_selections(selections: list[TickerSelection],
                       db_path: Path = DB_PATH) -> None:
    if not selections:
        return
    now_iso = datetime.now(UTC).isoformat()
    with sqlite3.connect(db_path) as conn:
        conn.executemany(
            """
            INSERT INTO asset_timeframes
                (ticker, chosen_timeframe, chosen_sharpe, sharpe_daily,
                 sharpe_1h, sharpe_4h, margin_vs_2nd, updated_at)
            VALUES (?, ?, ?, ?, ?, ?, ?, ?)
            ON CONFLICT(ticker) DO UPDATE SET
                chosen_timeframe=excluded.chosen_timeframe,
                chosen_sharpe=excluded.chosen_sharpe,
                sharpe_daily=excluded.sharpe_daily,
                sharpe_1h=excluded.sharpe_1h,
                sharpe_4h=excluded.sharpe_4h,
                margin_vs_2nd=excluded.margin_vs_2nd,
                updated_at=excluded.updated_at
            """,
            [
                (s.ticker, s.chosen_timeframe, s.chosen_sharpe,
                 s.sharpe_daily, s.sharpe_1h, s.sharpe_4h,
                 s.margin_vs_2nd, now_iso)
                for s in selections
            ],
        )
        conn.commit()


def load_universe(db_path: Path = DB_PATH) -> list[str]:
    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
        df = pd.read_sql("SELECT DISTINCT ticker FROM prices", conn)
        return df["ticker"].tolist()


def main() -> int:
    parser = argparse.ArgumentParser(description="Agente Timeframe Selector")
    parser.add_argument("--margin", type=float, default=DEFAULT_MARGIN_SHARPE,
                        help="Margen Sharpe minimo para cambiar de TF (default 0.20)")
    args = parser.parse_args()

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

    universe = load_universe(DB_PATH)
    log.info("Evaluando %d tickers con margin=%.2f", len(universe), args.margin)

    selections: list[TickerSelection] = []
    for i, ticker in enumerate(universe, 1):
        try:
            sel = select_for_ticker(ticker, DB_PATH, margin=args.margin)
            selections.append(sel)
            log.info("  %3d/%d %s → %s (Sharpe daily=%.2s 1h=%.2s 4h=%.2s)",
                     i, len(universe), ticker, sel.chosen_timeframe,
                     f"{sel.sharpe_daily:.2f}" if sel.sharpe_daily is not None else "  -",
                     f"{sel.sharpe_1h:.2f}" if sel.sharpe_1h is not None else "  -",
                     f"{sel.sharpe_4h:.2f}" if sel.sharpe_4h is not None else "  -")
        except Exception as exc:
            log.warning("Fallo %s: %s", ticker, exc)

    persist_selections(selections, DB_PATH)

    n_daily = sum(1 for s in selections if s.chosen_timeframe == "daily")
    n_1h = sum(1 for s in selections if s.chosen_timeframe == "1h")
    n_4h = sum(1 for s in selections if s.chosen_timeframe == "4h")

    summary = (f"Of {len(selections)} tickers: {n_daily} daily, {n_1h} 1h, {n_4h} 4h. "
               f"Anti-flip margin = {args.margin:.2f} Sharpe.")
    log.info(summary)

    report = SelectorReport(
        generated_at=datetime.now(UTC).isoformat(),
        n_tickers=len(selections),
        n_daily=n_daily, n_1h=n_1h, n_4h=n_4h,
        margin_threshold=args.margin,
        selections=selections,
        summary_note=summary,
    )
    payload = {k: v for k, v in asdict(report).items() if k != "selections"}
    payload["selections"] = [asdict(s) for s in report.selections]
    REPORT_PATH.write_text(
        json.dumps(payload, indent=2, ensure_ascii=False, default=str),
        encoding="utf-8",
    )
    log.info("Reporte exportado → %s", REPORT_PATH)

    sep = "=" * 64
    print(f"\n{sep}")
    print("Leonex -- Timeframe Selector")
    print(sep)
    print(summary)
    print()
    print(f"{'ticker':<10}{'chosen':<8}{'sharpe':<8}{'daily':<8}{'1h':<8}{'4h':<8}")
    for s in sorted(selections, key=lambda x: -x.chosen_sharpe):
        d = f"{s.sharpe_daily:+.2f}" if s.sharpe_daily is not None else "   -"
        h = f"{s.sharpe_1h:+.2f}" if s.sharpe_1h is not None else "   -"
        f4 = f"{s.sharpe_4h:+.2f}" if s.sharpe_4h is not None else "   -"
        print(f"  {s.ticker:<8}{s.chosen_timeframe:<8}{s.chosen_sharpe:+.2f}   "
              f"{d}   {h}   {f4}")
    print(sep)
    return 0


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