"""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]]