diff --git a/src/cast_value/types.py b/src/cast_value/types.py index 357428b..02dfbec 100644 --- a/src/cast_value/types.py +++ b/src/cast_value/types.py @@ -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] diff --git a/src/cast_value/zarr_compat/_parsing.py b/src/cast_value/zarr_compat/_parsing.py index 599e368..49a0c2f 100644 --- a/src/cast_value/zarr_compat/_parsing.py +++ b/src/cast_value/zarr_compat/_parsing.py @@ -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( diff --git a/src/cast_value/zarr_compat/v1/_base.py b/src/cast_value/zarr_compat/v1/_base.py index d5f56a3..5125f22 100644 --- a/src/cast_value/zarr_compat/v1/_base.py +++ b/src/cast_value/zarr_compat/v1/_base.py @@ -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 @@ -26,6 +30,7 @@ RoundingMode, ScalarMapEntry, ScalarMapJSON, + ScalarMapLike, ) @@ -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 @@ -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) @@ -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: diff --git a/tests/zarr_compat/v1/test_base.py b/tests/zarr_compat/v1/test_base.py index ce08bcc..3eee148 100644 --- a/tests/zarr_compat/v1/test_base.py +++ b/tests/zarr_compat/v1/test_base.py @@ -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 @@ -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 # ---------------------------------------------------------------------------