"""
Ensemble inference for X5 using all trained NN models.

Loads models:
    1. x5_lstm_v2 — LSTM, D1, pre-train+fine-tune (test acc 57.58%)
    2. x5_h1_lstm_v1 — LSTM, H1 (test acc 58.55%)
    3. x5_gru_v3 — GRU, D1, Focal Loss (test acc 48.48%, SELL recall 50%)

Produces:
    - Per-model predictions
    - Ensemble (weighted voting) prediction
    - Comparison with traditional TA

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

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

import numpy as np
import pandas as pd
import torch
import torch.nn as nn

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  # Go up 4 levels to project root
sys.path.insert(0, str(project_root))

from src.ml.models.lstm import LSTMPredictor, GRUPredictor
from src.ml.models.registry import ModelRegistry
from src.ml.features.pipeline import create_ml_dataset

# === Configuration ===
MODELS_DIR = Path('src/ml/models/saved')
TICKER = 'X5'

# Model weights for ensemble (based on test performance)
ENSEMBLE_WEIGHTS = {
    'x5_lstm_v2': 0.35,      # Best test accuracy (57.58%)
    'x5_h1_lstm_v1': 0.35,   # Best test accuracy (58.55%)
    'x5_gru_v3': 0.30,       # Best SELL recall (50%)
}

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


def load_model_for_inference(
    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', {})

        # Create model
        model = model_class(
            input_size=params.get('input_size', 61),
            hidden_size=params.get('hidden_size', 128),
            num_layers=params.get('num_layers', 2),
            dropout=params.get('dropout', 0.3),
        )

        # Load weights
        state_dict = torch.load(checkpoint_path, map_location=device, weights_only=True)
        model.load_state_dict(state_dict)
        model = model.to(device)
        model.eval()

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

        # Calibrated thresholds
        thresholds = config.get('calibrated_thresholds', {0: 0.0, 1: 0.0, 2: 0.0})

        return {
            'model': model,
            'config': config,
            'normalizer': normalizer,
            'thresholds': thresholds,
            'feature_cols': config.get('feature_cols', []),
        }

    except Exception as e:
        logger.error(f"Ошибка загрузки модели {model_name}: {e}")
        return None


def get_latest_data(
    ticker: str,
    tf: str,
    limit: int = 200,
) -> Optional[pd.DataFrame]:
    """Получить последние данные для инференса."""
    try:
        # Use feature pipeline to get properly formatted data
        data_dict = create_ml_dataset(
            ticker=ticker,
            tf=tf,
            seq_len=min(limit, 200),
            target_threshold=0.005,
            save=False,
        )
        # The dataset generates sequences; we need raw features.
        # Let's use fetch_ohlcv_combined instead for raw data
        from src.db.connection import fetch_ohlcv_combined
        df = fetch_ohlcv_combined(ticker, tf, limit=limit)
        if df is None or len(df) < 60:
            logger.error(f"Недостаточно данных для {ticker} {tf}: {len(df) if df is not None else 0}")
            return None
        return df
    except Exception as e:
        logger.error(f"Ошибка загрузки данных: {e}")
        return None


def prepare_features(df: pd.DataFrame, feature_cols: List[str]) -> Optional[np.ndarray]:
    """Подготовить фичи для инференса (Z-score по нормализатору)."""
    try:
        # Check which features are available
        available = [c for c in feature_cols if c in df.columns]
        missing = [c for c in feature_cols if c not in df.columns]

        if missing:
            logger.warning(f"Отсутствуют фичи: {missing[:5]}...")
            if len(available) < len(feature_cols) * 0.8:
                logger.error(f"Слишком много отсутствующих фич: {len(missing)}/{len(feature_cols)}")
                return None

        # Get features in correct order (matching normalizer)
        X_full = df[feature_cols].values.astype(np.float64)

        # Drop rows with NaN (indicators need warmup). Keep tail portion.
        # Find last row without NaN
        nan_mask = np.isnan(X_full).any(axis=1)
        if nan_mask.any():
            last_valid = np.where(~nan_mask)[0]
            if len(last_valid) == 0:
                logger.error("Все строки содержат NaN — недостаточно данных для индикаторов")
                return None
            X_full = X_full[last_valid[0]:]
            logger.info(f"  NaN удалены: осталось {len(X_full)} строк (из {len(df)} с NaN)")

        # Take last 150 rows max
        X = X_full[-150:] if len(X_full) > 150 else X_full
        return X
    except Exception as e:
        logger.error(f"Ошибка подготовки фич: {e}")
        import traceback
        traceback.print_exc()
        return None


def predict_model(
    model_artifacts: Dict,
    X: np.ndarray,
    seq_len: int,
    device: torch.device,
    use_thresholds: bool = True,
) -> Tuple[int, np.ndarray]:
    """
    Предсказание одной модели.

    Returns:
        (predicted_class, probabilities).
    """
    model = model_artifacts['model']
    normalizer = model_artifacts.get('normalizer')
    thresholds = model_artifacts.get('thresholds', {0: 0.0, 1: 0.0, 2: 0.0})

    # Normalize
    if normalizer:
        # Handle missing features by using 0 for missing
        X_norm = np.zeros_like(X)
        for i in range(X.shape[1]):
            if normalizer['std'][i] > 0:
                X_norm[:, i] = (X[:, i] - normalizer['mean'][i]) / normalizer['std'][i]
    else:
        X_norm = X

    # Take last seq_len samples and add batch dim
    if len(X_norm) < seq_len:
        X_seq = torch.FloatTensor(X_norm[-min(seq_len, len(X_norm)):]).unsqueeze(0)
        # Pad if needed
        if X_seq.shape[1] < seq_len:
            pad = torch.zeros(1, seq_len - X_seq.shape[1], X_seq.shape[2])
            X_seq = torch.cat([pad, X_seq], dim=1)
    else:
        X_seq = torch.FloatTensor(X_norm[-seq_len:]).unsqueeze(0)

    # Inference
    with torch.no_grad():
        output = model(X_seq.to(device))
        probs = torch.softmax(output, dim=1).squeeze().cpu().numpy()

    if len(probs.shape) == 0:
        probs = np.array([probs])

    # Apply thresholds
    if use_thresholds:
        if probs.max() < thresholds.get(int(probs.argmax()), 0.0):
            return 0, probs  # HOLD
        return int(probs.argmax()), probs

    return int(probs.argmax()), probs


def ensemble_vote(
    predictions: List[Tuple[str, int, np.ndarray]],
    weights: Dict[str, float],
) -> Tuple[int, np.ndarray, float]:
    """
    Взвешенное голосование ансамбля.

    Returns:
        (ensemble_class, averaged_probs, ensemble_confidence).
    """
    # Weighted average of probabilities
    avg_probs = np.zeros(3)
    total_weight = 0

    for model_name, pred_class, probs in predictions:
        w = weights.get(model_name, 1.0)
        avg_probs += probs * w
        total_weight += w

    avg_probs /= total_weight

    # Apply ensemble threshold (0.4 for BUY/SELL, 0.0 for HOLD)
    ensemble_class = int(avg_probs.argmax())
    threshold = 0.4 if ensemble_class in [1, 2] else 0.0
    if avg_probs.max() < threshold:
        ensemble_class = 0  # HOLD

    ensemble_confidence = float(avg_probs[ensemble_class]) * 100

    return ensemble_class, avg_probs, ensemble_confidence


def main():
    """Запуск ансамблевого инференса для X5."""
    logger.info("=" * 70)
    logger.info("АНСАМБЛЕВЫЙ НЕЙРОСЕТЕВОЙ ИНФЕРЕНС ДЛЯ X5")
    logger.info("=" * 70)

    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    logger.info(f"Устройство: {device}")

    # 1. Get latest data for D1 and H1
    logger.info("\nЗагрузка данных...")

    # For D1 models
    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

    # Fetch D1 data (need >260 for SMA200+seq_len=60)
    df_d1 = fetch_ohlcv_combined(TICKER, 'D1', limit=400)
    if df_d1 is None or len(df_d1) < 120:
        logger.error("Недостаточно D1 данных")
        sys.exit(1)
    logger.info(f"D1 данных: {len(df_d1)} баров")

    # Add features for D1
    df_d1 = add_price_features(df_d1)
    df_d1 = add_indicator_features(df_d1)
    df_d1 = add_calendar_features(df_d1)

    # Fetch H1 data (need >320 for SMA200+seq_len=120)
    df_h1 = fetch_ohlcv_combined(TICKER, 'H1', limit=500)
    if df_h1 is None or len(df_h1) < 120:
        logger.error("Недостаточно H1 данных")
    else:
        logger.info(f"H1 данных: {len(df_h1)} баров")
        df_h1 = add_price_features(df_h1)
        df_h1 = add_indicator_features(df_h1)
        df_h1 = add_calendar_features(df_h1)

    # 2. Load all models
    logger.info("\nЗагрузка моделей...")

    models_to_load = [
        ('x5_lstm_v2', LSTMPredictor, MODELS_DIR / 'x5_lstm_v2'),
        ('x5_h1_lstm_v1', LSTMPredictor, MODELS_DIR / 'x5_h1_lstm_v1'),
        ('x5_gru_v3', GRUPredictor, MODELS_DIR / 'x5_gru_v3'),
    ]

    loaded_models = {}
    for name, model_class, path in models_to_load:
        config_path = path / 'config.json'
        checkpoint_path = path / 'model.pt'
        artifacts = load_model_for_inference(name, model_class, config_path, checkpoint_path, device)
        if artifacts:
            loaded_models[name] = artifacts
            tf = artifacts['config'].get('timeframe', '?')
            logger.info(f"  ✅ {name} ({tf}) — загружена")
        else:
            logger.warning(f"  ❌ {name} — не загружена")

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

    # 3. Run inference per model
    logger.info("\n" + "-" * 60)
    logger.info("ИНФЕРЕНС")
    logger.info("-" * 60)

    predictions = []
    for model_name, artifacts in loaded_models.items():
        tf = artifacts['config'].get('timeframe', 'D1')
        feature_cols = artifacts['feature_cols']
        seq_len = artifacts['config'].get('params', {}).get('seq_len', 60)

        # Select correct dataframe
        df = df_d1 if tf == 'D1' else (df_h1 if tf == 'H1' else df_d1)

        # Prepare features
        X = prepare_features(df, feature_cols)
        if X is None:
            logger.warning(f"  {model_name}: пропускаем (нет фич)")
            continue

        # Predict
        pred_class, probs = predict_model(artifacts, X, seq_len, device, use_thresholds=True)
        predictions.append((model_name, pred_class, probs))

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

    # 4. Ensemble vote
    if len(predictions) > 1:
        logger.info("\n" + "-" * 60)
        logger.info("АНСАМБЛЕВОЕ ГОЛОСОВАНИЕ")
        logger.info("-" * 60)

        ens_class, avg_probs, ens_conf = ensemble_vote(predictions, ENSEMBLE_WEIGHTS)
        avg_pct = {CLASS_NAMES[i]: f"{p*100:.1f}%" for i, p in enumerate(avg_probs)}

        logger.info(f"  Ensemble: {CLASS_NAMES[ens_class]:>5} (conf={ens_conf:.1f}%)")
        logger.info(f"  Средние вероятности: HOLD={avg_pct['HOLD']} "
                    f"BUY={avg_pct['BUY']} SELL={avg_pct['SELL']}")
        logger.info(f"  Веса: {ENSEMBLE_WEIGHTS}")
    else:
        ens_class, avg_probs, ens_conf = predictions[0][1], predictions[0][2], 0.0
        logger.info("  Ensemble: только 1 модель — используем её предсказание")

    # 5. Per-model votes (for reporting)
    vote_details = {}
    for model_name, pred_class, probs in predictions:
        vote_details[model_name] = {
            'signal': CLASS_NAMES[pred_class],
            'confidence': f"{float(probs[pred_class])*100:.1f}%",
            'probabilities': {CLASS_NAMES[i]: f"{float(p)*100:.1f}%"
                              for i, p in enumerate(probs)},
        }

    # 6. Build final report
    results = {
        'ticker': TICKER,
        'date': '2026-06-25',
        'model_registry': {
            'active_models': [
                {'name': 'x5_lstm_v2', 'type': 'LSTMPredictor', 'tf': 'D1',
                 'test_acc': 57.58, 'SELL_recall': 0.0},
                {'name': 'x5_h1_lstm_v1', 'type': 'LSTMPredictor', 'tf': 'H1',
                 'test_acc': 58.55, 'SELL_recall': 5.2},
                {'name': 'x5_gru_v3', 'type': 'GRUPredictor', 'tf': 'D1',
                 'test_acc': 48.48, 'SELL_recall': 50.0, 'calibrated_acc': 54.55},
            ],
            'best_model': 'x5_lstm_v2' if len(predictions) > 0 else 'none',
        },
        'per_model': vote_details,
        'ensemble': {
            'signal': CLASS_NAMES[ens_class],
            'confidence': f"{ens_conf:.1f}%",
            'probabilities': {CLASS_NAMES[i]: f"{float(p)*100:.1f}%"
                              for i, p in enumerate(avg_probs)},
            'weights': ENSEMBLE_WEIGHTS,
            'num_models': len(predictions),
        },
        'latest_price': {
            'close': float(df_d1['Close'].iloc[-1]),
            'date': str(df_d1['Date'].iloc[-1]) if 'Date' in df_d1.columns else 'latest',
        },
    }

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

    # Print summary
    logger.info("\n" + "=" * 70)
    logger.info("ИТОГОВЫЙ НЕЙРОСЕТЕВОЙ СИГНАЛ")
    logger.info("=" * 70)
    logger.info(f"  Тикер: {TICKER}")
    logger.info(f"  Ensemble сигнал: {results['ensemble']['signal']} "
                f"(conf={results['ensemble']['confidence']})")
    logger.info(f"  Вероятности: {results['ensemble']['probabilities']}")
    logger.info(f"  Участвовало моделей: {results['ensemble']['num_models']}")
    logger.info(f"  Цена закрытия: {results['latest_price']['close']}")
    logger.info(f"  Дата: {results['latest_price']['date']}")
    logger.info("=" * 70)

    # Print comparison with v2 and v3
    logger.info("\nСравнение моделей:")
    for name, model_pred in vote_details.items():
        logger.info(f"  {name:<20}: {model_pred['signal']:>5} "
                    f"(conf={model_pred['confidence']})")

    return results


if __name__ == '__main__':
    main()
