diff --git a/haystack/components/joiners/branch.py b/haystack/components/joiners/branch.py index 40bf9ce346d..3715991459e 100644 --- a/haystack/components/joiners/branch.py +++ b/haystack/components/joiners/branch.py @@ -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]: """ diff --git a/haystack/components/routers/conditional_router.py b/haystack/components/routers/conditional_router.py index 5c5b248d51a..06d838f3b4c 100644 --- a/haystack/components/routers/conditional_router.py +++ b/haystack/components/routers/conditional_router.py @@ -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; @@ -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. " @@ -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]: """ diff --git a/releasenotes/notes/fix-conditional-router-from-dict-caller-mutation-da863fb2fd788ea9.yaml b/releasenotes/notes/fix-conditional-router-from-dict-caller-mutation-da863fb2fd788ea9.yaml new file mode 100644 index 00000000000..c672a6ceba7 --- /dev/null +++ b/releasenotes/notes/fix-conditional-router-from-dict-caller-mutation-da863fb2fd788ea9.yaml @@ -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. diff --git a/test/components/joiners/test_branch_joiner.py b/test/components/joiners/test_branch_joiner.py index 8cb2c6b818f..05e770f3812 100644 --- a/test/components/joiners/test_branch_joiner.py +++ b/test/components/joiners/test_branch_joiner.py @@ -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) diff --git a/test/components/routers/test_conditional_router.py b/test/components/routers/test_conditional_router.py index 12e9bf04f69..4d73994dd7f 100644 --- a/test/components/routers/test_conditional_router.py +++ b/test/components/routers/test_conditional_router.py @@ -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)