"""
Ансамбль Wyckoff-моделей H1 + D1 + W1 для определения фаз Вайкоффа.

Каждая модель (H1, D1, W1) предсказывает фазу независимо.
Финальное предсказание — взвешенное голосование с учётом confidence.

Usage:
    python src/ml/inference/wyckoff_ensemble.py SBER GAZP X5
"""

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

import numpy as np
import pandas as pd
import torch

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

from src.ml.inference.wyckoff_inference import WyckoffInference, WyckoffPhaseResult

logging.basicConfig(level=logging.INFO, format='%(asctime)s | %(message)s', datefmt='%H:%M:%S')
logger = logging.getLogger('wyckoff_ensemble')

# Веса моделей по качеству на валидации
TF_WEIGHTS = {
    'D1': 0.50,  # LSTM v2: Acc=71.5%, F1=64.1%
    'H1': 0.30,  # LSTM v1: Acc=66.4%, F1=53.0%
    'W1': 0.20,  # LSTM v1: Acc=40.6%, F1=30.8%
}

TF_MODELS = {
    'D1': 'wyckoff_mt_d1_lstm_v2',
    'H1': 'wyckoff_mt_h1_lstm_v1',
    'W1': 'wyckoff_mt_w1_lstm_v1',
}

PHASE_NAMES_SIMPLE = {
    0: 'Маркдаун', 1: 'Накопление (раннее)',
    2: 'Накопление (позднее)/Маркап', 3: 'Маркап', 4: 'Распределение',
}

PHASE_COLORS = {
    0: '\033[91m',  # Красный
    1: '\033[93m',  # Жёлтый
    2: '\033[94m',  # Синий
    3: '\033[92m',  # Зелёный
    4: '\033[95m',  # Фиолетовый
}
RESET_COLOR = '\033[0m'


