from collections.abc import Iterator, Sequence
import multiprocessing
import os
import typing
from typing import Protocol, SupportsIndex, TypeVar

# import jax
# import jax.numpy as jnp
import lerobot.common.datasets.lerobot_dataset as lerobot_dataset
import numpy as np
import torch
import pathlib
from lerobot.configs.policies import PreTrainedConfig
from typing import List

import openpi.models.model as _model
import openpi.training.config as _config
import openpi.transforms as _transforms

import time

T_co = TypeVar("T_co", covariant=True)


class Dataset(Protocol[T_co]):
    """Interface for a dataset with random access."""

    def __getitem__(self, index: SupportsIndex) -> T_co:
        raise NotImplementedError("Subclasses of Dataset should implement __getitem__.")

    def __len__(self) -> int:
        raise NotImplementedError("Subclasses of Dataset should implement __len__.")


class DataLoader(Protocol[T_co]):
    """Interface for a data loader."""

    def data_config(self) -> _config.DataConfig:
        """Get the data config for this data loader."""
        raise NotImplementedError("Subclasses of DataLoader should implement data_config.")

    def __iter__(self) -> Iterator[T_co]:
        raise NotImplementedError("Subclasses of DataLoader should implement __iter__.")


class TransformedDataset(Dataset[T_co]):
    def __init__(self, dataset: Dataset, transforms: Sequence[_transforms.DataTransformFn]):
        self._dataset = dataset
        self._transform = _transforms.compose(transforms)

    def __getitem__(self, index: SupportsIndex) -> T_co:
        return  self._transform(self._dataset[index])

    def __len__(self) -> int:
        return len(self._dataset)


class TransformedDataset_Late(Dataset[T_co]):
    def __init__(self, dataset: Dataset, transforms: Sequence[_transforms.DataTransformFn]):
        self._dataset = dataset
        self._transform = _transforms.compose(transforms)

    def __getitem__(self, index: SupportsIndex) -> T_co:
        batch =  self._transform(self._dataset[index])
        return _model.from_dict(batch), batch["actions"]

    def __len__(self) -> int:
        return len(self._dataset)
    

class TransformedPretrainDataset(Dataset[T_co]):
    def __init__(self, dataset: Dataset, transforms: List[Sequence[_transforms.DataTransformFn]]):
        self._dataset = dataset
        self._transform = []
        for transform in transforms:
            self._transform.append(_transforms.compose(transform))

    def __getitem__(self, index: SupportsIndex) -> T_co:
        dataset_idx = self._dataset[index]["dataset_index"]
        return self._transform[dataset_idx](self._dataset[index])

    def __len__(self) -> int:
        return len(self._dataset)
    

class TransformedPretrainDataset_Late(Dataset[T_co]):
    def __init__(self, dataset: Dataset, transforms: List[Sequence[_transforms.DataTransformFn]]):
        self._dataset = dataset
        self._transform = []
        for transform in transforms:
            self._transform.append(_transforms.compose(transform))

    def __getitem__(self, index: SupportsIndex) -> T_co:
        dataset_idx = self._dataset[index]["dataset_index"]
        batch = self._transform[dataset_idx](self._dataset[index])
        return _model.from_dict(batch), batch["actions"]

    def __len__(self) -> int:
        return len(self._dataset)
    

def create_dataset(data_config: _config.DataConfig, assets_dirs: pathlib.Path, model_config: PreTrainedConfig) -> Dataset:
    """Create a dataset for training."""
    repo_id = data_config.repo_id
    if repo_id is None:
        raise ValueError("Repo ID is not set. Cannot create dataset.")
    data_root = assets_dirs / repo_id
    dataset_meta = lerobot_dataset.LeRobotDatasetMetadata(repo_id, data_root, local_files_only=data_config.local_files_only)
    start_time = time.time()
    dataset = lerobot_dataset.LeRobotDataset(
        data_config.repo_id,
        data_root, 
        delta_timestamps={
            key: [t / dataset_meta.fps for t in range(model_config.n_action_steps)]
            for key in data_config.action_sequence_keys
        },
        local_files_only=data_config.local_files_only,
    )
    end_time = time.time()
    execution_time = end_time - start_time
    print(f"***********Time of data load: {execution_time} 秒")
    num_frames = dataset.num_frames
    num_episodes = dataset.num_episodes
    if data_config.prompt_from_task:
        dataset = TransformedDataset(dataset, [_transforms.PromptFromLeRobotTask(dataset_meta.tasks)])

    return dataset, num_frames, num_episodes


