"""
Историческое тестирование Wyckoff D1 модели (walk-forward) — оптимизированная версия.

Оптимизация: все фичи считаются 1 раз на полном датасете,
на каждом шаге только slice + normalize + predict.

Запуск:
    python _backtest_wyckoff_history_v2.py --ticker GAZP --step 20 --forward 20
    python _backtest_wyckoff_history_v2.py --all  # все тикеры с данными
"""

import sys
import json
import warnings
import argparse
from pathlib import Path
from collections import Counter

warnings.filterwarnings("ignore")
sys.path.insert(0, str(Path(__file__).resolve().parent))

import numpy as np
import pandas as pd
import torch

from src.db.connection import fetch_ohlcv_combined

# Feature modules (импортируем 1 раз)
from src.ml.features.price_features import add_price_features
from src.ml.features.indicator_features import add_indicator_features
from src.ml.features.wyckoff_features import add_wyckoff_features
from src.ml.data.wyckoff_labeling import WyckoffLabeler

from src.ml.models.registry import ModelRegistry

# ── Map phase IDs → readable names ──
PHASE_NAMES = {
    0: 'Маркдаун', 1: 'Накопление (раннее)',
    2: 'Накопление (позд)/Маркап', 3: 'Маркап', 4: 'Распределение',
}
PHASE_DIRECTION = {0: -1, 1: 1, 2: 1, 3: 1, 4: -1}


def load_model(model_name='wyckoff_mt_d1_lstm_v2'):
    """Загрузить модель из ModelRegistry."""
    registry = ModelRegistry()
    entry = registry.get_model(model_name)
    if entry is None:
        raise ValueError(f"Модель {model_name} не найдена в реестре")
    
    model_path = Path(entry['model_path'])
    device = torch.device('cpu')
    
    # Ищем файл модели
    model_file = model_path / 'best_model.pt'
    if not model_file.exists():
        for fname in ['wyckoff_final.pt', 'model.pt', 'model_checkpoint.pt']:
            alt = model_path / fname
            if alt.exists():
                model_file = alt
                break
    
    if not model_file.exists():
        raise FileNotFoundError(f"Файл модели не найден в {model_path}")
    
    checkpoint = torch.load(str(model_file), map_location=device, weights_only=False)
    
    # Извлекаем state_dict
    if 'model_state_dict' in checkpoint:
        state_dict = checkpoint['model_state_dict']
    else:
        state_dict = checkpoint
    
    # Параметры модели
    input_size = checkpoint.get('input_size', entry['params'].get('input_size', 122))
    num_classes = checkpoint.get('num_classes', entry['params'].get('num_classes', 5))
    seq_len = checkpoint.get('seq_len', entry['params'].get('seq_len', 30))
    arch = entry.get('params', {}).get('arch', 'lstm')
    
    # Создаём модель
    from src.ml.models.wyckoff import create_wyckoff_model
    model_params = {
        'input_size': input_size,
        'num_classes': num_classes,
        'seq_len': seq_len,
        'hidden_size': entry['params'].get('hidden_size', 128),
        'num_layers': entry['params'].get('num_layers', 2),
        'dropout': entry['params'].get('dropout', 0.4),
    }
    model = create_wyckoff_model(arch, **model_params)
    model.load_state_dict(state_dict)
    model.eval()
    
    metadata = {
        'seq_len': seq_len,
        'num_classes': num_classes,
        'feature_names': checkpoint.get('feature_names', []),
        'metrics': entry.get('metrics', {}),
    }
    
    return model, metadata


def prepare_features_once(df):
    """Посчитать все фичи 1 раз на полном датасете."""
    df = add_price_features(df)
    df = add_indicator_features(df)
    df = add_wyckoff_features(df)
    
    labeler = WyckoffLabeler(smooth_window=5)
    df = labeler.label_phases(df)
    
    return df


