"""
Fractional Differentiation — Lopez de Prado AFML cap. 5.

PROBLEMA:
    Los precios crudos son no-estacionarios (tienen tendencia). Si los usas
    como feature en un modelo, los algoritmos asumen estacionariedad y
    fallan. Solucion clasica: diferenciar (return = P_t - P_{t-1}). Pero
    esto borra TODA la memoria de la serie — info temporal valiosa.

SOLUCION (Lopez de Prado):
    Diferenciacion fraccional con orden d ∈ (0, 1):
        X_tilde_t = sum_{k=0}^∞ w_k * X_{t-k}
        w_0 = 1
        w_k = -w_{k-1} * (d - k + 1) / k

    d=0: serie original (no estacionaria, toda la memoria)
    d=1: serie diferenciada (estacionaria, sin memoria)
    d∈(0,1): equilibrio. Tipicamente d~0.4 mantiene >80% de la memoria
             y consigue estacionariedad para la mayoria de assets.

IMPLEMENTACION:
    Fixed-width Window FFD: truncar pesos cuando |w_k| < tau (1e-5 default).
    Esto es ESTABLE en tiempo (mismo numero de lags para cada t) y mas
    rapido que la convolucion infinita teorica.

USO:
    weights = get_weights_ffd(d=0.4, tau=1e-5)
    diff_series = frac_diff_ffd(price_series, d=0.4)
    d_min = find_min_d_for_stationarity(price_series)   # busqueda binaria

Diseno: funciones puras, sin side-effects, faciles de testar.
"""

from __future__ import annotations

import math
from typing import Optional

import numpy as np
import pandas as pd


DEFAULT_TAU = 1e-5            # tolerancia para truncar pesos
DEFAULT_ADF_PVALUE = 0.05     # umbral p-value test ADF (5% nivel significacion)
DEFAULT_D_GRID = [0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0]
DEFAULT_MAX_LAGS = 500        # tope superior de pesos a evaluar


# ─────────────────────────────────────────────────────────────────────────────
# Pesos fraccionales
# ─────────────────────────────────────────────────────────────────────────────

def get_weights_ffd(d: float, tau: float = DEFAULT_TAU,
                    max_lags: int = DEFAULT_MAX_LAGS) -> np.ndarray:
    """Pesos para Fractional Differentiation con ventana fija (FFD).

    Itera la recursion w_k = -w_{k-1} * (d - k + 1) / k hasta que el peso
    absoluto cae por debajo de tau. Devuelve array shape (K,) con
    w_{K-1}, ..., w_1, w_0 en orden CRECIENTE de lag (es decir, el ultimo
    elemento corresponde a X_t y el primero al lag mas antiguo).

    Args:
        d: orden fraccional ∈ (0, 1) habitualmente. d=0 → solo w_0=1.
        tau: tolerancia para truncar (1e-5 default).
        max_lags: tope superior por seguridad numerica.

    Returns:
        np.ndarray con los pesos ordenados de lag mas antiguo a t actual.
    """
    if d == 0:
        return np.array([1.0])
    weights = [1.0]
    k = 1
    while k < max_lags:
        w = -weights[-1] * (d - k + 1) / k
        if abs(w) < tau:
            break
        weights.append(w)
        k += 1
    return np.array(weights[::-1])   # invertimos: lag mas antiguo primero


# ─────────────────────────────────────────────────────────────────────────────
# Aplicacion de FFD a una serie
# ─────────────────────────────────────────────────────────────────────────────

def frac_diff_ffd(series: pd.Series, d: float, tau: float = DEFAULT_TAU,
                  max_lags: int = DEFAULT_MAX_LAGS) -> pd.Series:
    """Aplica Fractional Differentiation con ventana fija a una serie.

    Args:
        series: pd.Series indexada por fecha, valores numericos.
        d: orden fraccional.
        tau: tolerancia para truncar pesos.

    Returns:
        pd.Series del mismo indice que series, con NaN en los primeros
        K-1 valores (donde K = len(weights), no hay suficiente historia).
    """
    weights = get_weights_ffd(d, tau, max_lags)
    width = len(weights) - 1
    n = len(series)
    if n <= width or width <= 0:
        if d == 0:
            return series.copy()
        return pd.Series(np.nan, index=series.index, dtype=float)

    output = pd.Series(np.nan, index=series.index, dtype=float)
    values = series.values
    for i in range(width, n):
        window = values[i - width: i + 1]   # K valores
        # producto interno con pesos (w_{lag_K-1}, ..., w_1, w_0)
        if np.any(np.isnan(window)):
            continue
        output.iloc[i] = float(np.dot(weights, window))
    return output


# ─────────────────────────────────────────────────────────────────────────────
# Test de estacionariedad (ADF — Augmented Dickey-Fuller)
# ─────────────────────────────────────────────────────────────────────────────

