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
5 changes: 3 additions & 2 deletions docs/serializers/arrow.md
Original file line number Diff line number Diff line change
Expand Up @@ -99,8 +99,9 @@ ArrowSerializer supports:
- Scalar values (int, str, float)
- Lists of objects
- Dicts with a number or bool value (`{"id": 1}`)
- Dicts with a string, bytes or dict value (`{"name": "Alice"}`, `{"user": {"name": "x"}}`)

Every dict value must be a list or an array. A string or dict value is not rejected, but it is stored wrong: a string becomes a column of its characters, and a nested dict a column of its keys, with its values lost. Flatten nested dicts into columns first, or use [AutoSerializer](./auto.md).
Every dict value must be a column Arrow can convert: a list, tuple, NumPy array, pandas Series, pyarrow array or typed `memoryview`, or any object that implements the Arrow array protocol (`__arrow_array__` or `__arrow_c_array__`), a `Mapping` included. A `memoryview` is read as its element type, so a memoryview of `bytes` stores byte values. Flatten nested dicts into columns first, or use [AutoSerializer](./auto.md).

**Type checking example:**
```python
Expand All @@ -117,7 +118,7 @@ try:
serializer.serialize({"key": "value"})
except TypeError as e:
print(e)
# "ArrowSerializer only supports DataFrames. Use StandardSerializer for dict types."
# "... Got a dict that is not convertible to an Arrow table: value for 'key' is str, not a list or array. ..."
```

## Performance Benchmarks
Expand Down
9 changes: 9 additions & 0 deletions src/cachekit/serializers/arrow_serializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@

from __future__ import annotations

from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, BinaryIO, ClassVar

