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