from __future__ import annotations import random from collections.abc import Iterator, Sequence from pathlib import Path from typing import Any from .common import write_jsonl class DataObject(Sequence[dict[str, str]]): def __init__( self, records: Sequence[dict[str, str]] | None = None, *, name: str = "dataset", warnings: Sequence[str] | None = None, ) -> None: self.records = list(records or []) self.name = name self.warnings = list(warnings or []) def __len__(self) -> int: return len(self.records) def __iter__(self) -> Iterator[dict[str, str]]: return iter(self.records) def __getitem__(self, index: int | slice) -> dict[str, str] | list[dict[str, str]]: return self.records[index] def take(self, amount: int) -> list[dict[str, str]]: if amount < 0: raise ValueError("amount must be non-negative") return self.records[:amount] def shuffle(self, *, seed: int | None = 42) -> "DataObject": shuffled = list(self.records) rng = random.Random(seed) rng.shuffle(shuffled) return DataObject( shuffled, name=self.name, warnings=self.warnings, ) def to_jsonl(self, path: str | Path) -> None: write_jsonl(self.records, path) def extend(self, other: "DataObject") -> None: self.records.extend(other.records) self.warnings.extend(other.warnings) @classmethod def concat(cls, datasets: Sequence["DataObject"], *, name: str = "all") -> "DataObject": records: list[dict[str, str]] = [] warnings: list[str] = [] for dataset in datasets: records.extend(dataset.records) warnings.extend(dataset.warnings) return cls(records, name=name, warnings=warnings) def summary(self) -> dict[str, Any]: return { "name": self.name, "rows": len(self.records), "warnings": self.warnings, }