"""
Сравнение эффективности NN-ансамбля (3 модели) против скриптового анализа.

Запуск:
    python src/ml/inference/compare_nn_vs_script.py
"""

import json
import sys
import warnings
from pathlib import Path
from typing import Dict, List, Optional

import numpy as np
import torch

warnings.filterwarnings('ignore')

project_root = Path(__file__).resolve().parent.parent.parent.parent
sys.path.insert(0, str(project_root))

from src.db.connection import fetch_ohlcv_combined
from src.ml.features.price_features import add_price_features
from src.ml.features.indicator_features import add_indicator_features
from src.ml.features.calendar_features import add_calendar_features
from src.ml.models.lstm import GRUPredictor, LSTMPredictor

TICKER = 'X5'
CLASS_NAMES = {0: 'HOLD', 1: 'BUY', 2: 'SELL'}

# Все 3 модели ансамбля
ENSEMBLE_MODELS = {
    'x5_gru_v3': {
        'class': GRUPredictor,
        'weight': 1.0,
        'tf': 'D1',
    },
    'x5_lstm_v3_focal': {
        'class': LSTMPredictor,
        'weight': 1.0,
        'tf': 'D1',
    },
    'x5_gru_v4': {
        'class': GRUPredictor,
        'weight': 1.0,
        'tf': 'D1',
    }
}


def calc_rsi(prices, period=14):
    if len(prices) < period + 1:
        return 50.0
    deltas = np.diff(prices[-period-1:])
    gains = np.where(deltas > 0, deltas, 0)
    losses = np.where(deltas < 0, -deltas, 0)
    avg_gain = np.mean(gains)
    avg_loss = np.mean(losses)
    if avg_loss == 0:
        return 100.0
    rs = avg_gain / avg_loss
    return 100.0 - (100.0 / (1.0 + rs))


def calc_ema(prices, period):
    if len(prices) < period:
        return float(np.mean(prices))
    alpha = 2.0 / (period + 1)
    ema = float(np.mean(prices[:period]))
    for p in prices[period:]:
        ema = alpha * p + (1 - alpha) * ema
    return ema


def load_model(model_name: str, model_class, config_path: Path,
               checkpoint_path: Path, device: torch.device) -> Optional[Dict]:
    """Загрузить модель."""
    try:
        if not config_path.exists() or not checkpoint_path.exists():
            return None
        config = json.loads(config_path.read_text())
        params = config.get('params', {})
        model = model_class(
            input_size=params.get('input_size', 61),
            hidden_size=params.get('hidden_size', 96),
            num_layers=params.get('num_layers', 2),
            dropout=0.0,
        )
        state_dict = torch.load(checkpoint_path, map_location=device, weights_only=True)
        model.load_state_dict(state_dict)
        model = model.to(device)
        model.eval()

        norm_path = config_path.parent / 'normalizer.json'
        normalizer = None
        if norm_path.exists():
            nd = json.loads(norm_path.read_text())
            normalizer = {'mean': np.array(nd['mean_']), 'std': np.array(nd['std_'])}

        thresholds = config.get('calibrated_thresholds', {0: 0.0, 1: 0.4, 2: 0.4})
        return {
            'model': model, 'config': config,
            'normalizer': normalizer, 'thresholds': thresholds,
            'feature_cols': config.get('feature_cols', []),
            'weight': 1.0,
        }
    except Exception as e:
        print(f"  ❌ Ошибка загрузки {model_name}: {e}")
        return None


