import os
import pickle
import random

import numpy as np
import torch
from torch.utils.data import Dataset

from config import Config


class PKLDataset(Dataset):
    """Load chunked dict-pkl data lazily for predictor training."""

    def __init__(self, data_type: str = "train"):
        self.config = Config()
        if data_type not in ["train", "val"]:
            raise ValueError("data_type must be 'train' or 'val'")

        self.data_type = data_type
        self.window_stride = self.config.window_stride
        self.window = self.config.lookback_window + self.config.predict_window + 1
        self.time_feature_list = self.config.time_feature_list
        self.dataset_dir = self.config.dataset_path
        self.val_ratio = self.config.val_ratio
        self.py_rng = random.Random(self.config.seed)

        self.part_files = self._list_part_files()
        self.indices = []
        self.feature_list = None
        self.current_part_path = None
        self.current_part_data = None

        print(f"[{data_type.upper()}] Pre-computing sample indices from {len(self.part_files)} part files...")
        for part_path in self.part_files:
            with open(part_path, "rb") as f:
                part_data = pickle.load(f)

            part_indices = []
            for symbol, raw_df in part_data.items():
                df = self._prepare_dataframe(raw_df, keep_feature_list=False)
                df = self._split_dataframe(df)
                series_len = len(df)
                max_start = series_len - self.window

                if max_start >= 0:
                    for i in range(0, max_start + 1, self.window_stride):
                        part_indices.append((part_path, symbol, i))

            self.py_rng.shuffle(part_indices)
            self.indices.extend(part_indices)

        self.n_samples = len(self.indices)
        print(
            f"[{data_type.upper()}] Found {self.n_samples} ordered samples "
            f"with stride={self.window_stride}."
        )

    def _list_part_files(self):
        if not os.path.isdir(self.dataset_dir):
            raise RuntimeError(f"dataset_path is not a directory: {self.dataset_dir}")

        part_files = sorted(
            os.path.join(self.dataset_dir, file_name)
            for file_name in os.listdir(self.dataset_dir)
            if file_name.startswith("train_part_") and file_name.endswith(".pkl")
        )
        if not part_files:
            raise RuntimeError(
                f"在 {self.dataset_dir} 下没有找到 train_part_*.pkl"
            )
        return part_files

    def _get_time_column(self, df):
        if "datetime" in df.columns:
            return "datetime"
        if "timestamps" in df.columns:
            return "timestamps"
        raise KeyError("DataFrame must contain either 'datetime' or 'timestamps'")

    def _prepare_dataframe(self, raw_df, keep_feature_list):
        df = raw_df.reset_index(drop=True).copy()
        time_col = self._get_time_column(df)
        df = df.sort_values(time_col).reset_index(drop=True)

        df["minute"] = df[time_col].dt.minute
        df["hour"] = df[time_col].dt.hour
        df["weekday"] = df[time_col].dt.weekday
        df["day"] = df[time_col].dt.day
        df["month"] = df[time_col].dt.month

        if self.feature_list is None:
            self.feature_list = [
                col
                for col in df.columns
                if col not in {"code", "datetime", "timestamps", *self.time_feature_list}
            ]

        if keep_feature_list:
            return df[self.feature_list + self.time_feature_list]
        return df

    def _split_dataframe(self, df):
        split_idx = int(len(df) * (1 - self.val_ratio))
        split_idx = max(0, min(split_idx, len(df)))

        if self.data_type == "train":
            return df.iloc[:split_idx].reset_index(drop=True)
        return df.iloc[split_idx:].reset_index(drop=True)

    def _load_part(self, part_path):
        if self.current_part_path == part_path and self.current_part_data is not None:
            return

        with open(part_path, "rb") as f:
            raw_part_data = pickle.load(f)

        processed_part_data = {}
        for symbol, raw_df in raw_part_data.items():
            df = self._prepare_dataframe(raw_df, keep_feature_list=False)
            df = self._split_dataframe(df)
            processed_part_data[symbol] = df[self.feature_list + self.time_feature_list]

        self.current_part_data = processed_part_data
        self.current_part_path = part_path

    def set_epoch_seed(self, epoch: int):
        """Kept for interface compatibility; ordered traversal does not reseed."""
        return None

    def __len__(self) -> int:
        return self.n_samples

    def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]:
        part_path, symbol, start_idx = self.indices[idx]
        self._load_part(part_path)

        df = self.current_part_data[symbol]
        end_idx = start_idx + self.window
        win_df = df.iloc[start_idx:end_idx]

        factor_df = win_df[self.feature_list].copy()
        factor_df = factor_df.replace([np.inf, -np.inf], np.nan)
        factor_df = factor_df.ffill().fillna(0.0)
        x = factor_df.values.astype(np.float64)
        x_stamp = win_df[self.time_feature_list].values.astype(np.float32)

        x_mean = np.mean(x, axis=0)
        x_std = np.std(x, axis=0)
        x = (x - x_mean) / (x_std + 1e-5)
        x = np.clip(x, -self.config.clip, self.config.clip)

        x_tensor = torch.from_numpy(x.astype(np.float32))
        x_stamp_tensor = torch.from_numpy(x_stamp)
        return x_tensor, x_stamp_tensor
