Skip to content
Open
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
18 changes: 17 additions & 1 deletion burr/core/serde.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,13 @@
from typing import Any, Union

KEY = "__burr_serde__"
# Marker used to wrap an ordinary dictionary that itself contains KEY. Without
# it, such a dictionary is indistinguishable from a serde envelope and
# `deserialize` fails -- or dispatches to an unrelated deserializer -- when the
# value is read back.
ESCAPED_DICT = "burr.dict"
# Key an escaped dictionary stores its contents under.
PAYLOAD_KEY = "value"


class StringDispatch:
Expand Down Expand Up @@ -116,7 +123,16 @@ def serialize_primitive(value, **kwargs) -> Union[str, int, float, bool]:

@serialize.register(dict)
def serialize_dict(value: dict, **kwargs) -> dict[str, Any]:
return {k: serialize(v, **kwargs) for k, v in value.items()}
serialized = {k: serialize(v, **kwargs) for k, v in value.items()}
if KEY in value:
return {KEY: ESCAPED_DICT, PAYLOAD_KEY: serialized}
return serialized


@deserializer.register(ESCAPED_DICT)
def deserialize_escaped_dict(value: dict, **kwargs) -> dict[str, Any]:
"""Deserializes an ordinary dictionary that carries the serde marker key."""
return {k: deserialize(v, **kwargs) for k, v in value[PAYLOAD_KEY].items()}


@serialize.register(list)
Expand Down
32 changes: 31 additions & 1 deletion tests/core/test_serde.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,8 @@

import pytest

from burr.core.serde import StringDispatch, deserialize, serialize
from burr.core import State
from burr.core.serde import KEY, StringDispatch, deserialize, serialize


def test_serialize_primitive_types():
Expand Down Expand Up @@ -74,3 +75,32 @@ def test_string_dispatch_no_key_informative_message():
assert "nonexistent_key" in str(exc_info.value)
assert "known_key" in str(exc_info.value)
assert "imported" in str(exc_info.value)


def test_dict_with_serde_key_round_trips():
"""An ordinary dict carrying the serde marker must survive a round trip.

It used to be read back as a serde envelope, which raised "No deserializer
registered for key" instead of returning the value.
"""
value = {KEY: "hello", "nested": {KEY: {"deep": 1}}, "list": [{KEY: 1}]}

assert deserialize(serialize(value)) == value


def test_state_with_serde_key_round_trips():
"""State containing such a dict deserializes instead of failing."""
state = State({"payload": {KEY: "hello", "count": 2}})

restored = State.deserialize(state.serialize())

assert restored["payload"] == {KEY: "hello", "count": 2}


def test_envelope_without_imported_deserializer_still_raises():
"""A real envelope whose module was not imported keeps its helpful error."""
with pytest.raises(ValueError) as exc_info:
deserialize({KEY: "some.serde.that.is.not.imported"})

assert "some.serde.that.is.not.imported" in str(exc_info.value)
assert "imported" in str(exc_info.value)
Loading