Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions src/cast_value/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,20 @@ class ScalarMapJSON(TypedDict):
decode: NotRequired[list[tuple[object, object]]]


class ScalarMapLike(TypedDict):
"""
Accepted input forms for the ``scalar_map`` codec parameter.

Each direction may be given either as a mapping of source -> target or as
an iterable of ``(source, target)`` pairs. Both are normalized to
[`ScalarMapJSON`][cast_value.types.ScalarMapJSON] -- the canonical form the
cast_value spec uses -- at codec construction time.
"""

encode: NotRequired[Mapping[object, object] | Iterable[tuple[object, object]]]
decode: NotRequired[Mapping[object, object] | Iterable[tuple[object, object]]]


# Pre-parsed scalar map entry: (source_scalar, target_scalar)
ScalarMapEntry = tuple[NumericScalar, NumericScalar]

Expand Down
57 changes: 54 additions & 3 deletions src/cast_value/zarr_compat/_parsing.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,14 +6,65 @@

from __future__ import annotations

from collections.abc import Mapping
from typing import TYPE_CHECKING

if TYPE_CHECKING:
from collections.abc import Mapping

from zarr.core.dtype.wrapper import TBaseDType, TBaseScalar, ZDType

from cast_value.types import ScalarMapEntry, ScalarMapJSON
from cast_value.types import ScalarMapEntry, ScalarMapJSON, ScalarMapLike

_DIRECTIONS = ("encode", "decode")


def parse_scalar_map(
data: ScalarMapJSON | ScalarMapLike | None,
) -> ScalarMapJSON | None:
"""Normalize a scalar map to its canonical JSON form.

For each of the ``"encode"`` and ``"decode"`` keys, accepts either a
mapping of source -> target or an iterable of ``(source, target)`` pairs,
and normalizes to a list of pairs -- the form the cast_value spec uses
and ``to_dict`` serializes. Malformed maps raise here, at codec
construction time, rather than later at encode/decode time.
"""
if data is None:
return None
if not isinstance(data, Mapping):
msg = f"scalar_map must be a mapping, got {type(data).__name__}"
raise TypeError(msg)
unknown = {key for key in data if key not in _DIRECTIONS}
if unknown:
msg = (
f"scalar_map keys must be 'encode' or 'decode', "
f"got {sorted(map(str, unknown))}"
)
raise ValueError(msg)
result: ScalarMapJSON = {}
for direction in _DIRECTIONS:
if direction not in data:
continue
pairs = data[direction] # type: ignore[literal-required]
items = pairs.items() if isinstance(pairs, Mapping) else pairs
entries: list[tuple[object, object]] = []
for entry in items:
try:
pair = tuple(entry)
except TypeError:
msg = (
f"scalar_map {direction!r} entry {entry!r} is not a "
f"(source, target) pair"
)
raise TypeError(msg) from None
if len(pair) != 2:
msg = (
f"scalar_map {direction!r} entry {entry!r} must have "
f"exactly 2 elements, got {len(pair)}"
)
raise ValueError(msg)
entries.append((pair[0], pair[1]))
result[direction] = entries # type: ignore[literal-required]
return result


def extract_raw_map(
Expand Down
16 changes: 12 additions & 4 deletions src/cast_value/zarr_compat/v1/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,11 @@
from zarr.core.common import JSON, parse_named_configuration
from zarr.core.dtype import get_data_type_from_json

from cast_value.zarr_compat._parsing import extract_raw_map, parse_map_entries
from cast_value.zarr_compat._parsing import (
extract_raw_map,
parse_map_entries,
parse_scalar_map,
)

if TYPE_CHECKING:
from collections.abc import Iterable
Expand All @@ -26,6 +30,7 @@
RoundingMode,
ScalarMapEntry,
ScalarMapJSON,
ScalarMapLike,
)