def ensemble_predict(models_loaded: Dict, X_clean: np.ndarray, df_clean, seq_len: int = 60):
    """Прогноз ансамбля для каждого бара."""
    ensemble_signals = []
    
    for i in range(seq_len, len(X_clean)):
        # Получить прогноз от каждой модели
        avg_probs = np.zeros(3)
        total_weight = 0
        
        for name, artifacts in models_loaded.items():
            norm = artifacts['normalizer']
            thresholds = artifacts['thresholds']
            
            # Нормализация
            X_norm = np.zeros_like(X_clean)
            for j in range(X_clean.shape[1]):
                if norm['std'][j] > 0:
                    X_norm[:, j] = (X_clean[:, j] - norm['mean'][j]) / norm['std'][j]
            
            X_seq = torch.FloatTensor(X_norm[i - seq_len:i]).unsqueeze(0)
            with torch.no_grad():
                output = artifacts['model'](X_seq)
                probs = torch.softmax(output, dim=1).squeeze().cpu().numpy()
            
            avg_probs += probs * artifacts['weight']
            total_weight += artifacts['weight']
        
        avg_probs /= total_weight
        
        pred_class = int(avg_probs.argmax())
        threshold = 0.4 if pred_class in [1, 2] else 0.0
        if avg_probs[pred_class] < threshold:
            pred_class = 0
        
        ensemble_signals.append({
            'date': df_clean['Date'].iloc[i],
            'close': float(df_clean['Close'].iloc[i]),
            'nn_signal': CLASS_NAMES[pred_class],
            'nn_conf': float(avg_probs[pred_class]),
            'probs': [float(p) for p in avg_probs],
        })
    
    return ensemble_signals


