Skip to content
Closed
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
13 changes: 11 additions & 2 deletions haystack/core/pipeline/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
Original file line number Diff line number Diff line change
@@ -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``.
19 changes: 18 additions & 1 deletion test/core/pipeline/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@

import logging
import warnings
from collections import namedtuple
from collections import Counter, OrderedDict, defaultdict, namedtuple

import pytest

Expand Down Expand Up @@ -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
Expand Down