"""Dataset registry: decorator registration + name lookup with suggestions. Registration happens at import time ("import = register"). The registry maps name -> DatasetProvider; a Dataset is only *materialized* on first use. """ import difflib from typing import Callable, Dict, List, Optional, Union from .spec import DatasetSpec, FieldSpec # A provider factory: called once, returns how records become Samples. ProviderFactory = Callable[[], Union[FieldSpec, Callable, None]] class DatasetProvider: """A registered dataset: metadata + the recipe for converting records.""" def __init__(self, spec: DatasetSpec, factory: Optional[ProviderFactory] = None): self.spec = spec self._factory = factory self._record_fn = None self._resolved = False def resolve_record_fn(self) -> Callable: """Resolve the record->Sample converter, lazily and exactly once.""" if not self._resolved: result = self._factory() if self._factory else None if isinstance(result, FieldSpec): from .loader import field_spec_to_record_fn self._record_fn = field_spec_to_record_fn(result) elif callable(result): self._record_fn = result elif result is None: self._record_fn = field_spec_to_record_fn(FieldSpec()) else: raise TypeError(f'{self.spec.name}: factory must return FieldSpec or callable, got {type(result)}') self._resolved = True return self._record_fn class Registry: """Minimal dict-like registry with duplicate protection and suggestions.""" def __init__(self, kind: str): self.kind = kind self._items: Dict[str, DatasetProvider] = {} def register(self, name: str, item: DatasetProvider) -> DatasetProvider: if name in self._items: raise ValueError(f'{self.kind} {name!r} is already registered') self._items[name] = item return item def get(self, name: str) -> DatasetProvider: if name not in self._items: suggestions = difflib.get_close_matches(name, self._items.keys(), n=3) hint = f" Did you mean: {', '.join(suggestions)}?" if suggestions else '' raise KeyError(f'unknown {self.kind} {name!r}.{hint}') return self._items[name] def names(self) -> List[str]: return sorted(self._items) def __contains__(self, name: str) -> bool: return name in self._items def __len__(self) -> int: return len(self._items) DATASET_REGISTRY = Registry('dataset') def register_dataset(spec: DatasetSpec): """Decorator: register a dataset plugin. Usage: @register_dataset(DatasetSpec(name='gsm8k', source=...)) def gsm8k(): return lambda record: Sample(...) # or FieldSpec(...), or None """ def decorator(factory: ProviderFactory) -> ProviderFactory: DATASET_REGISTRY.register(spec.name, DatasetProvider(spec, factory)) return factory return decorator def get_dataset_provider(name: str) -> DatasetProvider: return DATASET_REGISTRY.get(name)