class WyckoffEnsemble:
    """
    Ансамбль Wyckoff-моделей по всем ТФ.
    
    Загружает H1, D1, W1 модели и комбинирует их предсказания
    через взвешенное голосование.
    """

    def __init__(self, use_rule_based_fallback: bool = True):
        self.inferers: Dict[str, WyckoffInference] = {}
        self.metadata: Dict[str, Any] = {}

        for tf, model_name in TF_MODELS.items():
            try:
                inferer = WyckoffInference(
                    model_name=model_name,
                    ticker='SBER',
                    tf=tf,
                    use_rule_based_fallback=use_rule_based_fallback,
                )
                self.inferers[tf] = inferer
                self.metadata[tf] = {
                    'model_name': model_name,
                    'f1': inferer.metadata.get('metrics', {}).get('val_macro_f1', 0),
                    'weight': TF_WEIGHTS.get(tf, 1/3),
                }
                logger.info(f"  ✅ {tf}: {model_name} (F1={self.metadata[tf]['f1']:.2%})")
            except Exception as e:
                logger.warning(f"  ⚠️ {tf}: не удалось загрузить — {e}")

    def analyze(self, ticker: str, tf: str = 'D1') -> 'EnsembleResult':
        """
        Проанализировать тикер по всем ТФ и объединить результаты.

        Args:
            ticker: тикер MOEX.
            tf: целевой ТФ (для rule-based fallback).

        Returns:
            EnsembleResult с комбинированным предсказанием.
        """
        tf_results = {}

        for model_tf, inferer in self.inferers.items():
            try:
                result = inferer.analyze(ticker=ticker, tf=model_tf)
                tf_results[model_tf] = result
            except Exception as e:
                logger.error(f"  {ticker} {model_tf}: ошибка — {e}")

        if not tf_results:
            logger.warning(f"  {ticker}: ни одна модель не сработала, fallback на D1")
            inferer = list(self.inferers.values())[0]
            result = inferer.analyze(ticker=ticker, tf='D1')
            tf_results = {'D1': result}

        return self._combine(ticker, tf_results)

    def _combine(self, ticker: str, tf_results: Dict[str, WyckoffPhaseResult]) -> 'EnsembleResult':
        """
        Комбинировать предсказания по всем ТФ через взвешенное голосование.
        """
        num_classes = 5
        weighted_probs = np.zeros(num_classes, dtype=np.float64)
        total_weight = 0.0
        details = {}

        for tf, result in tf_results.items():
            weight = TF_WEIGHTS.get(tf, 1/3)
            # Конвертируем str keys → int
            probs = np.zeros(num_classes)
            for k, v in result.phase_probs.items():
                probs[int(k)] = v

            weighted_probs += weight * probs
            total_weight += weight

            details[tf] = {
                'phase': int(result.phase),
                'phase_name': result.phase_name,
                'confidence': result.confidence,
                'probs': probs.tolist(),
            }

        if total_weight > 0:
            weighted_probs /= total_weight

        ensemble_phase = int(np.argmax(weighted_probs))
        ensemble_confidence = float(weighted_probs[ensemble_phase] * 100)
        ensemble_name = PHASE_NAMES_SIMPLE.get(ensemble_phase, f'Phase {ensemble_phase}')

        return EnsembleResult(
            ticker=ticker,
            phase=ensemble_phase,
            phase_name=ensemble_name,
            confidence=ensemble_confidence,
            phase_probs={i: float(p) for i, p in enumerate(weighted_probs)},
            details=details,
        )

    def batch_analyze(self, tickers: List[str]) -> List['EnsembleResult']:
        """Проанализировать список тикеров."""
        results = []
        for ticker in tickers:
            try:
                result = self.analyze(ticker)
                results.append(result)
            except Exception as e:
                logger.error(f"  {ticker}: ошибка — {e}")
        return results

    def print_summary(self, results: List['EnsembleResult']):
        """Цветной вывод сводки."""
        print()
        print("=" * 80)
        print(f"{'АНСАМБЛЬ WYCKOFF (H1+D1+W1)':^80}")
        print("=" * 80)

        for r in results:
            color = PHASE_COLORS.get(r.phase, '')
            print(
                f"  {r.ticker:6s} | "
                f"{color}{r.phase_name:40s}{RESET_COLOR} | "
                f"conf: {r.confidence:5.1f}%"
            )
            for tf, d in r.details.items():
                tf_color = PHASE_COLORS.get(d['phase'], '')
                print(
                    f"           {tf}: "
                    f"{tf_color}{d['phase_name']:30s}{RESET_COLOR} "
                    f"({d['confidence']:.0f}%)"
                )

        print("\n📊 Распределение:")
        phase_counts = {}
        for r in results:
            phase_counts[r.phase] = phase_counts.get(r.phase, 0) + 1
        for phase_id in sorted(phase_counts.keys()):
            pname = PHASE_NAMES_SIMPLE.get(phase_id, f'Phase {phase_id}')
            bar = '█' * phase_counts[phase_id]
            print(f"  {pname:40s}: {bar} {phase_counts[phase_id]}")


class EnsembleResult:
    """Результат ансамблевого предсказания."""

    def __init__(
        self,
        ticker: str,
        phase: int,
        phase_name: str,
        confidence: float,
        phase_probs: Dict[int, float],
        details: Dict[str, Any],
    ):
        self.ticker = ticker
        self.phase = phase
        self.phase_name = phase_name
        self.confidence = confidence
        self.phase_probs = phase_probs
        self.details = details
        self.is_accumulation = phase in (1, 2)
        self.is_markup = phase == 3
        self.is_distribution = phase == 4
        self.is_markdown = phase == 0

    def to_dict(self) -> Dict[str, Any]:
        return {
            'ticker': self.ticker,
            'phase': int(self.phase),
            'phase_name': self.phase_name,
            'confidence': round(float(self.confidence), 1),
            'phase_probs': {str(k): round(float(v), 3) for k, v in self.phase_probs.items()},
            'details': self.details,
        }


def main():
    import argparse
    parser = argparse.ArgumentParser(description='Ансамбль Wyckoff H1+D1+W1')
    parser.add_argument('tickers', nargs='+', help='Тикеры для анализа')
    args = parser.parse_args()

    ensemble = WyckoffEnsemble(use_rule_based_fallback=True)
    results = ensemble.batch_analyze(args.tickers)
    ensemble.print_summary(results)

    # JSON-вывод
    import json
    print("\n" + "=" * 80)
    print("JSON:")
    print(json.dumps(
        [r.to_dict() for r in results],
        indent=2, ensure_ascii=False,
    ))


if __name__ == '__main__':
    main()
