Source code for signal_dataset.dataset.view

"""Ordered dataset views with lightweight or materialized publication."""

from __future__ import annotations

from collections.abc import Sequence
from typing import TYPE_CHECKING, overload

from signal_dataset._internal.retry import RetryPolicy
from signal_dataset.config import PublicationOptions
from signal_dataset.dataset.layout import ShardEntry
from signal_dataset.record.model import AnnotationRecord, Record, RecordMetadata
from signal_dataset.storage import is_absolute_uri
from signal_dataset.storage.paths import join_uri
from signal_dataset.storage.reference import ObjectReference

if TYPE_CHECKING:
    from signal_dataset.dataset.reader import Dataset


[docs] class DatasetView(Sequence[Record]): def __init__(self, dataset: object, indices: Sequence[int]) -> None: from signal_dataset.dataset.reader import Dataset if not isinstance(dataset, Dataset): raise TypeError("dataset must be a Dataset") normalized = [] for index in indices: value = index + len(dataset) if index < 0 else index if value < 0 or value >= len(dataset): raise IndexError("selection index out of range") normalized.append(value) self.dataset = dataset self.indices = tuple(normalized) self.record_metadata = _MetadataView(self) def __len__(self) -> int: return len(self.indices) @overload def __getitem__(self, index: int) -> Record: ... @overload def __getitem__(self, index: slice) -> list[Record]: ... def __getitem__(self, index: int | slice) -> Record | list[Record]: if isinstance(index, slice): return [self[item] for item in range(*index.indices(len(self)))] return self.dataset[self.indices[index]] def _refuse_annotations_the_destination_cannot_carry(self, destination: str) -> None: """Check before writing anything, because root.json is written first. Both publication paths write the dataset and then loop over annotation sets. A destination that cannot compare-and-swap would therefore leave a readable dataset silently missing its annotations, with the failure arriving several steps after the commit that made it visible. """ if not self.dataset.annotations: return from signal_dataset.annotation.publisher import require_compare_and_swap from signal_dataset.storage import object_store_for require_compare_and_swap(object_store_for(destination, root_is_directory=True))
[docs] def publish( self, destination: str, *, dataset_id: str, snapshot_id: str, publication_options: PublicationOptions | None = None, retry: int | RetryPolicy | None = None, ) -> Dataset: from signal_dataset.annotation import publish_annotations from signal_dataset.dataset.references import publish_reference self._refuse_annotations_the_destination_cannot_carry(destination) result = publish_reference( self, destination, dataset_id=dataset_id, snapshot_id=snapshot_id, options=publication_options, storage_options=self.dataset.storage_options, shard_store=self.dataset.shard_store, retry=retry, ) for name, source in self.dataset.annotations.items(): projected = self._project(source) publish_annotations(result, name, projected, metadata=source.metadata, retry=retry) if self.dataset.annotations: from signal_dataset.dataset.reader import open return open( destination, options=self.dataset.storage_options, shard_store=self.dataset.shard_store, ) return result
[docs] def materialize( self, destination: str, *, dataset_id: str, snapshot_id: str, records_per_shard: int = 10_000, retry: int | RetryPolicy | None = None, ) -> Dataset: """Copy the view into a new dataset. `retry` is passed down to each shard write and to the publication rather than wrapping the whole loop: recovering from one failure on the last shard should not re-encode every shard before it. """ import signal_dataset as sds self._refuse_annotations_the_destination_cannot_carry(destination) if ( isinstance(records_per_shard, bool) or not isinstance(records_per_shard, int) or records_per_shard < 1 ): raise ValueError("records_per_shard must be a positive integer") descriptors = [] for first in range(0, len(self), records_per_shard): descriptors.append( sds.write_shard( self[first : first + records_per_shard], destination, work_id=f"{first:012d}", retry=retry, ) ) result = sds.publish( destination, descriptors, dataset_id=dataset_id, snapshot_id=snapshot_id, retry=retry, ) from signal_dataset.annotation import publish_annotations for name, source in self.dataset.annotations.items(): projected = self._project(source) publish_annotations(result, name, projected, metadata=source.metadata, retry=retry) if self.dataset.annotations: return sds.open(destination) return result
def _project(self, source: Sequence[AnnotationRecord]) -> list[AnnotationRecord]: projected = [] for ordinal, selected in enumerate(self.indices): item = source[selected] projected.append( AnnotationRecord( item.source_record_id, ordinal, item.status, item.values, item.fields, item.provenance, item.detail_status, ) ) return projected def _entries(self) -> tuple[ShardEntry, ...]: return self._entries_range(0, len(self)) def _entries_range(self, start: int, stop: int) -> tuple[ShardEntry, ...]: def resolved(path: str) -> str: return path if is_absolute_uri(path) else join_uri(self.dataset.root, path) entries = [] for ordinal in range(start, stop): source = self.indices[ordinal] shard, local = self.dataset._index.locate(source) # Copied from the source, but without its version: a reference # published today is a new document and records no generation, the # same as every other write path. data = ObjectReference( resolved(shard.data.path), None, shard.data.stored_bytes, ) metadata = ObjectReference( resolved(shard.metadata.path), None, shard.metadata.stored_bytes, ) entries.append( ShardEntry( ordinal, 1, data, metadata, (local,), ) ) return tuple(entries)
class _MetadataView(Sequence[RecordMetadata]): def __init__(self, view: DatasetView) -> None: self._view = view def __len__(self) -> int: return len(self._view) @overload def __getitem__(self, index: int) -> RecordMetadata: ... @overload def __getitem__(self, index: slice) -> list[RecordMetadata]: ... def __getitem__(self, index: int | slice) -> RecordMetadata | list[RecordMetadata]: if isinstance(index, slice): return [self[item] for item in range(*index.indices(len(self)))] return self._view.dataset.record_metadata[self._view.indices[index]]