"""
Ансамбль X5 v2 — 4 модели (включая BiGRU+Attention).

Модели:
    ✅ x5_gru_v3           — GRU, Focal γ=2, SELL recall 50%
    ✅ x5_lstm_v3_focal    — LSTM, Focal γ=2, SELL recall 40%
    ✅ x5_gru_v4           — GRU, Focal γ=3, SELL recall 42%
    ✅ x5_bigru_attn_v1    — BiGRU+Attention, SELL recall 30% (НОВЫЙ!)

Usage:
    python src/ml/inference/ensemble_predict_x5_v2.py
"""

import json
import logging
import sys
from pathlib import Path
from typing import Dict, List, Optional, Tuple

import numpy as np
import torch

logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s [%(levelname)s] %(message)s',
    datefmt='%H:%M:%S',
)
logger = logging.getLogger(__name__)

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

from src.ml.models.lstm import GRUPredictor, LSTMPredictor
from src.ml.models.gru_attention import BiGRUAttention
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.db.connection import fetch_ohlcv_combined

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

# Все 4 модели — без оверфиттинга, SELL recall > 0%
TRUSTED_MODELS = {
    'x5_gru_v3': {
        'class': GRUPredictor,
        'weight': 1.0,
        'tf': 'D1',
        'seq_len': 60,
    },
    'x5_lstm_v3_focal': {
        'class': LSTMPredictor,
        'weight': 1.0,
        'tf': 'D1',
        'seq_len': 60,
    },
    'x5_gru_v4': {
        'class': GRUPredictor,
        'weight': 1.0,
        'tf': 'D1',
        'seq_len': 60,
    },
    'x5_bigru_attn_v1': {
        'class': BiGRUAttention,
        'weight': 1.0,
        'tf': 'D1',
        'seq_len': 60,
    },
}

# De-duplicated: which columns from df to use as features (when feature_cols missing)
ALL_FEATURE_COLS = [
    'Open', 'High', 'Low', 'Close', 'Volume',
    'ret_1', 'ret_5', 'ret_10', 'ret_21',
    'log_ret_1', 'high_low_ratio', 'close_open_ratio',
    'upper_shadow', 'lower_shadow', 'spread_pct', 'close_position',
    'ema_9', 'ema_21', 'ema_50', 'sma_20', 'sma_50', 'sma_200',
    'price_to_ema9', 'price_to_sma50', 'price_to_sma200',
    'macd', 'macd_signal', 'macd_hist', 'macd_hist_pct',
    'adx', 'plus_di', 'minus_di',
    'rsi_14', 'rsi_7', 'stoch_k', 'stoch_d', 'cci_20', 'mfi_14', 'mfi_7',
    'bb_lower', 'bb_middle', 'bb_upper', 'bb_width', 'bb_position',
    'atr_14', 'atr_pct',
    'obv', 'volume_sma_20', 'volume_ratio', 'volume_change', 'volume_change_5',
    'day_of_week', 'month', 'quarter', 'day_of_month',
    'is_month_end', 'is_quarter_end',
    'day_sin', 'day_cos', 'month_sin', 'month_cos',
]


def load_model(model_name: str, model_class, config_path: Path,
               checkpoint_path: Path, device: torch.device, seq_len: int = 60) -> Optional[Dict]:
    """Загрузить модель — гибко для разных форматов конфига."""
    try:
        if not config_path.exists() or not checkpoint_path.exists():
            return None

        config = json.loads(config_path.read_text())

        # Определяем архитектурные параметры (разные форматы конфигов)
        if model_class == BiGRUAttention:
            arch = {
                'input_size': config.get('input_size', 61),
                'hidden_size': config.get('hidden_size', 64),
                'num_layers': config.get('num_layers', 2),
                'dropout': 0.0,
            }
        else:
            params = config.get('params', config)
            arch = {
                'input_size': params.get('input_size', 61),
                'hidden_size': params.get('hidden_size', 96),
                'num_layers': params.get('num_layers', 2),
                'dropout': 0.0,
            }

        model = model_class(**arch)
        state_dict = torch.load(checkpoint_path, map_location=device, weights_only=True)
        model.load_state_dict(state_dict)
        model = model.to(device)
        model.eval()

        # Пороги
        thresholds = config.get('calibrated_thresholds', {0: 0.0, 1: 0.4, 2: 0.4})
        # Для BiGRU thresholds могут быть {1: t, 2: t}, а не {0:0, 1:t, 2:t}
        if '0' not in thresholds and '1' in thresholds:
            fixed = {0: 0.0}
            fixed.update(thresholds)
            thresholds = fixed

        # Feature columns
        feature_cols = config.get('feature_cols', [])

        return {
            'model': model, 'config': config,
            'thresholds': thresholds,
            'feature_cols': feature_cols,
            'seq_len': seq_len,
        }
    except Exception as e:
        logger.error(f"Ошибка загрузки {model_name}: {e}")
        return None


def prepare_features(df, feature_cols):
    """Подготовить фичи — использовать feature_cols или все стандартные."""
    if feature_cols:
        cols = [c for c in feature_cols if c in df.columns]
    else:
        cols = [c for c in ALL_FEATURE_COLS if c in df.columns]

    X_full = df[cols].values.astype(np.float64)
    nan_mask = np.isnan(X_full).any(axis=1)
    if nan_mask.any():
        first_valid = np.where(~nan_mask)[0][0]
        X_full = X_full[first_valid:]
    X = X_full[-150:] if len(X_full) > 150 else X_full
    return X