def adf_pvalue(series: pd.Series, max_lags: int = 12) -> Optional[float]:
    """P-value del test ADF. Devuelve None si insuficiente data o falla.

    Requiere statsmodels. Si no esta disponible, intenta una implementacion
    aproximada simple. La interpretacion: p < 0.05 → rechazamos hipotesis
    nula de raiz unitaria → serie ESTACIONARIA.
    """
    clean = series.dropna()
    if len(clean) < 30:
        return None
    try:
        from statsmodels.tsa.stattools import adfuller
        result = adfuller(clean.values, maxlag=max_lags, autolag="AIC")
        return float(result[1])
    except ImportError:
        # Fallback: regression-based ADF aprox. NO recomendado en produccion.
        # Calcula t-stat sobre primera diferencia y compara con tabla critica.
        # Esto es burdo, pero permite que el modulo funcione sin statsmodels.
        x = clean.values
        dx = np.diff(x)
        x_lag = x[:-1]
        if len(dx) < 30:
            return None
        mu_dx = np.mean(dx)
        mu_x = np.mean(x_lag)
        cov = np.sum((dx - mu_dx) * (x_lag - mu_x))
        var = np.sum((x_lag - mu_x) ** 2)
        if var <= 0:
            return None
        rho = cov / var
        # se_rho aprox
        residuals = dx - rho * x_lag
        se = math.sqrt(np.sum(residuals ** 2) / (len(dx) - 1))
        if se <= 0:
            return None
        t_stat = rho / (se / math.sqrt(var))
        # tabla critica MacKinnon aprox (5% level, no constant): -1.95
        # Devolvemos un p-value aproximado lineal entre [0, 1]
        # |t| > 1.95 → p ≈ 0.05; |t| > 2.86 → p ≈ 0.01; |t| < 1.0 → p ≈ 0.30
        abs_t = abs(t_stat)
        if abs_t < 1.0:
            return 0.30
        if abs_t < 1.95:
            return 0.10
        if abs_t < 2.86:
            return 0.05
        if abs_t < 3.43:
            return 0.01
        return 0.001


# ─────────────────────────────────────────────────────────────────────────────
# Busqueda del minimo d que estabiliza la serie
# ─────────────────────────────────────────────────────────────────────────────

def find_min_d_for_stationarity(
    series: pd.Series,
    d_grid: list[float] = DEFAULT_D_GRID,
    adf_threshold: float = DEFAULT_ADF_PVALUE,
    tau: float = DEFAULT_TAU,
) -> dict:
    """Busca el menor d en la grilla que hace la serie estacionaria.

    Args:
        series: pd.Series (precios o log-precios).
        d_grid: valores de d a probar (default 0.0 a 1.0 en pasos de 0.1).
        adf_threshold: si p_value <= threshold → ESTACIONARIA.

    Returns:
        dict con:
            "d_min": float | None         → menor d que pasa ADF
            "all_results": list[dict]      → (d, adf_pvalue, n_valid, correlation_with_original)
            "verdict": "PASS" | "FAIL"
            "note": str
    """
    results = []
    d_min = None
    for d in sorted(d_grid):
        diff = frac_diff_ffd(series, d, tau=tau)
        clean = diff.dropna()
        if len(clean) < 30:
            results.append({
                "d": float(d), "adf_pvalue": None,
                "n_valid": int(len(clean)),
                "correlation_with_original": None,
                "stationary": False,
            })
            continue
        p = adf_pvalue(clean)
        # Correlacion con la serie original (memoria preservada)
        # Solo donde ambos tienen valor
        s_orig = series.reindex(clean.index).dropna()
        common = clean.loc[s_orig.index] if not s_orig.empty else clean
        common_orig = s_orig.loc[common.index]
        if len(common) > 1 and common_orig.std() > 0:
            corr = float(common.corr(common_orig))
        else:
            corr = None
        is_stat = (p is not None and p <= adf_threshold)
        results.append({
            "d": float(d),
            "adf_pvalue": p,
            "n_valid": int(len(clean)),
            "correlation_with_original": corr,
            "stationary": bool(is_stat),
        })
        if is_stat and d_min is None:
            d_min = float(d)

    if d_min is not None:
        # Buscar la fila correspondiente
        row = next((r for r in results if r["d"] == d_min), {})
        note = (f"d_min={d_min:.1f} hace la serie estacionaria "
                f"(ADF p={row.get('adf_pvalue', 0):.4f}). "
                f"Correlacion con original: {row.get('correlation_with_original', 0):.3f}.")
        return {
            "d_min": d_min,
            "all_results": results,
            "verdict": "PASS",
            "note": note,
        }
    return {
        "d_min": None,
        "all_results": results,
        "verdict": "FAIL",
        "note": "Ningun d en el grid logro estacionariedad. Considera ampliar el grid o revisar la serie.",
    }
