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
6 changes: 4 additions & 2 deletions haystack/components/joiners/branch.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,8 +113,10 @@ def from_dict(cls, data: dict[str, Any]) -> "BranchJoiner":
:returns:
A deserialized `BranchJoiner` instance.
"""
data["init_parameters"]["type_"] = deserialize_type(data["init_parameters"]["type_"])
return default_from_dict(cls, data)
# Copy so the caller's data keeps the serialized type and can be deserialized again
init_parameters = dict(data["init_parameters"])
init_parameters["type_"] = deserialize_type(init_parameters["type_"])
return default_from_dict(cls, {**data, "init_parameters": init_parameters})

def run(self, **kwargs: Any) -> dict[str, Any]:
"""
Expand Down
33 changes: 18 additions & 15 deletions haystack/components/routers/conditional_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -364,7 +364,8 @@ def from_dict(cls, data: dict[str, Any]) -> "ConditionalRouter":
:returns:
The deserialized component.
"""
init_params = data.get("init_parameters", {})
# Copy so the caller's data stays serialized; nested routes and filters are copied below too
init_params = dict(data.get("init_parameters", {}))

# `unsafe=True` swaps the Jinja sandbox for a NativeEnvironment that executes arbitrary code.
# Honor it from serialized data only when the whole pipeline is being loaded in unsafe mode;
Expand All @@ -375,7 +376,7 @@ def from_dict(cls, data: dict[str, Any]) -> "ConditionalRouter":
"If you trust the source of this data, load it with Pipeline.load(..., unsafe=True)."
)

custom_filters = init_params.get("custom_filters", {})
custom_filters = init_params.get("custom_filters")
if custom_filters and not _is_unsafe_deserialization():
raise DeserializationError(
"Refusing to deserialize a ConditionalRouter with custom filters while loading in safe mode. "
Expand All @@ -384,19 +385,21 @@ def from_dict(cls, data: dict[str, Any]) -> "ConditionalRouter":
)

routes = init_params.get("routes")
for route in routes:
# output_type needs to be deserialized from a string to a type
if isinstance(route["output_type"], list):
route["output_type"] = [deserialize_type(t) for t in route["output_type"]]
else:
route["output_type"] = deserialize_type(route["output_type"])

# Since the custom_filters are typed as optional in the init signature, we catch the
# case where they are not present in the serialized data and set them to an empty dict.
if custom_filters is not None:
for name, filter_func in custom_filters.items():
init_params["custom_filters"][name] = deserialize_callable(filter_func) if filter_func else None
return default_from_dict(cls, data)
if routes is not None:
init_params["routes"] = routes = [dict(route) for route in routes]
for route in routes:
# output_type needs to be deserialized from a string to a type
if isinstance(route["output_type"], list):
route["output_type"] = [deserialize_type(t) for t in route["output_type"]]
else:
route["output_type"] = deserialize_type(route["output_type"])

if custom_filters:
init_params["custom_filters"] = {
name: deserialize_callable(filter_func) if filter_func else None
for name, filter_func in custom_filters.items()
}
return default_from_dict(cls, {**data, "init_parameters": init_params})

def run(self, **kwargs: Any) -> dict[str, Any]:
"""
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
---
fixes:
- |
Fixed ``ConditionalRouter.from_dict`` mutating the caller's ``routes`` data in place: serialized
``output_type`` strings were deserialized directly inside the caller's dictionaries, so reusing the same
serialized pipeline dict afterwards yielded already-deserialized type objects. ``from_dict`` now works on a
copy and the caller's data is left untouched. ``BranchJoiner.from_dict`` received the same treatment.
14 changes: 14 additions & 0 deletions test/components/joiners/test_branch_joiner.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,20 @@
from haystack.components.joiners import BranchJoiner


class TestBranchJoinerDeserialization:
def test_from_dict_does_not_mutate_caller_data(self):
joiner = BranchJoiner(list[str])
data = joiner.to_dict()
serialized_type = data["init_parameters"]["type_"]
assert isinstance(serialized_type, str)

BranchJoiner.from_dict(data)

assert data["init_parameters"]["type_"] == serialized_type
# a second deserialization of the same dict must behave like the first
BranchJoiner.from_dict(data)


class TestBranchJoiner:
def test_one_value(self):
joiner = BranchJoiner(int)
Expand Down
18 changes: 18 additions & 0 deletions test/components/routers/test_conditional_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -1021,3 +1021,21 @@ def test_conditional_router_passthrough_skips_output_template_validation(self):
router = ConditionalRouter(routes)
result = router.run(**{"{{unclosed": "value"})
assert result == {"out": "value"}


class TestConditionalRouterDeserialization:
def test_from_dict_does_not_mutate_caller_data(self):
routes: list[Route] = [
{"condition": "{{ x > 1 }}", "output": "{{ x }}", "output_name": "big", "output_type": int},
{"condition": "{{ x <= 1 }}", "output": "{{ x }}", "output_name": "small", "output_type": int},
]
router = ConditionalRouter(routes)
data = router.to_dict()
serialized_types = [route["output_type"] for route in data["init_parameters"]["routes"]]
assert all(isinstance(t, str) for t in serialized_types)

ConditionalRouter.from_dict(data)

assert [route["output_type"] for route in data["init_parameters"]["routes"]] == serialized_types
# a second deserialization of the same dict must behave like the first
ConditionalRouter.from_dict(data)
Loading