# src/training/walk_forward.py
import logging
import time
import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from sklearn.preprocessing import RobustScaler
from sklearn.metrics import roc_auc_score
from config import (
    SEQ_LEN, BATCH_SIZE, EPOCHS, LR, PATIENCE,
    MIN_HORIZON, MAX_HORIZON, DEVICE, USE_CONTEXT, CONTEXT_CAPACITY, CONTEXT_HIDDEN_DIM,
    TRAIN_LR, TRAIN_BATCH_SIZE, TRAIN_LR_FACTOR, TRAIN_LR_PATIENCE,
    TRAIN_WEIGHT_DECAY, VAL_MIN_SAMPLES, TRAIN_MIN_SAMPLES,
)
from src.models.lstm import DualHeadLSTMModel, FocalLoss
from src.models.context import ContextEnhancedLSTMModel
from src.data.loader import fetch_ohlcv
from src.data.features import engineer_features

logger = logging.getLogger(__name__)

def _log_fold_start(dataset_size: int, train_size: float, val_size: float, gap: int, seed: int):
    logger.info(f"=== FOLD START === seed={seed} dataset={dataset_size} train_prop={train_size} val_prop={val_size} gap={gap}")

def _log_batch_progress(total_batches: int, batch_idx: int, epoch_loss: float, progress_bar: str):
    if batch_idx % max(1, total_batches // 10) == 0:
        logger.info(f"Train batch {batch_idx}/{total_batches} loss={epoch_loss:.4f} {progress_bar}")



def _log_early_stopping(stopped_epoch: int):
    logger.info(f"=== EARLY STOPPING === at epoch {stopped_epoch}")

def _log_fold_end(val_auc: float, val_loss: float):
    logger.info(f"=== FOLD END === AUC={val_auc:.4f} val_loss={val_loss:.4f}")

def _log_epoch_summary(epoch: int, epochs: int, train_loss: float, val_loss: float, patience_left: int):
    logger.info(f"Epoch {epoch+1}/{epochs} train={train_loss:.4f} val={val_loss:.4f} patience={patience_left}/{PATIENCE}")






class DualDirectionDataset(torch.utils.data.Dataset):
    def __init__(self, features: np.ndarray, labels_long: np.ndarray, labels_short: np.ndarray, seq_len: int):
        mask_long = ~np.isnan(labels_long)
        mask_short = ~np.isnan(labels_short)
        valid_mask = mask_long & mask_short
        self.features = features[valid_mask]
        self.labels_long = labels_long[valid_mask]
        self.labels_short = labels_short[valid_mask]
        self.seq_len = seq_len
        self.valid_indices = np.where(valid_mask)[0]

    def __len__(self) -> int:
        valid_len = len(self.features) - self.seq_len + 1
        return max(0, valid_len)

    def __getitem__(self, idx: int):
        seq_end_idx = idx + self.seq_len - 1
        if seq_end_idx >= len(self.features):
            raise IndexError(f"Index {seq_end_idx} out of bounds for features of length {len(self.features)}")
        x = torch.tensor(self.features[idx : seq_end_idx + 1], dtype=torch.float32)
        y_long = torch.tensor(self.labels_long[seq_end_idx], dtype=torch.float32)
        y_short = torch.tensor(self.labels_short[seq_end_idx], dtype=torch.float32)
        return x, y_long, y_short


def walk_forward_loop(
    df,
    feat_cols: list,
    model_class,
    criterion_class,
    lr: float,
    batch_size: int,
    seed: int,
    train_size: float = 0.6,
    val_size: float = 0.15,
    gap: int = 5,
    epochs: int = EPOCHS,
    patience: int = PATIENCE,
    rr_ratio: float = 3.0,
    use_conv: bool = True,
    use_attention: bool = True,
    use_context: bool = USE_CONTEXT,
    hidden_dim: int = 48,
    context_capacity: int = CONTEXT_CAPACITY,
    context_hidden_dim: int = CONTEXT_HIDDEN_DIM,
):
    torch.manual_seed(seed)
    np.random.seed(seed)

    X = df[feat_cols].values
    y_long = df['Label_Long'].values
    y_short = df['Label_Short'].values
    n = len(X)

    train_end = int(n * train_size)
    step = max(int(n * 0.1), 100)

    best_val_auc = 0.0
    best_val_loss = float('inf')
    best_state, best_scaler = None, None

    calib_logits_long, calib_labels_long = [], []
    calib_logits_short, calib_labels_short = [], []

    all_fold_losses = []  # [(train_losses, val_losses), ...] per fold

    start_idx = train_end

    logger.info(f"Training started with dataset size {n}")
    while start_idx < n - MAX_HORIZON - SEQ_LEN:
        # containers for loss curves of the current fold
        fold_train_losses = []
        fold_val_losses = []
        val_start = start_idx + gap
        val_end = min(val_start + int(n * val_size), n - MAX_HORIZON)
        if val_end <= val_start + SEQ_LEN:
            start_idx += step
            continue

        scaler = RobustScaler()
        X_train_scaled = scaler.fit_transform(X[:start_idx])
        X_val_scaled = scaler.transform(X[val_start:val_end])

        train_dataset = DualDirectionDataset(X_train_scaled, y_long[:start_idx], y_short[:start_idx], SEQ_LEN)
        val_dataset = DualDirectionDataset(X_val_scaled, y_long[val_start:val_end], y_short[val_start:val_end], SEQ_LEN)

        if len(train_dataset) < 30 or len(val_dataset) < 10:
            start_idx += step
            continue

        train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
        val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)

        if use_context:
            model = ContextEnhancedLSTMModel(
                input_dim=len(feat_cols),
                hidden_dim=hidden_dim,
                num_layers=2,
                dropout=0.35,
                use_conv=use_conv,
                use_attention=use_attention,
                use_context=True,
                context_capacity=context_capacity,
                context_hidden_dim=context_hidden_dim,
            ).to(DEVICE)
        else:
            model = model_class(
                input_dim=len(feat_cols),
                hidden_dim=hidden_dim,
                use_conv=use_conv,
                use_attention=use_attention,
            ).to(DEVICE)
        criterion = criterion_class()

        # --- Балансировка классов: вычисляем pos_weight из train-данных ---
        train_labels_long = train_dataset.labels_long[train_dataset.seq_len - 1: len(train_dataset)]
        train_labels_short = train_dataset.labels_short[train_dataset.seq_len - 1: len(train_dataset)]
        pos_long = np.sum(train_labels_long == 1.0)
        neg_long = np.sum(train_labels_long == 0.0)
        pos_short = np.sum(train_labels_short == 1.0)
        neg_short = np.sum(train_labels_short == 0.0)
        pw_long = neg_long / pos_long if pos_long > 0 else 1.0
        pw_short = neg_short / pos_short if pos_short > 0 else 1.0
        # Сохраняем pos_weight в criterion (FocalLoss использует его)
        if hasattr(criterion, 'set_pos_weight'):
            # Один общий pos_weight — среднее двух направлений
            pw = (pw_long + pw_short) / 2.0
            criterion.set_pos_weight(pw)
            logger.info(f"Class balance: LONG {pos_long:.0f}/{neg_long:.0f} (pw={pw_long:.2f}) | "
                        f"SHORT {pos_short:.0f}/{neg_short:.0f} (pw={pw_short:.2f}) → pw={pw:.2f}")
        optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=TRAIN_WEIGHT_DECAY)
        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=TRAIN_LR_FACTOR, patience=TRAIN_LR_PATIENCE)

        logger.info(f"🎲 Model initialized: seed={seed} lr={lr} device={DEVICE}")
        logger.info(f"📊 FOLD START: train {len(train_dataset)} samples, val {len(val_dataset)} samples, pw={pw:.2f}")

        fold_patience = 0
        best_val_loss = float('inf')
        
        for epoch in range(epochs):
            epoch_start = time.time()
            _log_epoch_summary(epoch, epochs, 0.0, 0.0, patience - fold_patience)
            if use_context and hasattr(model, 'clear_context'):
                model.clear_context()
                
            model.train()
            epoch_train_loss = 0.0
            batch_count = 0
            for xb, yb_long, yb_short in train_loader:
                batch_count += 1
                xb, yb_long, yb_short = xb.to(DEVICE), yb_long.to(DEVICE), yb_short.to(DEVICE)
                optimizer.zero_grad()
                
                if use_context:
                    features_raw = xb[:, -1, :]
                    logits_long, logits_short = model(xb, features_raw=features_raw, store_context=True)
                else:
                    logits_long, logits_short = model(xb)
                    
                loss_long = criterion(logits_long, yb_long)
                loss_short = criterion(logits_short, yb_short)
                loss = (loss_long + loss_short) / 2.0
                loss.backward()
                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
                optimizer.step()
                epoch_train_loss += loss.item() * len(xb)
            epoch_train_loss /= len(train_dataset)
            fold_train_losses.append(epoch_train_loss)
            # Progress bar every 5% of batches
            if batch_count % max(1, len(train_loader) // 20) == 0:
                progress_bar = "█" * (batch_count // max(1, len(train_loader) // 50)) + "░" * (20 - batch_count // max(1, len(train_loader) // 50))
                logger.info(f"  Epoch train: batch {batch_count}/{len(train_loader)} loss={epoch_train_loss:.5f} [{progress_bar}]")
            # Validation for this epoch
            if use_context and hasattr(model, 'clear_context'):
                model.clear_context()
                
            model.eval()
            epoch_val_loss = 0.0
            val_samples = 0
            val_batch_count = 0
            with torch.no_grad():
                for xb, yb_long, yb_short in val_loader:
                    if use_context:
                        features_raw = xb[:, -1, :]
                        logits_l, logits_s = model(xb.to(DEVICE), features_raw=features_raw, store_context=True)
                    else:
                        logits_l, logits_s = model(xb.to(DEVICE))
                        
                    val_loss_l = criterion(logits_l, yb_long.to(DEVICE))
                    val_loss_s = criterion(logits_s, yb_short.to(DEVICE))
                    val_loss = (val_loss_l + val_loss_s) / 2.0
                    epoch_val_loss += val_loss.item() * len(xb)
                    val_samples += len(xb)
                    val_batch_count += 1
            logger.info(f"  Epoch validation complete: {val_batch_count} batches, {val_samples} samples, val_loss={epoch_val_loss:.5f}")
            if len(val_dataset) > 0:
                epoch_val_loss /= len(val_dataset)
            else:
                epoch_val_loss = float('inf')
            
            # Minimize epoch logging
            if epoch == 0 or epoch == epochs - 1 or epoch % 5 == 0:
                logger.info(f"Epoch {epoch+1}/{epochs} train={epoch_train_loss:.4f} val={epoch_val_loss:.4f} patience={fold_patience}/{patience}")
            
            patience_left = max(patience - fold_patience, 0)

            improved_val_loss = epoch_val_loss < best_val_loss
            if improved_val_loss:
                best_val_loss = epoch_val_loss
                fold_patience = 0
                logger.info(f"  ✓ New best val loss: {best_val_loss:.4f}")
            else:
                fold_patience += 1
                logger.info(f"  Epoch validation complete: {val_batch_count} batches, val loss={epoch_val_loss:.4f}")

            if fold_patience >= patience:
                _log_early_stopping(epoch + 1)
                break

        model.eval()
        epoch_val_loss = 0.0
        fold_logits_long, fold_labels_long = [], []
        fold_logits_short, fold_labels_short = [], []
        
        if use_context and hasattr(model, 'clear_context'):
            model.clear_context()

        with torch.no_grad():
            for xb, yb_long, yb_short in val_loader:
                if use_context:
                    features_raw = xb[:, -1, :]
                    logits_l, logits_s = model(xb.to(DEVICE), features_raw=features_raw, store_context=True)
                else:
                    logits_l, logits_s = model(xb.to(DEVICE))
                    
                val_loss_l = criterion(logits_l, yb_long.to(DEVICE))
                val_loss_s = criterion(logits_s, yb_short.to(DEVICE))
                val_loss = (val_loss_l + val_loss_s) / 2.0
                epoch_val_loss += val_loss.item() * len(xb)
                fold_logits_long.extend(logits_l.cpu().numpy())
                fold_labels_long.extend(yb_long.cpu().numpy())
                fold_logits_short.extend(logits_s.cpu().numpy())
                fold_labels_short.extend(yb_short.cpu().numpy())
        logger.info(f"Validation complete: processed {len(fold_labels_long)} samples")

        if len(val_dataset) > 0:
            epoch_val_loss /= len(val_dataset)
        else:
            epoch_val_loss = float('inf')

        scheduler.step(epoch_val_loss)

        # AUC calculations
        auc_long = roc_auc_score(fold_labels_long, fold_logits_long) if len(set(fold_labels_long)) > 1 else 0.0
        auc_short = roc_auc_score(fold_labels_short, fold_logits_short) if len(set(fold_labels_short)) > 1 else 0.0
        current_auc = (auc_long + auc_short) / 2.0

        logger.info(f"=== VALIDATION RESULTS === val_loss={epoch_val_loss:.4f} AUC_long={auc_long:.4f} AUC_short={auc_short:.4f} AUC_avg={current_auc:.4f}")

        if current_auc > best_val_auc:
            best_val_auc = current_auc
            best_state = {k: v.cpu().clone() for k, v in model.state_dict().items()}
            best_state['val_auc'] = current_auc
            best_scaler = scaler
            logger.info(f"✓ New best AUC: {best_val_auc:.4f}")
            calib_logits_long = list(fold_logits_long)
            calib_labels_long = list(fold_labels_long)
            calib_logits_short = list(fold_logits_short)
            calib_labels_short = list(fold_labels_short)
        else:
            logger.info(f"⚠ AUC {current_auc:.4f} did not beat best {best_val_auc:.4f}, continuing...")

        # Сохраняем потери фолда для отчёта
        all_fold_losses.append((list(fold_train_losses), list(fold_val_losses)))

        _log_fold_end(current_auc, epoch_val_loss)
        start_idx += step

    logger.info("=== TRAINING COMPLETE ===")
    return best_state, best_scaler, (
        calib_logits_long, calib_labels_long,
        calib_logits_short, calib_labels_short
    ), all_fold_losses