def create_pretrain_dataset(data_config: List[_config.DataConfig], assets_dirs: List[pathlib.Path], model_config: List[PreTrainedConfig], data_weights: List[float]) -> Dataset:
    """Create a dataset for training."""
    start_time = time.time()
    repo_ids = []
    data_roots = []
    delta_timestamps = []
    local_files_only = []
    dataset_meta_tasks = []
    for i in range(len(data_config)):
        repo_id = data_config[i].repo_id
        repo_ids.append(repo_id)
        data_root = assets_dirs[i]/ repo_id
        data_roots.append(data_root)
        dataset_meta = lerobot_dataset.LeRobotDatasetMetadata(repo_id, data_root, local_files_only=data_config[i].local_files_only)
        dataset_meta_tasks.append(dataset_meta.tasks)
        delta_timestamp = {key: [t / dataset_meta.fps for t in range(model_config.n_action_steps)]
            for key in data_config[i].action_sequence_keys
        }
        delta_timestamps.append(delta_timestamp)
        local_files_only.append(data_config[i].local_files_only)

    dataset = lerobot_dataset.MultiLeRobotDataset(
        repo_ids,
        data_weights,
        data_roots, 
        delta_timestamps=delta_timestamps,
        local_files_only=local_files_only,
    )
    end_time = time.time()
    execution_time = end_time - start_time
    print(f"***********Time of data load: {execution_time} 秒")
    num_frames = dataset.num_frames
    num_episodes = dataset.num_episodes

    # if data_config.prompt_from_task:
    PromptFromLeRobotTask = []
    for i in range(len(dataset_meta_tasks)):
        PromptFromLeRobotTask.append([_transforms.PromptFromLeRobotTask(dataset_meta_tasks[i])])
    task_dataset = TransformedPretrainDataset(dataset, PromptFromLeRobotTask)

    return task_dataset, dataset, num_frames, num_episodes


def transform_dataset(dataset: Dataset, data_config: _config.DataConfig, *, skip_norm_stats: bool = False) -> Dataset:
    """Transform the dataset by applying the data transforms."""
    
    return TransformedDataset_Late(
        dataset,
        [
            *data_config.repack_transforms.inputs,
            *data_config.data_transforms.inputs,
            _transforms.Normalize(data_mask=data_config.data_mask, norm_stats=data_config.norm_stats),
            *data_config.model_transforms.inputs,
        ],
    )


def transform_pretrain_dataset(dataset: Dataset, data_config: List[_config.DataConfig], *, skip_norm_stats: bool = False) -> Dataset:
    """Transform the dataset by applying the data transforms."""

    dataset_transform = []
    for i in range(len(data_config)):
        sub_transform = [
            *data_config[i].repack_transforms.inputs,
            *data_config[i].data_transforms.inputs,
            _transforms.Normalize(norm_stats=data_config[i].norm_stats, data_mask=data_config[i].data_mask),
            *data_config[i].model_transforms.inputs,
        ]
        dataset_transform.append(sub_transform)
    return TransformedPretrainDataset_Late(dataset, dataset_transform)


def create_data_loader(
    config: _config.TrainConfig,
    *,
    skip_norm_stats: bool = False,
    shuffle: bool = False,
    num_workers: int = 0,
):
    """Create a data loader for training.

    Args:
        config: The training configuration.
        sharding: The sharding to use for the data loader. If None, the data loader will
            use a single device sharding.
        skip_norm_stats: Whether to skip data normalization.
        shuffle: Whether to shuffle the data.
        num_batches: Determines the number of batches to return. If the number exceeds the
            number of batches in the dataset, the data loader will loop over the dataset.
            If not provided, will iterate over the dataset indefinitely.
        num_workers: The number of worker processes to use. If zero, the data loader will
            execute in the main process.
    """
    # if int(os.environ["LOCAL_RANK"]) == 0:
    print("The root directory of dataset:", config.assets_dirs)
    data_config = config.data.create(config.assets_dirs, config.model)
    dataset, num_frames, num_episodes = create_dataset(data_config, config.assets_dirs, config.model)
    dataset = transform_dataset(dataset, data_config, skip_norm_stats=skip_norm_stats)

    mp_context = None
    if num_workers > 0:
        mp_context = multiprocessing.get_context("spawn")
    generator = torch.Generator()
    generator.manual_seed(config.seed)

    data_loader = torch.utils.data.DataLoader(
        typing.cast(torch.utils.data.Dataset, dataset),
        batch_size=config.batch_size,
        shuffle=shuffle,
        num_workers=num_workers,
        multiprocessing_context=mp_context,
        persistent_workers=num_workers > 0,
        # multiprocessing_context=mp_context,
        # persistent_workers=num_workers > 0,
        drop_last=True,
        generator=generator,
    )


    return data_loader, num_frames, num_episodes


def create_pretrain_data_loader(
    config: _config.PretrainConfig,
    *,
    skip_norm_stats: bool = False,
    shuffle: bool = False,
    num_batches: int | None = None,
    num_workers: int = 0,
):
    """Create a data loader for training.

    Args:
        config: The training configuration.
        sharding: The sharding to use for the data loader. If None, the data loader will
            use a single device sharding.
        skip_norm_stats: Whether to skip data normalization.
        shuffle: Whether to shuffle the data.
        num_batches: Determines the number of batches to return. If the number exceeds the
            number of batches in the dataset, the data loader will loop over the dataset.
            If not provided, will iterate over the dataset indefinitely.
        num_workers: The number of worker processes to use. If zero, the data loader will
            execute in the main process.
    """
    print("The root directory of datasets:", config.assets_dirs)
    sub_data_config = []
    sub_config_assets_dirs = []
    for subconfig in config.total_configs:
        sub_data_config.append(subconfig.data.create(subconfig.assets_dirs, subconfig.model))
        sub_config_assets_dirs.append(subconfig.assets_dirs)
    # data_config = config.data.create(config.assets_dirs, config.model)

    dataset, dataset_ori, num_frames, num_episodes = create_pretrain_dataset(sub_data_config, sub_config_assets_dirs, config.model, config.data_weights)
    dataset = transform_pretrain_dataset(dataset, sub_data_config, skip_norm_stats=skip_norm_stats)

    return dataset, dataset_ori, num_frames, num_episodes


