# SPDX-License-Identifier: Apache-2.0 """Host tests (no device, no weights): the request decoders shared by the API and the server, i.e. the vendored ``ttaw.io`` as this bundle binds it (``io.load_inputs`` / ``io.decode_inputs`` with ``INPUT_SCHEMA``; the full suite of ttaw.io is common/tests/host/test_io_outputs_host.py). TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_io_host.py """ from __future__ import annotations import base64 import io from pathlib import Path import numpy as np import pytest from tt_diffusion_planner import io as tio from tt_diffusion_planner.api import DiffusionPlanner from tt_diffusion_planner.reference import config as C from tt_diffusion_planner.tests.stubs import sample_inputs SAMPLES = ("kashiwanoha_dense", "straight_road") def _npz_b64(arrays) -> str: buf = io.BytesIO() np.savez_compressed(buf, **arrays) return base64.b64encode(buf.getvalue()).decode("ascii") def test_schema_is_the_node_input_map(): """The 15 raw tensors of ``DiffusionPlannerCore::create_input_data`` (core.cpp:414-595), batch 1, float32.""" assert DiffusionPlanner.INPUT_SCHEMA is C.INPUT_SCHEMA and DiffusionPlanner.INPUT_KIND == "planner" assert set(C.INPUT_SCHEMA) == { "sampled_trajectories", "ego_agent_past", "ego_current_state", "neighbor_agents_past", "static_objects", "lanes", "lanes_speed_limit", "route_lanes", "route_lanes_speed_limit", "polygons", "line_strings", "goal_pose", "ego_shape", "turn_indicators", "delay"} assert C.INPUT_SCHEMA["neighbor_agents_past"][0] == (1, 320, 31, 11) assert C.INPUT_SCHEMA["lanes"][0] == (1, 140, 20, 33) and C.INPUT_SCHEMA["line_strings"][0] == (1, 60, 20, 4) assert all(np.dtype(dt) == np.float32 for _, dt in C.INPUT_SCHEMA.values()) assert tio.DEFAULT_POINT_FIELDS == () and DiffusionPlanner.POINT_FIELDS == () @pytest.mark.parametrize("stem", SAMPLES) def test_shipped_samples_load_in_every_form(stem): path = Path(__file__).resolve().parents[1] / "samples" / f"{stem}.npz" arrays = tio.load_inputs(str(path)) assert set(arrays) == set(C.INPUT_NAMES) and all(a.dtype == np.float32 for a in arrays.values()) for form in (path.read_bytes(), {k: v.astype(np.float64) for k, v in arrays.items()}, {"format": "npz", "data": base64.b64encode(path.read_bytes()).decode()}): again = tio.load_inputs(form) for k in C.INPUT_NAMES: np.testing.assert_array_equal(again[k], arrays[k]) dec = tio.decode_inputs({"format": "npz", "data": _npz_b64(arrays)}) np.testing.assert_array_equal(dec["lanes"], arrays["lanes"]) def test_json_arrays_form(): raw = sample_inputs() spec = {"format": "json", "arrays": {k: v.tolist() for k, v in raw.items()}} got = tio.decode_inputs(spec) for k in C.INPUT_NAMES: np.testing.assert_array_equal(got[k], raw[k]) @pytest.mark.parametrize("mutate,match", [ (lambda r: r.pop("delay"), "missing"), (lambda r: r.update(extra=np.zeros(3, np.float32)), "unexpected"), (lambda r: r.update(lanes=np.zeros((1, 70, 20, 33), np.float32)), "shape"), (lambda r: r.update(ego_shape=np.zeros(3, np.float32)), "shape"), (lambda r: r["neighbor_agents_past"].__setitem__((0, 0, 0, 0), np.nan), "non-finite"), (lambda r: r["goal_pose"].__setitem__((0, 0), np.inf), "non-finite"), ]) def test_invalid_inputs_raise_input_error(mutate, match): raw = sample_inputs() mutate(raw) with pytest.raises(tio.InputError, match=match): tio.load_inputs(raw) with pytest.raises(tio.InputError, match=match): tio.decode_inputs({"format": "npz", "data": _npz_b64(raw)}) def test_undecodable_envelopes(): with pytest.raises(tio.InputError): tio.decode_inputs({"format": "npz", "data": "not base64!"}) with pytest.raises(tio.InputError): tio.decode_inputs({"format": "npz", "data": base64.b64encode(b"not an npz").decode()}) with pytest.raises(tio.InputError, match="format"): tio.decode_inputs({"format": "pickle", "data": ""}) with pytest.raises(tio.InputError): tio.load_inputs(42) @pytest.mark.parametrize("fmt", ["npz", "npz_compressed", "raw", "list"]) def test_encode_array_roundtrip(fmt): a = np.arange(2 * 80 * 5, dtype=np.float32).reshape(2, 80, 5) d = tio.encode_array(a, fmt=fmt, key="predicted_agents") if fmt.startswith("npz"): back = np.load(io.BytesIO(base64.b64decode(d["data"])))["predicted_agents"] elif fmt == "raw": back = np.frombuffer(base64.b64decode(d["data"]), dtype=d["dtype"]).reshape(d["shape"]) else: back = np.asarray(d["data"], dtype=d["dtype"]) np.testing.assert_array_equal(back, a)