def main():
    print("=" * 70)
    print("СРАВНЕНИЕ: NN-АНСАМБЛЬ (3 МОДЕЛИ) vs СКРИПТОВЫЙ АНАЛИЗ")
    print("=" * 70)

    # 1. Load data
    df = fetch_ohlcv_combined(TICKER, 'D1', limit=400)
    if df is None or len(df) < 100:
        print("ERROR: Not enough data")
        return

    print(f"Данных: {len(df)} баров, {df['Date'].iloc[0]} — {df['Date'].iloc[-1]}")

    # 2. Add features
    df_feat = add_price_features(df.copy())
    df_feat = add_indicator_features(df_feat)
    df_feat = add_calendar_features(df_feat)

    # 3. Load all 3 ensemble models
    print("\n--- Загрузка ансамбля (3 модели) ---")
    device = torch.device('cpu')

    models_loaded = {}
    for name, info in ENSEMBLE_MODELS.items():
        path = Path('src/ml/models/saved') / name
        arts = load_model(name, info['class'], path / 'config.json',
                          path / 'model.pt', device)
        if arts:
            arts['weight'] = info['weight']
            models_loaded[name] = arts
            print(f"  ✅ {name} — загружена (вес={info['weight']})")

    if not models_loaded:
        print("ERROR: Не загружена ни одна модель!")
        return

    print(f"\n  Всего моделей в ансамбле: {len(models_loaded)}")

    # 4. Prepare feature matrix
    # Используем фичи первой модели (у всех моделей они одинаковые)
    first_model = list(models_loaded.values())[0]
    feature_cols = first_model['feature_cols']
    
    X_full = df_feat[feature_cols].values.astype(np.float64)
    nan_mask = np.isnan(X_full).any(axis=1)
    first_valid = np.where(~nan_mask)[0][0]
    X_clean = X_full[first_valid:]
    df_clean = df_feat.iloc[first_valid:].reset_index(drop=True)
    print(f"  Чистых баров (после удаления NaN): {len(X_clean)}")

    seq_len = 60

    # 5. Generate ensemble NN signals
    print("\n--- Генерация NN-сигналов (ансамбль) ---")
    nn_signals = ensemble_predict(models_loaded, X_clean, df_clean, seq_len)
    print(f"  Сгенерировано {len(nn_signals)} NN-сигналов")

    # 6. Generate script-based signals
    print("\n--- Генерация скриптовых сигналов ---")
    closes = df_clean['Close'].values
    volumes = df_clean['Volume'].values

    script_signals = []
    for i in range(seq_len, len(df_clean)):
        close = closes[i]
        rsi_14 = calc_rsi(closes[:i+1], 14)
        rsi_7 = calc_rsi(closes[:i+1], 7)

        ema_9 = calc_ema(closes[:i+1], min(9, i+1))
        ema_21 = calc_ema(closes[:i+1], min(21, i+1))
        ema_50 = calc_ema(closes[:i+1], min(50, i+1))

        trend = 0
        reasons = []

        # RSA-based
        if rsi_14 < 30:
            trend = 0
            reasons.append(f'RSI14={rsi_14:.0f}<30→HOLD')
        elif rsi_14 > 70:
            trend = 2
            reasons.append(f'RSI14={rsi_14:.0f}>70→SELL')
        elif rsi_7 < 30 and rsi_14 < 40:
            trend = 0
            reasons.append(f'RSI7={rsi_7:.0f}<30+RSI14={rsi_14:.0f}<40→HOLD')
        elif rsi_7 > 70 and rsi_14 > 60:
            trend = 2
            reasons.append(f'RSI7={rsi_7:.0f}>70+RSI14={rsi_14:.0f}>60→SELL')

        # EMA trend
        if ema_9 < ema_21 < ema_50:
            trend = 2
            reasons.append('EMA_bear')
        elif ema_9 > ema_21 > ema_50:
            if trend != 2:
                trend = 1
                reasons.append('EMA_bull')

        # Volume
        vol_ratio = volumes[i] / np.mean(volumes[max(0, i-20):i+1]) if i >= 20 else 1.0
        if vol_ratio > 2.0 and close < closes[i-1]:
            if trend != 1:
                trend = 2
                reasons.append('HighVolDrop')
        elif vol_ratio > 2.0 and close > closes[i-1]:
            if trend != 2:
                trend = 1
                reasons.append('HighVolRise')

        script_signals.append({
            'date': df_clean['Date'].iloc[i],
            'close': float(close),
            'script_signal': ['HOLD', 'BUY', 'SELL'][trend],
            'rsi_14': float(rsi_14),
            'rsi_7': float(rsi_7),
            'reason': '; '.join(reasons) if reasons else 'no_signal',
        })
    print(f"  Сгенерировано {len(script_signals)} скриптовых сигналов")

    # 7. Compare
    print("\n" + "=" * 70)
    print("СРАВНЕНИЕ СИГНАЛОВ")
    print("=" * 70)

    min_len = min(len(nn_signals), len(script_signals))
    nn_s = nn_signals[-min_len:]
    sc_s = script_signals[-min_len:]

    # Statistics
    agree = 0
    nn_buy_sc_sell = 0
    nn_sell_sc_buy = 0
    nn_buy_sc_hold = 0
    nn_sell_sc_hold = 0
    nn_hold_sc_buy = 0
    nn_hold_sc_sell = 0

    nn_trades = {'BUY': 0, 'SELL': 0, 'HOLD': 0}
    sc_trades = {'BUY': 0, 'SELL': 0, 'HOLD': 0}

    for n, s in zip(nn_s, sc_s):
        nn_trades[n['nn_signal']] += 1
        sc_trades[s['script_signal']] += 1

        if n['nn_signal'] == s['script_signal']:
            agree += 1
        elif n['nn_signal'] == 'BUY' and s['script_signal'] == 'SELL':
            nn_buy_sc_sell += 1
        elif n['nn_signal'] == 'SELL' and s['script_signal'] == 'BUY':
            nn_sell_sc_buy += 1
        elif n['nn_signal'] == 'BUY' and s['script_signal'] == 'HOLD':
            nn_buy_sc_hold += 1
        elif n['nn_signal'] == 'SELL' and s['script_signal'] == 'HOLD':
            nn_sell_sc_hold += 1
        elif n['nn_signal'] == 'HOLD' and s['script_signal'] == 'BUY':
            nn_hold_sc_buy += 1
        elif n['nn_signal'] == 'HOLD' and s['script_signal'] == 'SELL':
            nn_hold_sc_sell += 1

    print(f"\nВсего сигналов: {min_len}")
    print(f"Совпадений (NN = Script): {agree} ({agree/min_len*100:.1f}%)")
    print(f"NN=BUY vs Script=SELL: {nn_buy_sc_sell} ({nn_buy_sc_sell/min_len*100:.1f}%)")
    print(f"NN=SELL vs Script=BUY: {nn_sell_sc_buy} ({nn_sell_sc_buy/min_len*100:.1f}%)")
    print(f"NN=BUY vs Script=HOLD: {nn_buy_sc_hold} ({nn_buy_sc_hold/min_len*100:.1f}%)")
    print(f"NN=SELL vs Script=HOLD: {nn_sell_sc_hold} ({nn_sell_sc_hold/min_len*100:.1f}%)")
    print(f"NN=HOLD vs Script=BUY: {nn_hold_sc_buy} ({nn_hold_sc_buy/min_len*100:.1f}%)")
    print(f"NN=HOLD vs Script=SELL: {nn_hold_sc_sell} ({nn_hold_sc_sell/min_len*100:.1f}%)")

    print(f"\nРаспределение NN-ансамбль:  BUY={nn_trades['BUY']}  HOLD={nn_trades['HOLD']}  SELL={nn_trades['SELL']}")
    print(f"Распределение Script:       BUY={sc_trades['BUY']}  HOLD={sc_trades['HOLD']}  SELL={sc_trades['SELL']}")

    # 8. Compare on last 60 bars
    print("\n" + "=" * 70)
    print("СРАВНЕНИЕ НА ПОСЛЕДНИХ 60 БАРАХ (ретроспектива)")
    print("=" * 70)

    lookback = 60
    nn_recent = nn_s[-lookback:]
    sc_recent = sc_s[-lookback:]

    print(f"\n{'Date':<14} {'Close':>8} {'NN(ens)':>8} {'NN%':>5} {'Script':>8} {'RSI14':>6}")
    print("-" * 55)

    for n, s in zip(nn_recent, sc_recent):
        print(f"{n['date']:<14} {n['close']:>8.0f} {n['nn_signal']:>8} "
              f"{n['nn_conf']*100:>4.0f}% {s['script_signal']:>8} "
              f"{s['rsi_14']:>5.0f}")

    # 9. Evaluate accuracy
    print("\n" + "=" * 70)
    print("ОЦЕНКА ТОЧНОСТИ ПРЕДСКАЗАНИЙ")
    print("=" * 70)

    def eval_accuracy(signals, price_col='close', signal_col='nn_signal'):
        correct = 0
        total = 0
        for idx in range(1, len(signals)):
            sig = signals[idx - 1][signal_col]
            if sig == 'HOLD':
                continue
            total += 1
            next_close = signals[idx]['close']
            curr_close = signals[idx - 1]['close']
            actual_up = next_close > curr_close
            predicted_up = (sig == 'BUY')
            if actual_up == predicted_up:
                correct += 1
        return correct, total, correct / total * 100 if total > 0 else 0

    nn_correct, nn_total, nn_acc = eval_accuracy(nn_s, signal_col='nn_signal')
    sc_correct, sc_total, sc_acc = eval_accuracy(sc_s, signal_col='script_signal')

    print(f"\nНЕЙРОСЕТЬ (ансамбль 3 моделей):")
    print(f"  Торговых сигналов (не HOLD): {nn_total}")
    print(f"  Правильных: {nn_correct} ({nn_acc:.1f}%)")
    print(f"  Неправильных: {nn_total - nn_correct} ({(nn_total-nn_correct)/nn_total*100:.1f}%)" if nn_total > 0 else "  Нет сигналов")

    print(f"\nСКРИПТОВЫЙ АНАЛИЗ:")
    print(f"  Торговых сигналов (не HOLD): {sc_total}")
    print(f"  Правильных: {sc_correct} ({sc_acc:.1f}%)")
    print(f"  Неправильных: {sc_total - sc_correct} ({(sc_total-sc_correct)/sc_total*100:.1f}%)" if sc_total > 0 else "  Нет сигналов")

    if nn_total > 0 and sc_total > 0:
        diff = nn_acc - sc_acc
        better = "НЕЙРОСЕТЬ" if diff > 0 else "СКРИПТ"
        print(f"\n{'=' * 40}")
        print(f"ВЕРДИКТ: {better} точнее на {abs(diff):.1f}%")
        print(f"{'=' * 40}")

    # 10. Last 10 days detail
    print("\n" + "=" * 70)
    print(f"СРАВНЕНИЕ СИГНАЛОВ ЗА {min(10, len(nn_s))} ДНЕЙ")
    print("=" * 70)

    last_n = min(10, len(nn_s))
    last_10_nn = nn_s[-last_n:]
    last_10_sc = sc_s[-last_n:]

    print(f"\n{'Date':<12} {'Close':>7} {'Chg%':>7} {'NN_ens':>7} {'NN%':>7} {'Sc_sig':>8} {'Sc_rsi':>7}")
    print("-" * 55)
    for i in range(len(last_10_nn)):
        n = last_10_nn[i]
        s = last_10_sc[i]
        chg = 0
        if i > 0:
            chg = (n['close'] - last_10_nn[i-1]['close']) / last_10_nn[i-1]['close'] * 100
        print(f"{n['date']:<12} {n['close']:>7.0f} {chg:>+6.2f}% "
              f"{n['nn_signal']:>7} {n['nn_conf']*100:>5.0f}% "
              f"{s['script_signal']:>8} {s['rsi_14']:>5.0f}")

    # 11. Summary
    print("\n" + "=" * 70)
    print("СВОДКА СРАВНЕНИЯ")
    print("=" * 70)
    
    nn_models_detail = []
    for name in ENSEMBLE_MODELS:
        if name in models_loaded:
            nn_models_detail.append(name)
    
    print(f"""
┌──────────────────────────────┬────────────────┬────────────┐
│ Показатель                   │ NN-ансамбль    │ Скрипт     │
│                              │ ({len(nn_models_detail)} модели)  │ _analyze   │
├──────────────────────────────┼────────────────┼────────────┤
│ Точность (не-HOLD сигналы)   │ {nn_acc:>5.1f}%          │ {sc_acc:>5.1f}%      │
│ Сигналов (не HOLD)           │ {nn_total:>5}            │ {sc_total:>5}        │
│ SELL-сигналов                │ {nn_trades['SELL']:>5}            │ {sc_trades['SELL']:>5}        │
│ BUY-сигналов                 │ {nn_trades['BUY']:>5}            │ {sc_trades['BUY']:>5}        │
│ HOLD-сигналов                │ {nn_trades['HOLD']:>5}            │ {sc_trades['HOLD']:>5}        │
│ Совпадений (NN=Script)       │ {agree:>5} ({agree/min_len*100:.0f}%)         │ —          │
│ Противоположных (NN≠Script)  │ {min_len-agree:>5} ({(min_len-agree)/min_len*100:.0f}%)         │ —          │
├──────────────────────────────┼────────────────┼────────────┤
│ Модели в ансамбле            │ {', '.join(nn_models_detail)} │ N/A (rules)│
│ SELL recall (тест, среднее)  │ ~44%           │ N/A        │
│ Параметры (суммарно)         │ ~527K          │ N/A        │
└──────────────────────────────┴────────────────┴────────────┘
""")

    # Save results
    results = {
        'comparison': {
            'total_signals': min_len,
            'agreement': agree,
            'agreement_pct': round(agree / min_len * 100, 1),
            'nn_accuracy': round(nn_acc, 1),
            'script_accuracy': round(sc_acc, 1),
            'nn_trades': nn_trades,
            'script_trades': sc_trades,
            'models_in_ensemble': nn_models_detail,
        },
        'test_metrics': {
            'nn_models': {
                'x5_gru_v3': {'test_acc': 48.48, 'calibrated_acc': 54.55, 'sell_recall': 50.0, 'params': 108067},
                'x5_lstm_v3_focal': {'test_acc': 54.55, 'calibrated_acc': 54.55, 'sell_recall': 40.0, 'params': 238483},
                'x5_gru_v4': {'test_acc': 57.58, 'calibrated_acc': 54.55, 'sell_recall': 41.7, 'params': 180867},
            },
            'script_type': 'rule-based (RSI+EMA+Volume)',
        }
    }

    output_path = Path('/tmp/x5_nn_vs_script_comparison.json')
    output_path.write_text(json.dumps(results, indent=2, ensure_ascii=False))
    print(f"\nРезультаты сохранены: {output_path}")


if __name__ == '__main__':
    main()
