Source code for signal_dataset.record.model

"""Immutable, use-case-agnostic logical primitives."""

from __future__ import annotations

from collections.abc import Iterator, Mapping
from dataclasses import dataclass, field
from enum import StrEnum
from types import MappingProxyType

import numpy as np

from signal_dataset._internal.json import Json, freeze_json, frozen_mapping


[docs] @dataclass(frozen=True, slots=True) class Coordinate: values: tuple[Json, ...] | None = None start: int | float | None = None step: int | float | None = None reference: str | None = None unit: str | None = None metadata: Mapping[str, Json] = field(default_factory=dict) def __post_init__(self) -> None: if self.values is not None and not isinstance(self.values, tuple): object.__setattr__(self, "values", tuple(self.values)) if self.values is not None: object.__setattr__(self, "values", tuple(freeze_json(item) for item in self.values)) object.__setattr__(self, "metadata", frozen_mapping(self.metadata))
[docs] @dataclass(frozen=True, slots=True) class Axis: name: str length: int role: str | None = None id: str | None = None coordinate: Coordinate | None = None metadata: Mapping[str, Json] = field(default_factory=dict) def __post_init__(self) -> None: object.__setattr__(self, "metadata", frozen_mapping(self.metadata))
[docs] @dataclass(frozen=True, slots=True, eq=False) class Field: data: np.ndarray axes: tuple[Axis, ...] = () metadata: Mapping[str, Json] = field(default_factory=dict) def __post_init__(self) -> None: source = np.asarray(self.data) if source.dtype.hasobject: raise ValueError("object arrays are unsupported; convert values to numeric tensors") contiguous = np.ascontiguousarray(source) # An immutable bytes owner prevents callers from re-enabling ndarray writes. data = np.frombuffer(contiguous.tobytes(order="C"), dtype=contiguous.dtype).reshape( contiguous.shape ) object.__setattr__(self, "data", data) object.__setattr__(self, "axes", tuple(self.axes)) object.__setattr__(self, "metadata", frozen_mapping(self.metadata))
[docs] @dataclass(frozen=True, slots=True, eq=False) class Record(Mapping[str, Field]): id: str fields: Mapping[str, Field] scene_id: str | None = None metadata: Mapping[str, Json] = field(default_factory=dict) def __post_init__(self) -> None: object.__setattr__(self, "fields", MappingProxyType(dict(self.fields))) object.__setattr__(self, "metadata", frozen_mapping(self.metadata)) def __getitem__(self, name: str) -> Field: return self.fields[name] def __iter__(self) -> Iterator[str]: return iter(self.fields) def __len__(self) -> int: return len(self.fields)
[docs] @dataclass(frozen=True, slots=True) class FieldMetadata: dtype: str shape: tuple[int, ...] axes: tuple[Axis, ...] = () metadata: Mapping[str, Json] = field(default_factory=dict) def __post_init__(self) -> None: object.__setattr__(self, "shape", tuple(self.shape)) object.__setattr__(self, "axes", tuple(self.axes)) object.__setattr__(self, "metadata", frozen_mapping(self.metadata))
[docs] @dataclass(frozen=True, slots=True) class RecordMetadata(Mapping[str, FieldMetadata]): id: str fields: Mapping[str, FieldMetadata] scene_id: str | None = None metadata: Mapping[str, Json] = field(default_factory=dict) def __post_init__(self) -> None: object.__setattr__(self, "fields", MappingProxyType(dict(self.fields))) object.__setattr__(self, "metadata", frozen_mapping(self.metadata)) def __getitem__(self, name: str) -> FieldMetadata: return self.fields[name] def __iter__(self) -> Iterator[str]: return iter(self.fields) def __len__(self) -> int: return len(self.fields)
[docs] @classmethod def from_record(cls, record: Record) -> RecordMetadata: fields = { name: FieldMetadata( dtype=item.data.dtype.str, shape=item.data.shape, axes=item.axes, metadata=item.metadata, ) for name, item in record.fields.items() } return cls(record.id, fields, record.scene_id, record.metadata)
[docs] @dataclass(frozen=True, slots=True) class PublishedShard: data_uri: str metadata_uri: str work_id: str attempt: int record_count: int data_bytes: int metadata_bytes: int data_generation: int | None = None metadata_generation: int | None = None def __post_init__(self) -> None: if not self.data_uri or not self.metadata_uri or not self.work_id: raise ValueError("shard URIs and work_id must be non-empty") for name in ( "attempt", "record_count", "data_bytes", "metadata_bytes", ): value = getattr(self, name) if isinstance(value, bool) or not isinstance(value, int) or value < 0: raise ValueError(f"shard {name} must be a nonnegative integer")
[docs] class AnnotationStatus(StrEnum): SUCCESS = "success" SKIPPED = "skipped" FAILED = "failed"
[docs] @dataclass(frozen=True, slots=True, eq=False) class AnnotationRecord: source_record_id: str source_index: int status: AnnotationStatus values: Mapping[str, Json] = field(default_factory=dict) fields: Mapping[str, Field] = field(default_factory=dict) provenance: Mapping[str, Json] = field(default_factory=dict) detail_status: str | None = None def __post_init__(self) -> None: if self.source_index < 0: raise ValueError("source_index must be nonnegative") if self.detail_status is not None and not self.detail_status: raise ValueError("detail_status must be non-empty when provided") object.__setattr__(self, "values", frozen_mapping(self.values)) object.__setattr__(self, "fields", MappingProxyType(dict(self.fields))) object.__setattr__(self, "provenance", frozen_mapping(self.provenance))