def run_backtest(ticker, model, metadata, step=20, forward_bars=20, min_bars=200):
    """
    Walk-forward с предвычисленными фичами.
    
    Returns: list of test results
    """
    seq_len = metadata.get('seq_len', 30)
    feature_names = metadata.get('feature_names', [])
    
    # ── 1. Загружаем и подготавливаем данные ──
    print(f"  Загрузка {ticker} D1...", end=' ')
    df_raw = fetch_ohlcv_combined(ticker, 'D1', limit=5000)
    if df_raw is None or len(df_raw) < min_bars:
        print(f"❌ Недостаточно данных ({len(df_raw) if df_raw is not None else 0})")
        return None
    
    df_raw = df_raw.sort_values('timestamp').reset_index(drop=True)
    print(f"{len(df_raw)} свечей ({df_raw['Date'].iloc[0]} — {df_raw['Date'].iloc[-1]})")
    
    # ── 2. Предвычисляем фичи ──
    print(f"  Расчёт фичей...", end=' ')
    df_feat = prepare_features_once(df_raw.copy())
    print(f"✓ ({len(df_feat.columns)} колонок)")
    
    # ── 3. Отбираем нужные колонки ──
    if feature_names:
        feature_cols = [c for c in feature_names if c in df_feat.columns]
        missing = set(feature_names) - set(df_feat.columns)
        if missing:
            print(f"  ⚠️ Отсутствуют {len(missing)} признаков, заполняем нулями")
            for c in missing:
                df_feat[c] = 0.0
            feature_cols = list(feature_names)
    else:
        exclude_cols = {
            'timestamp', 'Date', 'Time', 'Open', 'High', 'Low', 'Close', 'Volume',
            'target', 'wyckoff_phase', 'wyckoff_phase_name',
            'wyckoff_phase_simple', 'wyckoff_phase_simple_name',
            'hh_hl', 'lh_ll',
        }
        feature_cols = [c for c in df_feat.columns if c not in exclude_cols
                        and not c.startswith('target_')
                        and df_feat[c].dtype in [np.float64, np.float32, np.int64, np.int32]]
    
    # Заполняем NaN
    df_feat = df_feat.ffill().bfill().fillna(0)
    
    # Матрица признаков
    X_full = df_feat[feature_cols].values.astype(np.float32)
    N = len(X_full)
    
    # ── 4. Walk-forward ──
    test_points = list(range(min_bars, N, step))
    results = []
    
    # Кэш для нормализации: на каждом шаге используем mean/std последних 200 баров ДО точки
    for idx, test_idx in enumerate(test_points):
        # Берем данные до test_idx (не включая будущее)
        X_slice = X_full[:test_idx + 1]
        
        # Нормализация (как в inference: mean/std последних 200 баров)
        window = max(200, seq_len + 10)
        norm_slice = X_slice[-window:] if len(X_slice) >= window else X_slice
        mean = norm_slice.mean(axis=0)
        std = norm_slice.std(axis=0) + 1e-8
        X_norm = (X_slice - mean) / std
        
        # Берём последние seq_len баров
        if len(X_norm) < seq_len:
            continue
        
        X_seq = X_norm[-seq_len:]
        X_seq = np.expand_dims(X_seq, axis=0).astype(np.float32)
        
        # Предсказание
        with torch.no_grad():
            x_tensor = torch.FloatTensor(X_seq)
            logits = model(x_tensor)
            probs = torch.softmax(logits, dim=1).numpy()[0]
        
        phase = int(np.argmax(probs))
        confidence = float(probs[phase] * 100)
        
        # Forward return
        current_close = float(df_raw.iloc[test_idx]['Close'])
        future_idx = min(test_idx + forward_bars, N - 1)
        future_close = float(df_raw.iloc[future_idx]['Close'])
        forward_return = (future_close / current_close - 1) * 100
        
        # Правильность направления
        pred_dir = PHASE_DIRECTION.get(phase, 0)
        actual_dir = 1 if forward_return > 0 else -1 if forward_return < 0 else 0
        direction_correct = (pred_dir * actual_dir) > 0
        
        current_date = df_raw.iloc[test_idx]['Date']
        
        results.append({
            'idx': test_idx,
            'date': current_date,
            'price': round(current_close, 2),
            'phase': phase,
            'phase_name': PHASE_NAMES.get(phase, f'Phase {phase}'),
            'confidence': round(confidence, 1),
            'forward_return': round(forward_return, 2),
            'direction_correct': direction_correct,
        })
        
        # Прогресс
        total = len(test_points)
        if total <= 20 or (idx + 1) % max(1, total // 20) == 0:
            mark = '✓' if direction_correct else '✗'
            print(f"    [{idx+1}/{total}] {current_date} | {PHASE_NAMES.get(phase, '?'):25s} | "
                  f"conf={confidence:5.1f}% | fwd={forward_return:+6.2f}% | {mark}")
    
    return results


def analyze_results(results, ticker, forward_bars):
    """Анализ и вывод результатов."""
    if not results:
        print("  Нет результатов.")
        return None
    
    df = pd.DataFrame(results)
    total = len(df)
    
    # Buy & Hold
    # (не можем посчитать без всех цен, пропускаем)
    
    # По фазам
    print(f"\n{'=' * 90}")
    print(f"      {ticker} — РЕЗУЛЬТАТЫ (forward {forward_bars}д)")
    print(f"{'=' * 90}")
    print(f"{'Фаза':<35} {'Кол-во':<8} {'%':<7} {'Ср.Conf':<9} {'Ср.Fwd':<10} {'WinRate':<10}")
    print("-" * 79)
    
    phase_counts = Counter(r['phase'] for r in results)
    for phase_id in sorted(phase_counts.keys()):
        subset = df[df['phase'] == phase_id]
        c = len(subset)
        pct = c / total * 100
        avg_conf = subset['confidence'].mean()
        avg_fwd = subset['forward_return'].mean()
        wr = subset['direction_correct'].mean() * 100
        print(f"{PHASE_NAMES.get(phase_id, f'Phase {phase_id}'):<35} "
              f"{c:<8} {pct:<7.1f} {avg_conf:<9.1f} {avg_fwd:<+10.2f} {wr:<10.1f}")
    
    # Итого
    total_win = df['direction_correct'].sum()
    avg_fwd_all = df['forward_return'].mean()
    avg_conf_all = df['confidence'].mean()
    print(f"\n{'ИТОГО':<35} {total:<8} {100:<7.1f} "
          f"{avg_conf_all:<9.1f} {avg_fwd_all:<+10.2f} {total_win/total*100:<10.1f}")
    
    # По уверенности
    print(f"\n{'─' * 50}")
    print(f"  Качество по уровням confidence:")
    print(f"{'Conf':<12} {'Кол-во':<8} {'WinRate':<10} {'Ср.Fwd':<10}")
    print(f"{'─' * 40}")
    
    buckets = [(40, 55), (55, 70), (70, 85), (85, 101)]
    for lo, hi in buckets:
        sub = df[(df['confidence'] >= lo) & (df['confidence'] < hi)]
        if len(sub) == 0:
            continue
        print(f"{lo:.0f}–{hi:.0f}%{'':<9} {len(sub):<8} "
              f"{sub['direction_correct'].mean()*100:<10.1f} {sub['forward_return'].mean():<+10.2f}")
    
    # Вывод
    print(f"\n  {'▶'*3} Выводы:")
    
    # Самая точная фаза
    best_phase = None
    best_wr = -1
    for phase_id in sorted(phase_counts.keys()):
        subset = df[df['phase'] == phase_id]
        wr = subset['direction_correct'].mean() * 100
        if len(subset) >= 2 and wr > best_wr:
            best_wr = wr
            best_phase = phase_id
    
    if best_phase is not None:
        print(f"      Лучшая фаза: {PHASE_NAMES.get(best_phase, '?')} ({best_wr:.1f}% accuracy)")
    
    # Если бы торговали только по сигналам с conf > 70%
    high_conf = df[df['confidence'] >= 70]
    if len(high_conf) >= 3:
        hc_wr = high_conf['direction_correct'].mean() * 100
        hc_fwd = high_conf['forward_return'].mean()
        print(f"      Сигналы с conf≥70%: {len(high_conf)} шт, WinRate={hc_wr:.1f}%, "
              f"средняя доходность={hc_fwd:+.2f}%")
    
    # Общая оценка
    overall_wr = total_win / total * 100
    if overall_wr >= 60:
        print(f"      ⭐ Модель показывает значимую предсказательную силу ({overall_wr:.1f}%)")
    elif overall_wr >= 50:
        print(f"      📊 Модель на уровне случайного ({overall_wr:.1f}%) — требуется доработка")
    else:
        print(f"      ❌ Модель показывает результат ниже случайного ({overall_wr:.1f}%)")
    
    print()
    
    summary = {
        'ticker': ticker,
        'total_points': total,
        'avg_confidence': round(avg_conf_all, 1),
        'avg_forward_return': round(avg_fwd_all, 2),
        'direction_accuracy': round(total_win / total * 100, 1),
        'phase_distribution': {str(k): v for k, v in phase_counts.items()},
        'best_phase': int(best_phase) if best_phase is not None else -1,
        'best_phase_accuracy': round(best_wr, 1) if best_phase is not None else 0,
    }
    
    return summary


def main():
    parser = argparse.ArgumentParser(description='Backtest Wyckoff D1 — v2 оптимизированный')
    parser.add_argument('--ticker', default='X5', help='Тикер')
    parser.add_argument('--all', action='store_true', help='Все тикеры с данными')
    parser.add_argument('--step', type=int, default=20, help='Шаг')
    parser.add_argument('--forward', type=int, default=20, help='Forward баров')
    parser.add_argument('--min-bars', type=int, default=200, help='Минимум баров')
    args = parser.parse_args()
    
    # Загружаем модель 1 раз
    print(f"Загрузка модели wyckoff_mt_d1_lstm_v2...")
    model, metadata = load_model()
    seq_len = metadata.get('seq_len', 30)
    print(f"  seq_len={seq_len}, features={len(metadata.get('feature_names', []))}")
    print()
    
    tickers = [args.ticker]
    if args.all:
        tickers = ['SBER', 'GAZP', 'LKOH', 'ROSN', 'NVTK', 'PLZL', 'PHOR',
                   'MTSS', 'VTBR', 'MOEX', 'X5', 'ASTR']
    
    all_summaries = []
    
    for ticker in tickers:
        print(f"{'=' * 90}")
        print(f"      ТЕСТИРОВАНИЕ: {ticker}")
        print(f"{'=' * 90}")
        
        results = run_backtest(ticker, model, metadata, 
                               step=args.step, forward_bars=args.forward, 
                               min_bars=args.min_bars)
        
        if results:
            summary = analyze_results(results, ticker, args.forward)
            if summary:
                all_summaries.append(summary)
            
            # Сохраняем
            out_path = Path(__file__).resolve().parent / 'reports' / f'wyckoff_backtest_{ticker}.json'
            with open(out_path, 'w', encoding='utf-8') as f:
                json.dump({
                    'ticker': ticker,
                    'config': {'step': args.step, 'forward_bars': args.forward, 'min_bars': args.min_bars},
                    'results': results,
                    'summary': summary,
                }, f, indent=2, ensure_ascii=False)
            print(f"  📁 Сохранено: {out_path}")
    
    # Сводка по всем тикерам
    if len(all_summaries) > 1:
        print(f"\n\n{'=' * 90}")
        print(f"      СВОДКА ПО ВСЕМ ТИКЕРАМ")
        print(f"{'=' * 90}")
        print(f"{'Тикер':<8} {'Точек':<8} {'Accur':<8} {'Ср.Fwd':<10} {'Ср.Conf':<10} {'Лучшая фаза':<30}")
        print("-" * 74)
        for s in all_summaries:
            best_name = PHASE_NAMES.get(s.get('best_phase', -1), '—')
            print(f"{s['ticker']:<8} {s['total_points']:<8} {s['direction_accuracy']:<8.1f} "
                  f"{s['avg_forward_return']:<+10.2f} {s['avg_confidence']:<10.1f} {best_name:<30}")
        
        print(f"\n{'=' * 90}")


if __name__ == '__main__':
    main()
