"""
Очищенный ансамбль X5 — ТОЛЬКО модели без оверфиттинга.

Удалены:
    ❌ x5_lstm_v2  — SELL recall 0%, но выдаёт SELL 81.3% (некалиброван)
    ❌ x5_h1_lstm_v1 — BUY bias 34%, SELL recall 10%

Оставлена:
    ✅ x5_gru_v3 — SELL recall 50%, Focal Loss, calibrated thresholds

Usage:
    python src/ml/inference/ensemble_predict_x5_clean.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.registry import ModelRegistry
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'}

# Модели БЕЗ оверфиттинга — все с Focal Loss, все с SELL recall > 0%
TRUSTED_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 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', []),
        }
    except Exception as e:
        logger.error(f"Ошибка загрузки {model_name}: {e}")
        return None


def prepare_features(df, feature_cols):
    """Подготовить фичи."""
    X_full = df[feature_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, seq_len, device):
    """Inference одной модели."""
    model = model_artifacts['model']
    norm = model_artifacts.get('normalizer')
    thresholds = model_artifacts.get('thresholds', {0: 0.0, 1: 0.4, 2: 0.4})

    if norm:
        X_norm = np.zeros_like(X)
        for i in range(X.shape[1]):
            if norm['std'][i] > 0:
                X_norm[:, i] = (X[:, i] - norm['mean'][i]) / norm['std'][i]
    else:
        X_norm = X

    if len(X_norm) < seq_len:
        pad = np.zeros((seq_len - len(X_norm), X_norm.shape[1]))
        X_seq = np.concatenate([pad, X_norm], axis=0)
    else:
        X_seq = X_norm[-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 (БЕЗ ОВЕРФИТТИНГА)")
    logger.info("=" * 70)
    logger.info(f"  Удалены:     x5_lstm_v2 (SELL recall=0%), x5_h1_lstm_v1 (BUY bias=34%)")
    logger.info(f"  Ансамбль 3х: x5_gru_v3 (SELL recall=50%) + x5_lstm_v3_focal (SELL recall=40%) + x5_gru_v4 (SELL recall=42%)")

    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)}")

    # Загружаем все 3 модели
    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)
        if arts:
            models_loaded[name] = {**arts, 'weight': info['weight'], 'tf': info['tf']}
            logger.info(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']
        seq_len = artifacts['config'].get('params', {}).get('seq_len', 60)

        X = prepare_features(df_feat, feature_cols)
        if X is None or len(X) < seq_len:
            logger.warning(f"  {name}: недостаточно данных ({len(X) if X is not None else 0})")
            continue

        pred_class, probs = predict_one(artifacts, X, seq_len, 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} → {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 (calibrated threshold)
    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% → HOLD"
    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("  Причина: confidence ниже порога 60%")
        logger.info(f"  SELL={avg_pct['SELL']} < 60% → вне рынка")

    # Обновлённый ensemble-результат
    results = {
        'ticker': TICKER,
        'date': '2026-06-25',
        'models_removed': [
            {'name': 'x5_lstm_v2', 'reason': 'SELL recall=0% на тесте, но даёт SELL 81.3% — некалиброванный оверфиттинг'},
            {'name': 'x5_h1_lstm_v1', 'reason': 'BUY bias=34% (324/952 HOLD→BUY), SELL recall=10%, gap val-test 6.5%'},
        ],
        'models_kept': [
            {'name': 'x5_gru_v3', 'reason': 'SELL recall=50%, Focal Loss γ=2.0, GRU 108K params'},
            {'name': 'x5_lstm_v3_focal', 'reason': 'SELL recall=40%, Focal Loss γ=2.0, LSTM 238K params'},
            {'name': 'x5_gru_v4', 'reason': 'SELL recall=42%, Focal Loss γ=3.0, GRU 181K params'},
        ],
        'clean_ensemble': {
            'signal': CLASS_NAMES[final_class],
            'confidence': f"{final_conf:.1f}%",
            'passes_threshold': final_conf >= 60,
            'final_recommendation': final_signal,
            'probabilities': {CLASS_NAMES[i]: f"{p*100:.1f}%"
                              for i, p in enumerate(avg_probs)},
        },
    }

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


if __name__ == '__main__':
    main()
