Skip to content

Commit f16dd9a

Browse files
committed
fix: preserve explicit runtime retriever overrides
1 parent 6f75d43 commit f16dd9a

8 files changed

Lines changed: 73 additions & 12 deletions

haystack/components/retrievers/filter_retriever.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,8 @@ class FilterRetriever:
2929
doc_store.write_documents(docs)
3030
retriever = FilterRetriever(doc_store, filters={"field": "lang", "operator": "==", "value": "en"})
3131
32-
# if passed in the run method, filters override those provided at initialization
32+
# If passed in the run method, filters override those provided at initialization.
33+
# Passing an empty dictionary explicitly clears the initialization filters.
3334
result = retriever.run(filters={"field": "lang", "operator": "==", "value": "de"})
3435
3536
print(result["documents"])
@@ -44,6 +45,7 @@ def __init__(self, document_store: DocumentStore, filters: dict[str, Any] | None
4445
An instance of a Document Store to use with the Retriever.
4546
:param filters:
4647
A dictionary with filters to narrow down the search space.
48+
Passing an empty dictionary explicitly clears filters provided at initialization.
4749
"""
4850
self.document_store = document_store
4951
self.filters = filters
@@ -82,11 +84,12 @@ def run(self, filters: dict[str, Any] | None = None) -> dict[str, Any]:
8284
8385
:param filters:
8486
A dictionary with filters to narrow down the search space.
87+
Passing an empty dictionary explicitly clears filters provided at initialization.
8588
If not specified, the FilterRetriever uses the values provided at initialization.
8689
:returns:
8790
A list of retrieved documents.
8891
"""
89-
return {"documents": self.document_store.filter_documents(filters=filters or self.filters)}
92+
return {"documents": self.document_store.filter_documents(filters=self.filters if filters is None else filters)}
9093

9194
@component.output_types(documents=list[Document])
9295
async def run_async(self, filters: dict[str, Any] | None = None) -> dict[str, Any]:
@@ -100,7 +103,9 @@ async def run_async(self, filters: dict[str, Any] | None = None) -> dict[str, An
100103
A list of retrieved documents.
101104
"""
102105
# 'ignore' since filter_documents_async is not defined in the Protocol but exists in the implementations
103-
out_documents = await self.document_store.filter_documents_async(filters=filters or self.filters) # type: ignore[attr-defined]
106+
out_documents = await self.document_store.filter_documents_async( # type: ignore[attr-defined]
107+
filters=self.filters if filters is None else filters
108+
)
104109
return {"documents": out_documents}
105110

106111
def close(self) -> None:

haystack/components/retrievers/sentence_window_retriever.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -104,8 +104,7 @@ def __init__(
104104
metadata fields. If False, it will skip retrieving the context for documents that are missing
105105
the required metadata fields, but will still include the original document in the results.
106106
"""
107-
if window_size < 1:
108-
raise ValueError("The window_size parameter must be greater than 0.")
107+
self._validate_window_size(window_size)
109108

110109
self.window_size = window_size
111110
self.document_store = document_store
@@ -197,8 +196,8 @@ def run(self, retrieved_documents: list[Document], window_size: int | None = Non
197196
meta field.
198197
199198
"""
200-
window_size = window_size or self.window_size
201-
SentenceWindowRetriever._raise_if_windows_size_is_negative(window_size)
199+
window_size = self.window_size if window_size is None else window_size
200+
self._validate_window_size(window_size)
202201
self._raise_if_documents_do_not_have_expected_metadata(retrieved_documents)
203202

204203
context_text = []
@@ -230,8 +229,8 @@ async def run_async(self, retrieved_documents: list[Document], window_size: int
230229
meta field.
231230
232231
"""
233-
window_size = window_size or self.window_size
234-
SentenceWindowRetriever._raise_if_windows_size_is_negative(window_size)
232+
window_size = self.window_size if window_size is None else window_size
233+
self._validate_window_size(window_size)
235234
self._raise_if_documents_do_not_have_expected_metadata(retrieved_documents)
236235

237236
context_text = []
@@ -244,7 +243,7 @@ async def run_async(self, retrieved_documents: list[Document], window_size: int
244243
return {"context_windows": context_text, "context_documents": context_documents}
245244

246245
@staticmethod
247-
def _raise_if_windows_size_is_negative(window_size: int) -> None:
246+
def _validate_window_size(window_size: int) -> None:
248247
if window_size < 1:
249248
raise ValueError("The window_size parameter must be greater than 0.")
250249

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
---
2+
fixes:
3+
- |
4+
FilterRetriever now treats an explicitly provided empty runtime filter as an override
5+
of initialization filters in both synchronous and asynchronous execution.
Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
---
2+
fixes:
3+
- |
4+
SentenceWindowRetriever now validates an explicitly provided runtime `window_size=0`
5+
instead of silently falling back to the constructor value.

test/components/retrievers/test_filter_retriever.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -122,6 +122,13 @@ def test_retriever_init_filter_run_filter_override(self, sample_document_store,
122122
assert len(result["documents"]) == 2
123123
assert TestFilterRetriever._documents_equal(result["documents"], sample_docs["de_docs"])
124124

125+
def test_retriever_runtime_empty_filter_clears_init_filter(self, sample_document_store, sample_docs):
126+
retriever = FilterRetriever(sample_document_store, filters={"field": "lang", "operator": "==", "value": "en"})
127+
result = retriever.run(filters={})
128+
129+
assert len(result["documents"]) == len(sample_docs["all_docs"])
130+
assert TestFilterRetriever._documents_equal(result["documents"], sample_docs["all_docs"])
131+
125132
@pytest.mark.integration
126133
def test_run_with_pipeline(self, sample_document_store, sample_docs):
127134
retriever = FilterRetriever(sample_document_store, filters={"field": "lang", "operator": "==", "value": "de"})
@@ -146,6 +153,12 @@ def test_run_with_pipeline(self, sample_document_store, sample_docs):
146153
assert results_docs
147154
assert TestFilterRetriever._documents_equal(results_docs, sample_docs["en_docs"])
148155

156+
result: dict[str, Any] = pipeline.run(data={"retriever": {"filters": {}}})
157+
158+
results_docs = result["retriever"]["documents"]
159+
assert len(results_docs) == len(sample_docs["all_docs"])
160+
assert TestFilterRetriever._documents_equal(results_docs, sample_docs["all_docs"])
161+
149162
def test_close(self):
150163
closable_document_store = Mock(spec=["close"])
151164
retriever = FilterRetriever(document_store=closable_document_store)

test/components/retrievers/test_filter_retriever_async.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -71,6 +71,14 @@ async def test_retriever_init_filter_run_filter_override(self, sample_document_s
7171
assert len(result["documents"]) == 2
7272
assert TestFilterRetrieverAsync._documents_equal(result["documents"], sample_docs["de_docs"])
7373

74+
@pytest.mark.asyncio
75+
async def test_retriever_runtime_empty_filter_clears_init_filter(self, sample_document_store, sample_docs):
76+
retriever = FilterRetriever(sample_document_store, filters={"field": "lang", "operator": "==", "value": "en"})
77+
result = await retriever.run_async(filters={})
78+
79+
assert len(result["documents"]) == len(sample_docs["all_docs"])
80+
assert TestFilterRetrieverAsync._documents_equal(result["documents"], sample_docs["all_docs"])
81+
7482
@pytest.mark.asyncio
7583
@pytest.mark.integration
7684
async def test_run_with_pipeline(self, sample_document_store, sample_docs):
@@ -96,6 +104,12 @@ async def test_run_with_pipeline(self, sample_document_store, sample_docs):
96104
assert results_docs
97105
assert TestFilterRetrieverAsync._documents_equal(results_docs, sample_docs["en_docs"])
98106

107+
result: dict[str, Any] = await pipeline.run_async(data={"retriever": {"filters": {}}})
108+
109+
results_docs = result["retriever"]["documents"]
110+
assert len(results_docs) == len(sample_docs["all_docs"])
111+
assert TestFilterRetrieverAsync._documents_equal(results_docs, sample_docs["all_docs"])
112+
99113
@pytest.mark.asyncio
100114
async def test_close_async(self):
101115
closable_document_store = Mock(spec=["close_async"])

test/components/retrievers/test_sentence_window_retriever.py

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@ def test_init_with_parameters(self, in_memory_doc_store):
2424
retriever = SentenceWindowRetriever(in_memory_doc_store, window_size=5)
2525
assert retriever.window_size == 5
2626

27-
def test_init_with_invalid_window_size_parameter(self, in_memory_doc_store):
27+
def test_init_invalid_window_size(self, in_memory_doc_store):
2828
with pytest.raises(ValueError):
2929
SentenceWindowRetriever(in_memory_doc_store, window_size=-2)
3030

@@ -174,12 +174,21 @@ def test_document_without_all_source_ids(self, in_memory_doc_store):
174174
)
175175
retriever.run(retrieved_documents=docs)
176176

177-
def test_run_invalid_window_size(self, in_memory_doc_store):
177+
def test_init_rejects_zero_window_size(self, in_memory_doc_store):
178178
docs = [Document(content="This is a text with some words. There is a ", meta={"id": "doc_0", "split_id": 0})]
179179
with pytest.raises(ValueError):
180180
retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, window_size=0)
181181
retriever.run(retrieved_documents=docs)
182182

183+
def test_run_rejects_zero_runtime_window_size(self, in_memory_doc_store):
184+
retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, window_size=3)
185+
186+
with pytest.raises(ValueError, match="window_size parameter must be greater than 0"):
187+
retriever.run(retrieved_documents=[], window_size=0)
188+
189+
with pytest.raises(ValueError, match="window_size parameter must be greater than 0"):
190+
retriever.run(retrieved_documents=[], window_size=-1)
191+
183192
def test_constructor_parameter_does_not_change(self, in_memory_doc_store):
184193
retriever = SentenceWindowRetriever(in_memory_doc_store, window_size=5)
185194
assert retriever.window_size == 5

test/components/retrievers/test_sentence_window_retriever_async.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515

1616

1717
class TestSentenceWindowRetrieverAsync:
18+
@pytest.mark.asyncio
1819
async def test_document_without_split_id(self, in_memory_doc_store):
1920
docs = [
2021
Document(content="This is a text with some words. There is a ", meta={"id": "doc_0"}),
@@ -68,6 +69,16 @@ async def test_run_async_invalid_window_size(self, in_memory_doc_store):
6869
retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, window_size=0)
6970
await retriever.run_async(retrieved_documents=docs)
7071

72+
@pytest.mark.asyncio
73+
async def test_run_async_rejects_zero_runtime_window_size(self, in_memory_doc_store):
74+
retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, window_size=3)
75+
76+
with pytest.raises(ValueError, match="window_size parameter must be greater than 0"):
77+
await retriever.run_async(retrieved_documents=[], window_size=0)
78+
79+
with pytest.raises(ValueError, match="window_size parameter must be greater than 0"):
80+
await retriever.run_async(retrieved_documents=[], window_size=-1)
81+
7182
@pytest.mark.asyncio
7283
async def test_constructor_parameter_does_not_change(self, in_memory_doc_store):
7384
retriever = SentenceWindowRetriever(in_memory_doc_store, window_size=5)

0 commit comments

Comments
 (0)