Source code for signal_dataset.record.validation

"""Universal validation independent of signal modality."""

from __future__ import annotations

import json
from typing import Any

import numpy as np

from signal_dataset._internal.json import thaw_json
from signal_dataset._internal.policy import ResourcePolicy
from signal_dataset.errors import ValidationError
from signal_dataset.record.model import Axis, Coordinate, Field, Record

SUPPORTED_DTYPES = frozenset(
    np.dtype(name)
    for name in (
        "bool",
        "int8",
        "int16",
        "int32",
        "int64",
        "uint8",
        "uint16",
        "uint32",
        "uint64",
        "float16",
        "float32",
        "float64",
        "complex64",
    )
)


def _json_size(value: Any, subject: str) -> int:
    try:
        return len(json.dumps(thaw_json(value), ensure_ascii=False, allow_nan=False).encode())
    except (TypeError, ValueError, OverflowError) as exc:
        raise ValidationError(f"{subject} must be JSON-compatible: {exc}") from exc


def _validate_dtype(name: str, array: np.ndarray) -> None:
    if array.dtype == np.dtype("O"):
        raise ValidationError(
            f"field {name!r} uses object dtype; encode values as a numerical array or JSON metadata"
        )
    if array.dtype == np.dtype("complex128"):
        raise ValidationError(
            f"field {name!r} uses complex128; explicitly convert to complex64 for version 0.1"
        )
    if array.dtype not in SUPPORTED_DTYPES:
        raise ValidationError(f"field {name!r} has unsupported dtype {array.dtype}")


[docs] def validate(record: Record, *, policy: ResourcePolicy | None = None) -> None: """Validate structural, reference, JSON, dtype, and resource invariants.""" limits = policy or ResourcePolicy() if not isinstance(record.id, str) or not record.id: raise ValidationError("record id must be a non-empty string") if record.scene_id is not None and not isinstance(record.scene_id, str): raise ValidationError("scene id must be a string or None") if not record.fields: raise ValidationError("record must contain at least one field") if len(record.fields) > limits.max_fields: raise ValidationError(f"record exceeds the {limits.max_fields}-field policy") metadata_bytes = _json_size(dict(record.metadata), "record metadata") tensor_bytes = 0 shared_axes: dict[str, tuple[int, object]] = {} axis_ids: set[str] = set() for field_name, field_value in record.fields.items(): if not isinstance(field_name, str) or not field_name: raise ValidationError("field names must be non-empty strings") if not isinstance(field_value, Field): raise ValidationError(f"field {field_name!r} must be a Field") array = field_value.data _validate_dtype(field_name, array) if array.ndim != len(field_value.axes): raise ValidationError( f"field {field_name!r} rank {array.ndim} does not match " f"its {len(field_value.axes)} axes" ) if len(field_value.axes) > limits.max_axes_per_field: raise ValidationError(f"field {field_name!r} exceeds max_axes_per_field") if array.nbytes > limits.max_tensor_bytes: raise ValidationError(f"field {field_name!r} exceeds max_tensor_bytes") tensor_bytes += array.nbytes metadata_bytes += _json_size(dict(field_value.metadata), f"field {field_name!r} metadata") axis_names: set[str] = set() for dimension, axis in enumerate(field_value.axes): if not isinstance(axis, Axis): raise ValidationError(f"field {field_name!r} axes must contain Axis values") if not isinstance(axis.name, str) or not axis.name: raise ValidationError(f"field {field_name!r} has an invalid axis name") if axis.name in axis_names: raise ValidationError(f"field {field_name!r} repeats axis name {axis.name!r}") axis_names.add(axis.name) if isinstance(axis.length, bool) or not isinstance(axis.length, int) or axis.length < 0: raise ValidationError("axis length must be a nonnegative integer") for attribute in ("role", "id"): value = getattr(axis, attribute) if value is not None and not isinstance(value, str): raise ValidationError(f"axis {attribute} must be a string or None") if axis.length != array.shape[dimension]: raise ValidationError( f"field {field_name!r} axis {axis.name!r} length does not match dimension" ) metadata_bytes += _json_size(dict(axis.metadata), f"axis {axis.name!r} metadata") coordinate = axis.coordinate if coordinate is not None: if not isinstance(coordinate, Coordinate): raise ValidationError("axis coordinate must be a Coordinate or None") if coordinate.unit is not None and not isinstance(coordinate.unit, str): raise ValidationError("coordinate unit must be a string or None") for coordinate_name in ("start", "step"): coordinate_value = getattr(coordinate, coordinate_name) if coordinate_value is not None and ( isinstance(coordinate_value, bool) or not isinstance(coordinate_value, (int, float)) ): raise ValidationError( f"coordinate {coordinate_name} must be a number or None" ) explicit = coordinate.values is not None linear = coordinate.start is not None or coordinate.step is not None referenced = coordinate.reference is not None if sum((explicit, linear, referenced)) > 1: raise ValidationError( f"axis {axis.name!r} coordinate forms are mutually exclusive" ) if (coordinate.start is None) != (coordinate.step is None): raise ValidationError( f"axis {axis.name!r} linear coordinate requires both start and step" ) if explicit and len(coordinate.values or ()) != axis.length: raise ValidationError( f"axis {axis.name!r} explicit coordinate length does not match axis" ) if explicit and len(coordinate.values or ()) > limits.max_coordinate_values: raise ValidationError(f"axis {axis.name!r} exceeds max_coordinate_values") metadata_bytes += _json_size( { "values": coordinate.values, "start": coordinate.start, "step": coordinate.step, "reference": coordinate.reference, "unit": coordinate.unit, }, f"axis {axis.name!r} coordinate", ) metadata_bytes += _json_size( dict(coordinate.metadata), f"axis {axis.name!r} coordinate metadata" ) if axis.id: signature = (axis.length, coordinate) previous = shared_axes.setdefault(axis.id, signature) if previous != signature: raise ValidationError(f"shared axis {axis.id!r} declarations disagree") axis_ids.add(axis.id) for field_value in record.fields.values(): for axis in field_value.axes: reference = axis.coordinate.reference if axis.coordinate else None if reference is not None and not isinstance(reference, str): raise ValidationError("coordinate references must be strings") if reference and reference not in record.fields and reference not in axis_ids: raise ValidationError(f"coordinate reference {reference!r} does not resolve") if tensor_bytes > limits.max_total_tensor_bytes: raise ValidationError("record exceeds max_total_tensor_bytes") if metadata_bytes > limits.max_metadata_bytes: raise ValidationError("record exceeds max_metadata_bytes")