"""
Agente de Régimen de Leonex.

Lee datos de SQLite, calcula indicadores (ADX, ATR, SMA) y determina
el régimen de mercado (Tendencia, Lateral, Crisis) para cada activo.
"""

from __future__ import annotations

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 Iterable

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"
REGIMES_SNAPSHOT_PATH = DASHBOARD_DATA_DIR / "regimes_snapshot.json"


@dataclass(frozen=True)
class RegimeSnapshot:
    ticker: str
    date: str
    adx: float
    atr: float
    sma_50: float
    close: float
    regime: str


def wilders_smoothing(s: pd.Series, n: int) -> pd.Series:
    """Suavizado de Wilder usado para ADX y ATR."""
    return s.ewm(alpha=1/n, adjust=False).mean()


def calculate_indicators(df: pd.DataFrame, n: int = 14) -> pd.DataFrame:
    """Calcula ADX, ATR y SMAs."""
    df = df.copy()
    up = df['high'] - df['high'].shift(1)
    down = df['low'].shift(1) - df['low']
    
    plus_dm = pd.Series(np.where((up > down) & (up > 0), up, 0.0), index=df.index)
    minus_dm = pd.Series(np.where((down > up) & (down > 0), down, 0.0), index=df.index)
    
    tr1 = df['high'] - df['low']
    tr2 = (df['high'] - df['close'].shift(1)).abs()
    tr3 = (df['low'] - df['close'].shift(1)).abs()
    tr = pd.concat([tr1, tr2, tr3], axis=1).max(axis=1)
    
    df['atr'] = wilders_smoothing(tr, n)
    plus_di = 100 * wilders_smoothing(plus_dm, n) / df['atr']
    minus_di = 100 * wilders_smoothing(minus_dm, n) / df['atr']
    
    dx = 100 * (plus_di - minus_di).abs() / (plus_di + minus_di)
    df['adx'] = wilders_smoothing(dx, n)
    
    df['sma_50'] = df['close'].rolling(window=50).mean()
    df['sma_200'] = df['close'].rolling(window=200).mean()
    
    # Promedio historico de ATR para detectar crisis (volatilidad anormal)
    df['atr_mean_60'] = df['atr'].rolling(window=60).mean()
    
    return df


def classify_regime(row: pd.Series) -> str:
    """Clasifica el regimen basándose en los indicadores."""
    if pd.isna(row['adx']) or pd.isna(row['sma_50']):
        return "DESCONOCIDO"
        
    # Detectar Crisis: Si la volatilidad es más del doble del promedio reciente
    if not pd.isna(row['atr_mean_60']) and row['atr'] > 2.0 * row['atr_mean_60']:
        return "CRISIS"
        
    if row['adx'] > 25:
        if row['close'] > row['sma_50']:
            return "TENDENCIA_ALCISTA"
        else:
            return "TENDENCIA_BAJISTA"
    else:
        return "LATERAL"


class AgenteRegimen:
    def __init__(self, db_path: Path = DB_PATH, snapshot_path: Path = REGIMES_SNAPSHOT_PATH) -> None:
        self.db_path = db_path
        self.snapshot_path = snapshot_path
        self.logger = self._build_logger()
        self._ensure_directories()
        self._ensure_schema()

    def run(self) -> list[RegimeSnapshot]:
        self.logger.info("Iniciando Agente de Régimen")
        snapshots: list[RegimeSnapshot] = []
        
        with sqlite3.connect(self.db_path) as conn:
            # Obtener todos los tickers
            tickers_df = pd.read_sql("SELECT DISTINCT ticker FROM prices", conn)
            tickers = tickers_df['ticker'].tolist()
            
            for ticker in tickers:
                # Leer precios historicos
                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("No hay suficientes datos para %s (%d filas)", ticker, len(prices))
                    continue
                    
                # Calcular indicadores
                prices = calculate_indicators(prices)
                
                # Clasificar regímenes
                prices['regime'] = prices.apply(classify_regime, axis=1)
                
                # Guardar en SQLite
                self.save_regimes(ticker, prices)
                
                # Tomar snapshot del ultimo dia
                last_row = prices.iloc[-1]
                snapshot = RegimeSnapshot(
                    ticker=ticker,
                    date=prices.index[-1].strftime("%Y-%m-%d"),
                    adx=float(last_row['adx']) if not pd.isna(last_row['adx']) else 0.0,
                    atr=float(last_row['atr']) if not pd.isna(last_row['atr']) else 0.0,
                    sma_50=float(last_row['sma_50']) if not pd.isna(last_row['sma_50']) else 0.0,
                    close=float(last_row['close']),
                    regime=str(last_row['regime'])
                )
                snapshots.append(snapshot)
                self.logger.info("%s clasificado como %s (ADX: %.1f)", ticker, snapshot.regime, snapshot.adx)

        self.export_snapshot(snapshots)
        self.logger.info("Agente de Régimen terminado. Activos procesados=%s", len(snapshots))
        return snapshots

    def save_regimes(self, ticker: str, df: pd.DataFrame) -> None:
        """Guarda la serie histórica de regímenes en SQLite."""
        rows = df.reset_index(names="date").copy()
        rows["ticker"] = ticker
        rows["date"] = rows["date"].dt.strftime("%Y-%m-%d")
        
        # Filtrar NaN para la base de datos
        rows = rows.fillna(0.0)

        records = rows[
            [
                "ticker", "date", "adx", "atr", "sma_50", "sma_200", "regime"
            ]
        ].to_records(index=False)

        with sqlite3.connect(self.db_path) as conn:
            conn.executemany(
                """
                INSERT OR REPLACE INTO regimes (
                    ticker, date, adx, atr, sma_50, sma_200, regime
                )
                VALUES (?, ?, ?, ?, ?, ?, ?)
                """,
                records,
            )
            conn.commit()

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

    def _ensure_directories(self) -> None:
        DASHBOARD_DATA_DIR.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 regimes (
                    ticker TEXT NOT NULL,
                    date TEXT NOT NULL,
                    adx REAL,
                    atr REAL,
                    sma_50 REAL,
                    sma_200 REAL,
                    regime TEXT NOT NULL,
                    PRIMARY KEY (ticker, date)
                )
                """
            )
            conn.commit()

    def _build_logger(self) -> logging.Logger:
        LOGS_DIR.mkdir(parents=True, exist_ok=True)
        logger = logging.getLogger("agente_regimen")
        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_regimen.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:
    snapshots = AgenteRegimen().run()
    print(f"\nRegímenes actualizados: {len(snapshots)}")
    print(f"Base de datos Leonex: {DB_PATH}")
    print(f"Snapshot dashboard: {REGIMES_SNAPSHOT_PATH}")


if __name__ == "__main__":
    main()
