diff --git a/src/hflow/episode.py b/src/hflow/episode.py index dd956228..93b9fad9 100644 --- a/src/hflow/episode.py +++ b/src/hflow/episode.py @@ -136,7 +136,9 @@ def _is_numeric_sequence(value: Any) -> bool: if isinstance(value, np.ndarray): return value.size > 0 and value.dtype.kind in "iuf" if isinstance(value, (list, tuple)): - return len(value) > 0 and _is_numeric_scalar(value[0]) + for item in value: + if item is not None: + return _is_numeric_scalar(item) return False @@ -154,7 +156,9 @@ def _is_arrow_scalar(value: Any) -> bool: def _is_empty_sequence(value: Any) -> bool: if isinstance(value, np.ndarray): return value.size == 0 - return isinstance(value, (list, tuple)) and len(value) == 0 + if isinstance(value, (list, tuple)): + return len(value) == 0 or all(item is None for item in value) + return False def _arrow_field_values(messages: Sequence[Any]) -> dict[str, list[Any]]: @@ -213,7 +217,7 @@ def _arrow_column_values(topic: str, field_name: str, values: Sequence[Any]) -> # An empty numeric array keeps its dtype, so it types the column. saw_list = True elif _is_empty_sequence(value): - # `[]` carries no element type; an empty numeric ndarray above + # `[]` or an all-null sequence carries no element type; an empty numeric ndarray above # does. Keep the slot null unless a typed sample arrives. continue else: diff --git a/tests/test_episode.py b/tests/test_episode.py index 027d502c..e9c05eeb 100644 --- a/tests/test_episode.py +++ b/tests/test_episode.py @@ -266,3 +266,46 @@ def test_empty_channel_to_arrow(tmp_path: Path) -> None: assert table.num_rows == 0 assert table.column_names == ["log_time_ns"] assert table.schema.field("log_time_ns").type == pyarrow.int64() + + +def test_channel_to_arrow_numeric_list_with_null_elements() -> None: + import json + + import numpy as np + + from hflow.episode import ChannelData + from hflow.reader import TopicInfo + + info = TopicInfo( + topic="/arm/joint_state", + channel_id=1, + schema_name="JointState", + schema_encoding="jsonschema", + message_encoding="json", + message_count=4, + schema_data=b"{}", + ) + cd = ChannelData( + topic="/arm/joint_state", + channel_id=1, + info=info, + log_times=np.array([1000, 2000, 3000, 4000], dtype=np.int64), + publish_times=np.array([1000, 2000, 3000, 4000], dtype=np.int64), + raw=[ + b'{"position": [1.0, 2.0, null], "untyped": [null, null]}', + b'{"position": [null, 2.0, 3.0], "untyped": [null, null]}', + b'{"position": [null, null, null], "untyped": [null, null]}', + b'{"position": [1.0, 2.0, 3.0], "untyped": [null, null]}', + ], + decoder=lambda b: json.loads(b.decode()), + ) + table = cd.to_arrow() + assert "position" in table.column_names + # Untyped all-null list column carries no element type and is omitted like an empty list + assert "untyped" not in table.column_names + assert table["position"].to_pylist() == [ + [1.0, 2.0, None], + [None, 2.0, 3.0], + [None, None, None], + [1.0, 2.0, 3.0], + ]