# src/models/calibration.py
import numpy as np
import pandas as pd
from scipy.optimize import minimize
from sklearn.preprocessing import RobustScaler
from sklearn.isotonic import IsotonicRegression
from sklearn.linear_model import LogisticRegression
from typing import Dict

from config import TEMP_BOUNDS, DRIFT_THRESHOLD


class TemperatureScaling:
    """Калибровка логитов через Temperature Scaling с L2-регуляризацией."""
    def __init__(self, bounds: tuple = TEMP_BOUNDS, default_temp: float = 1.0):
        self.temperature = default_temp
        self.bounds = bounds
        self.fitted = False

    def fit(self, logits: np.ndarray, labels: np.ndarray) -> None:
        if len(logits) < 30:
            self.fitted = False
            return

        def nll_loss(T: float) -> float:
            T_clipped = np.clip(T, *self.bounds)
            probs = 1.0 / (1.0 + np.exp(-logits / T_clipped))
            probs = np.clip(probs, 1e-7, 1.0 - 1e-7)
            nll = -np.mean(labels * np.log(probs) + (1.0 - labels) * np.log(1.0 - probs))
            reg = 0.05 * (T_clipped - 1.0) ** 2
            return nll + reg
        result = minimize(nll_loss, x0=[1.0], bounds=[self.bounds], method='L-BFGS-B')
        self.temperature = float(np.clip(result.x[0], *self.bounds))
        self.fitted = True


class PlattScaling:
    """Platt Scaling (sigmoid calibration)."""
    def __init__(self):
        self.model = None

    def fit(self, logits: np.ndarray, labels: np.ndarray) -> None:
        if len(logits) < 30:
            self.model = None
            return
        self.model = LogisticRegression(C=1.0, solver='lbfgs', max_iter=1000)
        self.model.fit(logits.reshape(-1, 1), labels)

    def transform(self, logits: np.ndarray) -> np.ndarray:
        if self.model is None:
            return 1.0 / (1.0 + np.exp(-logits))
        return self.model.predict_proba(logits.reshape(-1, 1))[:, 1]


class IsotonicCalibration:
    """Isotonic Regression calibration."""
    def __init__(self):
        self.model = None

    def fit(self, logits: np.ndarray, labels: np.ndarray) -> None:
        if len(logits) < 30:
            self.model = None
            return
        self.model = IsotonicRegression(y_min=0.01, y_max=0.99, out_of_bounds='clip')
        probs = 1.0 / (1.0 + np.exp(-logits))
        self.model.fit(probs, labels)

    def transform(self, logits: np.ndarray) -> np.ndarray:
        if self.model is None:
            return 1.0 / (1.0 + np.exp(-logits))
        probs = 1.0 / (1.0 + np.exp(-logits))
        return self.model.predict(probs)


class MultiCalibrator:
    """Комбинирует несколько методов калибровки."""
    def __init__(self, use_temp: bool = True, use_platt: bool = False, use_isotonic: bool = False):
        self.use_temp = use_temp
        self.use_platt = use_platt
        self.use_isotonic = use_isotonic
        self.temp = TemperatureScaling() if use_temp else None
        self.platt = PlattScaling() if use_platt else None
        self.isotonic = IsotonicCalibration() if use_isotonic else None

    def fit(self, logits: np.ndarray, labels: np.ndarray) -> None:
        if self.temp:
            self.temp.fit(logits, labels)
        if self.platt:
            self.platt.fit(logits, labels)
        if self.isotonic:
            self.isotonic.fit(logits, labels)

    def transform(self, logits: np.ndarray) -> np.ndarray:
        probs = []
        if self.temp and self.temp.fitted:
            probs.append(1.0 / (1.0 + np.exp(-logits / self.temp.temperature)))
        if self.platt:
            probs.append(self.platt.transform(logits))
        if self.isotonic:
            probs.append(self.isotonic.transform(logits))
        
        if not probs:
            probs.append(1.0 / (1.0 + np.exp(-logits)))
        
        return np.mean(probs, axis=0)


class DriftAdaptiveScaler:
    """Смешивает глобальный и локальный RobustScaler при дрейфе."""
    def __init__(self, global_scaler: RobustScaler, df_features: pd.DataFrame,
                 feat_cols: list, threshold: float = DRIFT_THRESHOLD, local_window: int = 50):
        self.global_scaler = global_scaler
        self.local_scaler = RobustScaler()
        self.threshold = threshold
        actual_window = min(local_window, len(df_features))
        self.local_scaler.fit(df_features.iloc[-actual_window:][feat_cols].values)

    def transform(self, X_latest: np.ndarray) -> np.ndarray:
        global_transformed = self.global_scaler.transform(X_latest)
        global_z = np.abs(global_transformed[-1])
        if global_z.max() > self.threshold:
            local_transformed = self.local_scaler.transform(X_latest)
            return 0.7 * local_transformed + 0.3 * global_transformed
        return global_transformed


class RobustUncertaintyEstimator:
    """Оценка неопределённости на основе дисперсии ансамбля."""
    def estimate(self, ensemble_probs_long: np.ndarray, ensemble_probs_short: np.ndarray) -> Dict:
        mean_long = float(np.mean(ensemble_probs_long))
        mean_short = float(np.mean(ensemble_probs_short))
        std_long = float(np.std(ensemble_probs_long))
        std_short = float(np.std(ensemble_probs_short))

        margin_long = max(0.06, std_long * 2.5 + 0.04)
        margin_short = max(0.06, std_short * 2.5 + 0.04)

        conf_long = 'HIGH' if (std_long < 0.04 and margin_long < 0.14) else \
                    'MEDIUM' if margin_long < 0.20 else 'LOW'
        conf_short = 'HIGH' if (std_short < 0.04 and margin_short < 0.14) else \
                     'MEDIUM' if margin_short < 0.20 else 'LOW'

        return {
            'prob_long': mean_long, 'margin_long': margin_long, 'std_long': std_long, 'conf_long': conf_long,
            'prob_short': mean_short, 'margin_short': margin_short, 'std_short': std_short, 'conf_short': conf_short
        }
