#!/usr/bin/env python3
# -*- coding: utf-8 -*-

from __future__ import annotations

import argparse
import gzip
import io
import json
import math
import os
import random
import sys
from dataclasses import dataclass
from typing import Any, Dict, Iterable, Iterator, List, Optional

import torch
from torch.utils.data import IterableDataset, DataLoader, get_worker_info

def _open_maybe_gzip(path: str):
    if path.endswith(".gz"):
        return gzip.open(path, "rb")
    return open(path, "rb")

def _iter_jsonl(path: str) -> Iterator[Dict[str, Any]]:
    fh = _open_maybe_gzip(path)
    try:
        for raw in fh:
            line = raw.decode("utf-8") if isinstance(raw, (bytes, bytearray)) else raw
            line = line.rstrip("\n")
            if not line:
                continue
            try:
                rec = json.loads(line)
            except Exception:
                continue
            yield rec
    finally:
        fh.close()

def _buffered_shuffle(records: Iterable[Dict[str, Any]], buffer_size: int, rng: random.Random) -> Iterator[Dict[str, Any]]:
    if buffer_size <= 0:
        for r in records:
            yield r
        return
    buf: List[Dict[str, Any]] = []
    for r in records:
        buf.append(r)
        if len(buf) >= buffer_size:
            i = rng.randrange(len(buf))
            buf[i], buf[-1] = buf[-1], buf[i]
            yield buf.pop()
    while buf:
        i = rng.randrange(len(buf))
        buf[i], buf[-1] = buf[-1], buf[i]
        yield buf.pop()

@dataclass
class StreamConfig:
    shards_index: Optional[str] = None  # path to *.index.json created by packer
    shard_paths: Optional[List[str]] = None  # explicit list of shards
    shuffle_shards: bool = True
    shuffle_records_buffer: int = 10_000
    seed: int = 42
    repeat: int = 1  # repeat dataset for N epochs when used standalone

class JsonlShardDataset(IterableDataset):
    def __init__(self, cfg: StreamConfig) -> None:
        super().__init__()
        if not cfg.shard_paths and not cfg.shards_index:
            raise ValueError("Either shard_paths or shards_index must be provided")
        self.cfg = cfg

        if cfg.shards_index:
            with open(cfg.shards_index, "r", encoding="utf-8") as f:
                data = json.load(f)
            shards = data.get("shards")
            if not shards:
                raise ValueError("Index JSON missing 'shards' field")
            self.all_shards = [str(p) for p in shards]
        else:
            self.all_shards = list(cfg.shard_paths or [])

        if not self.all_shards:
            raise ValueError("No shards to read")

    def _select_shards_for_worker(self) -> List[str]:
        worker = get_worker_info()
        world_size = int(os.environ.get("WORLD_SIZE", "1"))
        rank = int(os.environ.get("RANK", "0"))

        rng = random.Random(self.cfg.seed)
        shards = list(self.all_shards)
        if self.cfg.shuffle_shards:
            rng.shuffle(shards)

        # Split shards by distributed rank first
        shards = shards[rank::world_size]

        # Then split by worker id if using multiple DataLoader workers
        if worker is not None and worker.num_workers > 1:
            wid = worker.id
            n = worker.num_workers
            shards = shards[wid::n]

        return shards

    def _iter_worker(self) -> Iterator[Dict[str, Any]]:
        rng = random.Random(self.cfg.seed + (get_worker_info().id if get_worker_info() else 0))
        shard_list = self._select_shards_for_worker()
        for _ in range(max(1, self.cfg.repeat)):
            for shard in shard_list:
                for rec in _iter_jsonl(shard):
                    yield rec

    def __iter__(self) -> Iterator[Dict[str, Any]]:
        stream = self._iter_worker()
        rng = random.Random(self.cfg.seed + 17)
        return _buffered_shuffle(stream, self.cfg.shuffle_records_buffer, rng)

def collate_dynamic_padding(batch: List[Dict[str, Any]]) -> Dict[str, torch.Tensor]:
    # Expect keys: input_ids, attention_mask, token_type_ids, mlm_labels, optional nsp_label/sop_label
    max_len = 0
    for rec in batch:
        max_len = max(max_len, len(rec["input_ids"]))

    def pad(seq: List[int], pad_val: int, L: int) -> List[int]:
        if len(seq) >= L:
            return seq[:L]
        return seq + [pad_val] * (L - len(seq))

    input_ids, attention_mask, token_type_ids, mlm_labels = [], [], [], []
    nsp_labels, sop_labels = [], []
    for rec in batch:
        input_ids.append(pad(rec["input_ids"], 0, max_len))
        attention_mask.append(pad(rec.get("attention_mask", [1] * len(rec["input_ids"])), 0, max_len))
        token_type_ids.append(pad(rec.get("token_type_ids", [0] * len(rec["input_ids"])), 0, max_len))
        mlm_labels.append(pad(rec.get("mlm_labels", [-100] * len(rec["input_ids"])), -100, max_len))
        if "nsp_label" in rec:
            nsp_labels.append(int(rec["nsp_label"]))
        if "sop_label" in rec:
            sop_labels.append(int(rec["sop_label"]))

    batch_tensors: Dict[str, torch.Tensor] = {
        "input_ids": torch.tensor(input_ids, dtype=torch.long),
        "attention_mask": torch.tensor(attention_mask, dtype=torch.long),
        "token_type_ids": torch.tensor(token_type_ids, dtype=torch.long),
        "mlm_labels": torch.tensor(mlm_labels, dtype=torch.long),
    }
    if nsp_labels:
        batch_tensors["nsp_label"] = torch.tensor(nsp_labels, dtype=torch.long)
    if sop_labels:
        batch_tensors["sop_label"] = torch.tensor(sop_labels, dtype=torch.long)
    return batch_tensors

def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="PyTorch DataLoader demo for sharded JSONL pretraining data")
    g = p.add_mutually_exclusive_group(required=True)
    g.add_argument("--shards_index", type=str, help="Path to *.index.json from packer")
    g.add_argument("--shard_paths", type=str, nargs="*", help="Explicit shard paths (jsonl or jsonl.gz)")
    p.add_argument("--batch_size", type=int, default=32)
    p.add_argument("--num_workers", type=int, default=2)
    p.add_argument("--shuffle_buffer", type=int, default=10_000)
    p.add_argument("--seed", type=int, default=42)
    p.add_argument("--repeat", type=int, default=1)
    p.add_argument("--max_batches", type=int, default=5)
    return p.parse_args()

def main() -> None:
    args = parse_args()

    cfg = StreamConfig(
        shards_index=args.shards_index,
        shard_paths=args.shard_paths,
        shuffle_shards=True,
        shuffle_records_buffer=args.shuffle_buffer,
        seed=args.seed,
        repeat=args.repeat,
    )
    ds = JsonlShardDataset(cfg)

    loader = DataLoader(
        ds,
        batch_size=args.batch_size,
        num_workers=args.num_workers,
        collate_fn=collate_dynamic_padding,
        pin_memory=True,
        prefetch_factor=2 if args.num_workers > 0 else None,
        persistent_workers=True if args.num_workers > 0 else False,
    )

    it = iter(loader)
    for i in range(args.max_batches):
        try:
            batch = next(it)
        except StopIteration:
            break
        print(f"Batch {i}:", {k: tuple(v.shape) for k, v in batch.items()})

    print("✅ DataLoader demo finished.")

if __name__ == "__main__":
    main()