Expand All @@ -48,7 +53,10 @@ class _CastValueBaseV1(ArrayArrayCodec):
None means error. "clamp" clips to range. "wrap" uses modular arithmetic
(only valid for integer types).
scalar_map : dict or None
Explicit value overrides as JSON: {"encode": [[src, tgt], ...], "decode": [[src, tgt], ...]}.
Explicit value overrides. Each of the optional "encode"/"decode" keys
accepts either a mapping of source -> target or an iterable of
(source, target) pairs; both are normalized to the spec's
list-of-pairs form: {"encode": [[src, tgt], ...], "decode": [[src, tgt], ...]}.
"""

is_fixed_size = True
Expand All @@ -64,7 +72,7 @@ def __init__(
data_type: str | ZDType[TBaseDType, TBaseScalar],
rounding: RoundingMode = "nearest-even",
out_of_range: OutOfRangeMode | None = None,
scalar_map: ScalarMapJSON | None = None,
scalar_map: ScalarMapJSON | ScalarMapLike | None = None,
) -> None:
if isinstance(data_type, str):
dtype = get_data_type_from_json(data_type, zarr_format=3)
Expand All @@ -73,7 +81,7 @@ def __init__(
object.__setattr__(self, "dtype", dtype)
object.__setattr__(self, "rounding", rounding)
object.__setattr__(self, "out_of_range", out_of_range)
object.__setattr__(self, "scalar_map", scalar_map)
object.__setattr__(self, "scalar_map", parse_scalar_map(scalar_map))

@classmethod
def from_dict(cls, data: dict[str, JSON]) -> Self:
Expand Down
103 changes: 103 additions & 0 deletions tests/zarr_compat/v1/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from zarr.core.chunk_grids import RegularChunkGrid
from zarr.core.dtype import get_data_type_from_json

from cast_value.zarr_compat._parsing import parse_scalar_map
from cast_value.zarr_compat.v1 import CastValueNumpyV1, parse_map_entries
from zarr_compat.v1._helpers import arrays_bytes_equal, make_spec

Expand Down Expand Up @@ -252,6 +253,108 @@ def test_parse_map_entries(
assert rt == et


# ---------------------------------------------------------------------------
# parse_scalar_map
# ---------------------------------------------------------------------------


@pytest.mark.parametrize(
"case",
[
Expect(id="none", input=None, expected=None),
Expect(
id="dict-form",
input={"encode": {"NaN": -32768}, "decode": {-32768: "NaN"}},
expected={"encode": [("NaN", -32768)], "decode": [(-32768, "NaN")]},
),
Expect(
id="pairs-of-tuples",
input={"encode": [("NaN", 0)]},
expected={"encode": [("NaN", 0)]},
),
Expect(
id="pairs-of-lists",
input={"encode": [["NaN", 0]]},
expected={"encode": [("NaN", 0)]},
),
Expect(
id="mixed-forms",
input={"encode": {"NaN": 0}, "decode": [(0, "NaN")]},
expected={"encode": [("NaN", 0)], "decode": [(0, "NaN")]},
),
Expect(id="empty-map", input={}, expected={}),
Expect(
id="empty-direction",
input={"encode": {}},
expected={"encode": []},
),
],
ids=lambda case: case.id,
)
def test_parse_scalar_map(case: Expect[Any, Any]) -> None:
"""Test that parse_scalar_map normalizes every accepted input form to the
spec's list-of-pairs form."""
assert parse_scalar_map(case.input) == case.expected


def test_parse_scalar_map_not_a_mapping() -> None:
"""Test that parse_scalar_map rejects a scalar_map that is not a mapping."""
with pytest.raises(TypeError, match="must be a mapping"):
parse_scalar_map([("NaN", 0)]) # ty: ignore[invalid-argument-type]


def test_parse_scalar_map_unknown_key() -> None:
"""Test that parse_scalar_map rejects direction keys other than
'encode'/'decode'."""
with pytest.raises(ValueError, match="keys must be 'encode' or 'decode'"):
parse_scalar_map({"Encode": [("NaN", 0)]}) # ty: ignore[invalid-argument-type]


def test_parse_scalar_map_entry_wrong_arity() -> None:
"""Test that parse_scalar_map rejects entries that are not 2-element pairs."""
with pytest.raises(ValueError, match="exactly 2 elements"):
parse_scalar_map({"encode": [(1, 2, 3)]}) # ty: ignore[invalid-argument-type]


def test_parse_scalar_map_entry_not_a_pair() -> None:
"""Test that parse_scalar_map rejects entries that cannot be unpacked
into a pair."""
with pytest.raises(TypeError, match="is not a"):
parse_scalar_map({"encode": [5]}) # ty: ignore[invalid-argument-type]


def test_init_normalizes_dict_scalar_map() -> None:
"""Test that dict-form scalar_map is normalized at construction so the
codec works and serializes to the spec's list-of-pairs form.

Regression test for zarr-developers/cast-value.rs#24: dict-form maps were
accepted at construction but crashed with 'too many values to unpack' when
the codec was used.
"""
codec = CastValueNumpyV1(
data_type="int16",
scalar_map={"encode": {"NaN": -32768}, "decode": {-32768: "NaN"}},
)
assert codec.scalar_map == {
"encode": [("NaN", -32768)],
"decode": [(-32768, "NaN")],
}
config = codec.to_dict()["configuration"]
assert isinstance(config, dict)
assert config["scalar_map"] == {
"encode": [("NaN", -32768)],
"decode": [(-32768, "NaN")],
}
# The crash fired on first use; exercise the encode path.
spec = make_spec("float32", 0, shape=(3,))
arr = np.array([1.0, np.nan, 3.0], dtype=np.float32)
buf = NDBuffer.from_ndarray_like(arr) # ty: ignore[invalid-argument-type]
result_buf = asyncio.run(codec._encode_single(buf, spec))
assert result_buf is not None
result = np.asarray(result_buf.as_ndarray_like())
assert arrays_bytes_equal(result, np.array([1, -32768, 3], dtype=np.int16))


# ---------------------------------------------------------------------------
# compute_encoded_size
# ---------------------------------------------------------------------------
Expand Down
Loading