diff --git a/haystack/core/pipeline/utils.py b/haystack/core/pipeline/utils.py index d1263dec97..9fc75d3926 100644 --- a/haystack/core/pipeline/utils.py +++ b/haystack/core/pipeline/utils.py @@ -4,7 +4,7 @@ import heapq from collections.abc import Callable -from copy import deepcopy +from copy import copy, deepcopy from functools import wraps from itertools import count from typing import Any @@ -42,7 +42,16 @@ def _deepcopy_with_exceptions(obj: Any) -> Any: return type(obj)(_deepcopy_with_exceptions(v) for v in obj) if isinstance(obj, dict): - return {k: _deepcopy_with_exceptions(v) for k, v in obj.items()} + if type(obj) is dict: + return {k: _deepcopy_with_exceptions(v) for k, v in obj.items()} + try: + copied = copy(obj) + copied.clear() + for key, value in obj.items(): + copied[key] = _deepcopy_with_exceptions(value) + except Exception: + copied = {k: _deepcopy_with_exceptions(v) for k, v in obj.items()} + return copied # Components and Tools often contain objects that we do not want to deepcopy or are not deepcopyable # (e.g. models, clients, etc.). In this case we return the object as-is. diff --git a/releasenotes/notes/preserve-dict-subclasses-67205ab6180303fa.yaml b/releasenotes/notes/preserve-dict-subclasses-67205ab6180303fa.yaml new file mode 100644 index 0000000000..0e0f764095 --- /dev/null +++ b/releasenotes/notes/preserve-dict-subclasses-67205ab6180303fa.yaml @@ -0,0 +1,6 @@ +--- +fixes: + - | + Preserve dictionary subclasses such as ``defaultdict``, ``OrderedDict`` and + ``Counter`` when copying pipeline inputs. Their subclass behavior and state + are now retained instead of being replaced by a plain ``dict``. diff --git a/test/core/pipeline/test_utils.py b/test/core/pipeline/test_utils.py index c9e1d9831f..f544fd3ad2 100644 --- a/test/core/pipeline/test_utils.py +++ b/test/core/pipeline/test_utils.py @@ -4,7 +4,7 @@ import logging import warnings -from collections import namedtuple +from collections import Counter, OrderedDict, defaultdict, namedtuple import pytest @@ -269,6 +269,23 @@ def test_deepcopy_with_fallback_namedtuple(self): # Its contents are deep-copied, matching how plain tuples are handled. assert copy["point"].x is not original["point"].x + def test_deepcopy_preserves_dict_subclasses(self): + originals = [ + OrderedDict([("value", {"nested": []})]), + defaultdict[str, object](list, {"value": {"nested": []}}), + Counter({"value": 2}), + ] + + copies = [_deepcopy_with_exceptions(original) for original in originals] + + for original, copy in zip(originals, copies, strict=True): + assert type(copy) is type(original) + assert copy == original + assert copy is not original + + assert copies[1]["missing"] == [] + assert "missing" in copies[1] + class TestArgsDeprecated: @pytest.fixture