Skip to content

Commit 12ec79b

Browse files
committed
fix: preserve explicit runtime retriever overrides
1 parent 038931a commit 12ec79b

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,5 +103,7 @@ 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}

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
@@ -121,6 +121,13 @@ def test_retriever_init_filter_run_filter_override(self, sample_document_store,
121121
assert len(result["documents"]) == 2
122122
assert TestFilterRetriever._documents_equal(result["documents"], sample_docs["de_docs"])
123123

124+
def test_retriever_runtime_empty_filter_clears_init_filter(self, sample_document_store, sample_docs):
125+
retriever = FilterRetriever(sample_document_store, filters={"field": "lang", "operator": "==", "value": "en"})
126+
result = retriever.run(filters={})
127+
128+
assert len(result["documents"]) == len(sample_docs["all_docs"])
129+
assert TestFilterRetriever._documents_equal(result["documents"], sample_docs["all_docs"])
130+
124131
@pytest.mark.integration
125132
def test_run_with_pipeline(self, sample_document_store, sample_docs):
126133
retriever = FilterRetriever(sample_document_store, filters={"field": "lang", "operator": "==", "value": "de"})
@@ -144,3 +151,9 @@ def test_run_with_pipeline(self, sample_document_store, sample_docs):
144151
results_docs = result["retriever"]["documents"]
145152
assert results_docs
146153
assert TestFilterRetriever._documents_equal(results_docs, sample_docs["en_docs"])
154+
155+
result: dict[str, Any] = pipeline.run(data={"retriever": {"filters": {}}})
156+
157+
results_docs = result["retriever"]["documents"]
158+
assert len(results_docs) == len(sample_docs["all_docs"])
159+
assert TestFilterRetriever._documents_equal(results_docs, sample_docs["all_docs"])

test/components/retrievers/test_filter_retriever_async.py

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

73+
@pytest.mark.asyncio
74+
async def test_retriever_runtime_empty_filter_clears_init_filter(self, sample_document_store, sample_docs):
75+
retriever = FilterRetriever(sample_document_store, filters={"field": "lang", "operator": "==", "value": "en"})
76+
result = await retriever.run_async(filters={})
77+
78+
assert len(result["documents"]) == len(sample_docs["all_docs"])
79+
assert TestFilterRetrieverAsync._documents_equal(result["documents"], sample_docs["all_docs"])
80+
7381
@pytest.mark.asyncio
7482
@pytest.mark.integration
7583
async def test_run_with_pipeline(self, sample_document_store, sample_docs):
@@ -94,3 +102,9 @@ async def test_run_with_pipeline(self, sample_document_store, sample_docs):
94102
results_docs = result["retriever"]["documents"]
95103
assert results_docs
96104
assert TestFilterRetrieverAsync._documents_equal(results_docs, sample_docs["en_docs"])
105+
106+
result: dict[str, Any] = await pipeline.run_async(data={"retriever": {"filters": {}}})
107+
108+
results_docs = result["retriever"]["documents"]
109+
assert len(results_docs) == len(sample_docs["all_docs"])
110+
assert TestFilterRetrieverAsync._documents_equal(results_docs, sample_docs["all_docs"])

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
@@ -14,6 +14,7 @@
1414

1515

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

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

0 commit comments

Comments
 (0)