diff --git a/haystack/components/retrievers/in_memory/bm25_retriever.py b/haystack/components/retrievers/in_memory/bm25_retriever.py index 05618f7d810..450f70cc561 100644 --- a/haystack/components/retrievers/in_memory/bm25_retriever.py +++ b/haystack/components/retrievers/in_memory/bm25_retriever.py @@ -6,7 +6,7 @@ from haystack import Document, component, default_from_dict, default_to_dict from haystack.document_stores.in_memory import InMemoryDocumentStore -from haystack.document_stores.types import FilterPolicy +from haystack.document_stores.types import FilterPolicy, apply_filter_policy @component @@ -143,10 +143,7 @@ def run( :raises ValueError: If the specified DocumentStore is not found or is not a InMemoryDocumentStore instance. """ - if self.filter_policy == FilterPolicy.MERGE and filters: - filters = {**(self.filters or {}), **filters} - else: - filters = filters or self.filters + filters = apply_filter_policy(self.filter_policy, self.filters, filters) if top_k is None: top_k = self.top_k if scale_score is None: @@ -181,10 +178,7 @@ async def run_async( :raises ValueError: If the specified DocumentStore is not found or is not a InMemoryDocumentStore instance. """ - if self.filter_policy == FilterPolicy.MERGE and filters: - filters = {**(self.filters or {}), **filters} - else: - filters = filters or self.filters + filters = apply_filter_policy(self.filter_policy, self.filters, filters) if top_k is None: top_k = self.top_k if scale_score is None: diff --git a/haystack/components/retrievers/in_memory/embedding_retriever.py b/haystack/components/retrievers/in_memory/embedding_retriever.py index f185d33dc2f..5c2bf0dcf8f 100644 --- a/haystack/components/retrievers/in_memory/embedding_retriever.py +++ b/haystack/components/retrievers/in_memory/embedding_retriever.py @@ -6,7 +6,7 @@ from haystack import Document, component, default_from_dict, default_to_dict from haystack.document_stores.in_memory import InMemoryDocumentStore -from haystack.document_stores.types import FilterPolicy +from haystack.document_stores.types import FilterPolicy, apply_filter_policy @component @@ -163,10 +163,7 @@ def run( :raises ValueError: If the specified DocumentStore is not found or is not an InMemoryDocumentStore instance. """ - if self.filter_policy == FilterPolicy.MERGE and filters: - filters = {**(self.filters or {}), **filters} - else: - filters = filters or self.filters + filters = apply_filter_policy(self.filter_policy, self.filters, filters) if top_k is None: top_k = self.top_k if scale_score is None: @@ -214,10 +211,7 @@ async def run_async( :raises ValueError: If the specified DocumentStore is not found or is not an InMemoryDocumentStore instance. """ - if self.filter_policy == FilterPolicy.MERGE and filters: - filters = {**(self.filters or {}), **filters} - else: - filters = filters or self.filters + filters = apply_filter_policy(self.filter_policy, self.filters, filters) if top_k is None: top_k = self.top_k if scale_score is None: diff --git a/releasenotes/notes/filter-policy-merge-in-memory-retrievers-f49dfcb2a1332e8b.yaml b/releasenotes/notes/filter-policy-merge-in-memory-retrievers-f49dfcb2a1332e8b.yaml new file mode 100644 index 00000000000..c4eb313961c --- /dev/null +++ b/releasenotes/notes/filter-policy-merge-in-memory-retrievers-f49dfcb2a1332e8b.yaml @@ -0,0 +1,6 @@ +--- +fixes: + - | + Fix ``FilterPolicy.MERGE`` in ``InMemoryBM25Retriever`` and + ``InMemoryEmbeddingRetriever`` so initialization filters are combined with + runtime comparison filters instead of being silently overwritten. \ No newline at end of file diff --git a/test/components/retrievers/test_in_memory_bm25_retriever.py b/test/components/retrievers/test_in_memory_bm25_retriever.py index 42cbcb56ce8..3452667d95d 100644 --- a/test/components/retrievers/test_in_memory_bm25_retriever.py +++ b/test/components/retrievers/test_in_memory_bm25_retriever.py @@ -139,6 +139,47 @@ def test_retriever_valid_run(self, in_memory_doc_store, mock_docs): assert len(result["documents"]) == 5 assert result["documents"][0].content == "PHP is a popular programming language" + def test_run_with_filter_policy_merge_combines_init_and_runtime_filters(self, in_memory_doc_store): + in_memory_doc_store.write_documents( + [ + Document(content="python article current", meta={"type": "article", "year": 2020}), + Document(content="python blog current", meta={"type": "blog", "year": 2021}), + Document(content="python article archived", meta={"type": "article", "year": 2019}), + ] + ) + + retriever = InMemoryBM25Retriever( + in_memory_doc_store, + filters={"field": "meta.type", "operator": "==", "value": "article"}, + filter_policy=FilterPolicy.MERGE, + ) + + result = retriever.run(query="python", filters={"field": "meta.year", "operator": ">=", "value": 2020}) + + assert [doc.content for doc in result["documents"]] == ["python article current"] + + @pytest.mark.asyncio + async def test_run_async_with_filter_policy_merge_combines_init_and_runtime_filters(self, in_memory_doc_store): + in_memory_doc_store.write_documents( + [ + Document(content="python article current", meta={"type": "article", "year": 2020}), + Document(content="python blog current", meta={"type": "blog", "year": 2021}), + Document(content="python article archived", meta={"type": "article", "year": 2019}), + ] + ) + + retriever = InMemoryBM25Retriever( + in_memory_doc_store, + filters={"field": "meta.type", "operator": "==", "value": "article"}, + filter_policy=FilterPolicy.MERGE, + ) + + result = await retriever.run_async( + query="python", filters={"field": "meta.year", "operator": ">=", "value": 2020} + ) + + assert [doc.content for doc in result["documents"]] == ["python article current"] + def test_invalid_run_wrong_store_type(self): SomeOtherDocumentStore = document_store_class("SomeOtherDocumentStore") with pytest.raises(TypeError, match="document_store must be an instance of InMemoryDocumentStore"): diff --git a/test/components/retrievers/test_in_memory_embedding_retriever.py b/test/components/retrievers/test_in_memory_embedding_retriever.py index 024495c312b..4cf6d9dceef 100644 --- a/test/components/retrievers/test_in_memory_embedding_retriever.py +++ b/test/components/retrievers/test_in_memory_embedding_retriever.py @@ -141,6 +141,67 @@ def test_valid_run(self): assert len(result["documents"]) == top_k assert result["documents"][0].embedding == [1.0, 1.0, 1.0, 1.0] + def test_run_with_filter_policy_merge_combines_init_and_runtime_filters(self): + ds = InMemoryDocumentStore(embedding_similarity_function="cosine") + ds.write_documents( + [ + Document( + content="python article current", + embedding=[1.0, 0.0, 0.0, 0.0], + meta={"type": "article", "year": 2020}, + ), + Document( + content="python blog current", embedding=[1.0, 0.0, 0.0, 0.0], meta={"type": "blog", "year": 2021} + ), + Document( + content="python article archived", + embedding=[1.0, 0.0, 0.0, 0.0], + meta={"type": "article", "year": 2019}, + ), + ] + ) + + retriever = InMemoryEmbeddingRetriever( + ds, filters={"field": "meta.type", "operator": "==", "value": "article"}, filter_policy=FilterPolicy.MERGE + ) + + result = retriever.run( + query_embedding=[1.0, 0.0, 0.0, 0.0], filters={"field": "meta.year", "operator": ">=", "value": 2020} + ) + + assert [doc.content for doc in result["documents"]] == ["python article current"] + + @pytest.mark.asyncio + async def test_run_async_with_filter_policy_merge_combines_init_and_runtime_filters(self): + ds = InMemoryDocumentStore(embedding_similarity_function="cosine") + ds.write_documents( + [ + Document( + content="python article current", + embedding=[1.0, 0.0, 0.0, 0.0], + meta={"type": "article", "year": 2020}, + ), + Document( + content="python blog current", embedding=[1.0, 0.0, 0.0, 0.0], meta={"type": "blog", "year": 2021} + ), + Document( + content="python article archived", + embedding=[1.0, 0.0, 0.0, 0.0], + meta={"type": "article", "year": 2019}, + ), + ] + ) + + retriever = InMemoryEmbeddingRetriever( + ds, filters={"field": "meta.type", "operator": "==", "value": "article"}, filter_policy=FilterPolicy.MERGE + ) + + result = await retriever.run_async( + query_embedding=[1.0, 0.0, 0.0, 0.0], filters={"field": "meta.year", "operator": ">=", "value": 2020} + ) + + assert [doc.content for doc in result["documents"]] == ["python article current"] + def test_invalid_run_wrong_store_type(self): SomeOtherDocumentStore = document_store_class("SomeOtherDocumentStore") with pytest.raises(TypeError, match="document_store must be an instance of InMemoryDocumentStore"):