from .base import SerializationError, SerializationFormat, SerializationMetadata
Expand Down Expand Up @@ -263,6 +264,14 @@ def _to_table(obj: Any) -> pa.Table: # type: ignore[name-defined]
# (e.g. dict-of-scalars -> "'int' object is not iterable") into the
# documented TypeError so callers get a consistent, actionable message.
try:
for name, column in obj.items():
# pa.table iterates these instead of rejecting them, storing a wrong column:
# a str as its characters, bytes as byte values, a dict as its keys only.
# A value with an Arrow array protocol is converted through it, not iterated.
if isinstance(column, (str, bytes, bytearray, Mapping)) and not (
hasattr(column, "__arrow_array__") or hasattr(column, "__arrow_c_array__")
):
raise TypeError(f"value for {name!r} is {type(column).__name__}, not a list or array")
return pa.table(obj)
except (pa.ArrowInvalid, pa.ArrowTypeError, TypeError, ValueError) as e:
raise TypeError(
Expand Down
127 changes: 119 additions & 8 deletions tests/unit/test_arrow_serializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,17 +5,57 @@

from __future__ import annotations

import io
from collections.abc import Mapping

import pytest

# ArrowSerializer requires the [data] extra — absent e.g. in the free-threaded
# CI lane until pandas/pyarrow ship free-threaded wheels (LAB-511).
pd = pytest.importorskip("pandas")
pa = pytest.importorskip("pyarrow")
np = pytest.importorskip("numpy")

from cachekit.serializers.arrow_serializer import ArrowSerializer # noqa: E402
from cachekit.serializers.base import SerializationError, SerializationFormat, SerializationMetadata # noqa: E402


def _serialize_via(path: str, serializer: ArrowSerializer, obj: object) -> tuple[bytes, SerializationMetadata]:
"""Serialize through the buffered (``serialize``) or streaming (``serialize_to_sink``) path."""
if path == "buffered":
return serializer.serialize(obj)
sink = io.BytesIO()
metadata = serializer.serialize_to_sink(obj, sink)
return sink.getvalue(), metadata


class _KeysMapping(Mapping):
"""A Mapping whose iteration yields its keys, so pa.table alone would store ["k1", "k2"]."""

def __getitem__(self, key):
raise KeyError(key)

def __iter__(self):
return iter(["k1", "k2"])

def __len__(self):
return 2


class _ArrowArrayMapping(_KeysMapping):
"""Also an Arrow column adapter: pyarrow converts it through ``__arrow_array__`` to [10, 20]."""

def __arrow_array__(self, type=None):
return pa.array([10, 20])


class _ArrowCArrayMapping(_KeysMapping):
"""Same adapter through the Arrow PyCapsule protocol (``__arrow_c_array__``)."""

def __arrow_c_array__(self, requested_schema=None):
return pa.array([10, 20]).__arrow_c_array__(requested_schema)


class TestArrowSerializerBasics:
"""Test basic ArrowSerializer functionality."""

Expand Down Expand Up @@ -205,6 +245,21 @@ def test_dict_of_arrays_with_arrow_return_format(self):
assert isinstance(result, pa.Table)
assert result.column_names == ["a", "b"]

@pytest.mark.parametrize("path", ["buffered", "streaming"])
@pytest.mark.parametrize(
"make_column",
[list, np.array, pd.Series, pa.array],
ids=["list", "numpy", "series", "pyarrow"],
)
def test_dict_of_columns_round_trips_on_both_paths(self, path, make_column):
"""Every accepted column container round-trips, buffered and streamed."""
serializer = ArrowSerializer()
data_dict = {"n": make_column([1, 2, 3]), "s": make_column(["x", "y", "z"])}

data, metadata = _serialize_via(path, serializer, data_dict)

assert serializer.deserialize(data, metadata).to_dict("list") == {"n": [1, 2, 3], "s": ["x", "y", "z"]}


class TestErrorHandling:
"""Test error handling for unsupported types and corrupted data."""
Expand All @@ -221,16 +276,72 @@ def test_scalar_value_raises_type_error(self):
assert "Got: int" in error_msg
assert "For scalar values or nested dicts, use AutoSerializer" in error_msg

def test_non_columnar_dict_successfully_serialized(self):
"""Arrow can handle certain dict structures (converts to struct/list types)."""
@pytest.mark.parametrize("path", ["buffered", "streaming"])
@pytest.mark.parametrize(
("obj", "detail"),
[
({"name": "Alice"}, "value for 'name' is str"),
({"b": b"ab"}, "value for 'b' is bytes"),
({"b": bytearray(b"ab")}, "value for 'b' is bytearray"),
({"user": {"name": "x"}}, "value for 'user' is dict"),
({"ok": [1, 2], "name": "Alice"}, "value for 'name' is str"),
],
ids=["str", "bytes", "bytearray", "dict", "one-bad-column"],
)
def test_non_column_dict_value_raises_type_error(self, path, obj, detail):
"""pyarrow would iterate these values into a wrong column (a str into its characters,
bytes into their byte values, a dict into its keys), so they are rejected instead."""
serializer = ArrowSerializer()
nested = {"key": {"nested": "value"}}

# Arrow will convert this successfully (struct/list types)
# This is actually valid - Arrow has flexible schema support
data, metadata = serializer.serialize(nested)
assert isinstance(data, bytes)
assert isinstance(metadata, SerializationMetadata)
with pytest.raises(TypeError) as exc_info:
_serialize_via(path, serializer, obj)

error_msg = str(exc_info.value)
assert "Got a dict that is not convertible to an Arrow table" in error_msg
assert f"{detail}, not a list or array" in error_msg

@pytest.mark.parametrize("path", ["buffered", "streaming"])
@pytest.mark.parametrize("adapter", [_ArrowArrayMapping, _ArrowCArrayMapping], ids=["arrow_array", "arrow_c_array"])
def test_mapping_with_arrow_array_protocol_round_trips(self, path, adapter):
"""pyarrow converts a value through its Arrow array protocol rather than iterating it,
so a Mapping that implements one is a correct column and is not rejected."""
serializer = ArrowSerializer()

data, metadata = _serialize_via(path, serializer, {"c": adapter()})

assert serializer.deserialize(data, metadata).to_dict("list") == {"c": [10, 20]}

def test_rejected_dict_writes_nothing_to_sink(self):
"""The streaming path rejects the value before writing a byte."""
sink = io.BytesIO()

with pytest.raises(TypeError):
ArrowSerializer().serialize_to_sink({"name": "Alice"}, sink)

assert sink.getvalue() == b""

def test_decorator_reruns_instead_of_returning_a_wrong_hit(self, tmp_path):
"""Through @cache with a File backend, a rejected dict is never cached, so the
second call runs the function and returns the original dict, not a DataFrame."""
from cachekit import cache
from cachekit.backends.file import FileBackend
from cachekit.backends.file.config import FileBackendConfig

backend = FileBackend(FileBackendConfig(cache_dir=str(tmp_path), max_size_mb=10, max_value_mb=5))
calls = 0

@cache(serializer="arrow", backend=backend, ttl=60)
def load_user() -> dict:
nonlocal calls
calls += 1
return {"name": "Alice"}

assert load_user() == {"name": "Alice"}
second = load_user()

assert calls == 2
assert isinstance(second, dict)
assert second == {"name": "Alice"}

def test_string_raises_type_error(self):
"""String value raises TypeError."""
Expand Down
Loading