def predict_one(model_artifacts, X, device):
    """Inference одной модели."""
    model = model_artifacts['model']
    thresholds = model_artifacts.get('thresholds', {0: 0.0, 1: 0.4, 2: 0.4})
    seq_len = model_artifacts.get('seq_len', 60)

    if len(X) < seq_len:
        pad = np.zeros((seq_len - len(X), X.shape[1]))
        X_seq = np.concatenate([pad, X], axis=0)
    else:
        X_seq = X[-seq_len:]

    X_tensor = torch.FloatTensor(X_seq).unsqueeze(0).to(device)
    with torch.no_grad():
        output = model(X_tensor)
        probs = torch.softmax(output, dim=1).squeeze().cpu().numpy()

    pred_class = int(probs.argmax())
    if pred_class in [1, 2] and probs[pred_class] < thresholds.get(pred_class, 0.0):
        pred_class = 0
    return pred_class, probs


def main():
    logger.info("=" * 70)
    logger.info("АНСАМБЛЬ X5 v2 — 4 МОДЕЛИ (С BiGRU+ATTENTION)")
    logger.info("=" * 70)
    logger.info(f"  Модели: x5_gru_v3 (SELL 50%) + x5_lstm_v3_focal (SELL 40%) + "
                f"x5_gru_v4 (SELL 42%) + x5_bigru_attn_v1 (SELL 30%)")

    device = torch.device('cpu')
    df = fetch_ohlcv_combined(TICKER, 'D1', limit=400)
    if df is None or len(df) < 100:
        logger.error("Недостаточно данных")
        return
    logger.info(f"  Данных: {len(df)} баров")

    df_feat = add_price_features(df.copy())
    df_feat = add_indicator_features(df_feat)
    df_feat = add_calendar_features(df_feat)
    logger.info(f"  Фич: {len(df_feat.columns)}")

    # Загружаем все 4 модели
    models_loaded = {}
    for name, info in TRUSTED_MODELS.items():
        path = Path('src/ml/models/saved') / name
        arts = load_model(name, info['class'], path / 'config.json',
                          path / 'model.pt', device, seq_len=info['seq_len'])
        if arts:
            models_loaded[name] = {**arts, 'weight': info['weight'], 'tf': info['tf']}
            logger.info(f"  ✅ {name} — загружена")
        else:
            logger.warning(f"  ❌ {name} — НЕ загружена")

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

    # Инференс
    logger.info("\n" + "-" * 50)
    logger.info("РЕЗУЛЬТАТЫ ИНФЕРЕНСА")
    logger.info("-" * 50)

    predictions = []
    for name, artifacts in models_loaded.items():
        feature_cols = artifacts['feature_cols']
        X = prepare_features(df_feat, feature_cols)
        if X is None or len(X) < 10:
            logger.warning(f"  {name}: недостаточно данных")
            continue

        pred_class, probs = predict_one(artifacts, X, device)
        predictions.append((name, pred_class, probs, artifacts['weight']))

        probs_pct = {CLASS_NAMES[i]: f"{p*100:.1f}%" for i, p in enumerate(probs)}
        logger.info(f"  {name:25s} → {CLASS_NAMES[pred_class]:>5} | "
                    f"HOLD={probs_pct['HOLD']} BUY={probs_pct['BUY']} SELL={probs_pct['SELL']}")

    # Финальный сигнал — взвешенное усреднение вероятностей
    logger.info("\n" + "-" * 50)
    logger.info("ФИНАЛЬНЫЙ СИГНАЛ")
    logger.info("-" * 50)

    avg_probs = np.zeros(3)
    total_weight = 0
    for name, pred_class, probs, weight in predictions:
        avg_probs += probs * weight
        total_weight += weight
    avg_probs /= total_weight

    # Порог 40% для BUY/SELL
    final_class = int(avg_probs.argmax())
    threshold = 0.4 if final_class in [1, 2] else 0.0
    if avg_probs[final_class] < threshold:
        final_class = 0

    final_conf = float(avg_probs[final_class]) * 100
    avg_pct = {CLASS_NAMES[i]: f"{p*100:.1f}%" for i, p in enumerate(avg_probs)}

    signal_status = "✅ ПОРОГ 60% ПРОЙДЕН" if final_conf >= 60 else "⬇ НИЖЕ ПОРОГА 60%"
    logger.info(f"  Сигнал: {CLASS_NAMES[final_class]:>5} (conf={final_conf:.1f}%) — {signal_status}")
    logger.info(f"  Средние вероятности: HOLD={avg_pct['HOLD']} BUY={avg_pct['BUY']} SELL={avg_pct['SELL']}")

    # Итоговый вердикт
    final_signal = "HOLD" if final_conf < 60 else CLASS_NAMES[final_class]

    logger.info("\n" + "=" * 70)
    logger.info(f"ИТОГОВЫЙ ВЕРДИКТ: {final_signal}")
    logger.info("=" * 70)

    if final_signal == 'HOLD':
        logger.info(f"  Причина: confidence {final_conf:.1f}% < 60% порога")
    else:
        logger.info(f"  ✅ Уверенность {final_conf:.1f}% — порог 60% пройден!")

    # Сохранение
    results = {
        'ticker': TICKER,
        'date': '2026-06-25',
        'models': [
            {
                'name': name,
                'prediction': CLASS_NAMES[pred_class],
                'probabilities': {CLASS_NAMES[i]: f"{p*100:.1f}%"
                                  for i, p in enumerate(probs)},
            }
            for name, pred_class, probs, _ in predictions
        ],
        'ensemble': {
            'signal': CLASS_NAMES[final_class],
            'confidence': f"{final_conf:.1f}%",
            'passes_threshold': final_conf >= 60,
            'final_recommendation': final_signal,
            'probabilities': avg_pct,
        },
    }

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


if __name__ == '__main__':
    main()
