diff --git a/haystack/components/retrievers/sentence_window_retriever.py b/haystack/components/retrievers/sentence_window_retriever.py index 0eecfb1297..3e60a1c39e 100644 --- a/haystack/components/retrievers/sentence_window_retriever.py +++ b/haystack/components/retrievers/sentence_window_retriever.py @@ -201,14 +201,11 @@ def run(self, retrieved_documents: list[Document], window_size: int | None = Non self._validate_window_size(window_size) self._raise_if_documents_do_not_have_expected_metadata(retrieved_documents) - context_text = [] - context_documents = [] - for doc in retrieved_documents: - text, docs = self._retrieve_context_for_document(doc, window_size) - context_text.append(text) - context_documents.extend(docs) + windows = self._get_windows(retrieved_documents) + filters = self._build_filters(windows, window_size) + fetched_documents = self.document_store.filter_documents(filters) if filters else [] - return {"context_windows": context_text, "context_documents": context_documents} + return self._assemble_context(retrieved_documents, windows, window_size, fetched_documents) @component.output_types(context_windows=list[str], context_documents=list[Document]) async def run_async(self, retrieved_documents: list[Document], window_size: int | None = None) -> dict[str, Any]: @@ -235,14 +232,16 @@ async def run_async(self, retrieved_documents: list[Document], window_size: int self._validate_window_size(window_size) self._raise_if_documents_do_not_have_expected_metadata(retrieved_documents) - context_text = [] - context_documents = [] - for doc in retrieved_documents: - text, docs = await self._retrieve_context_for_document_async(doc, window_size) - context_text.append(text) - context_documents.extend(docs) + windows = self._get_windows(retrieved_documents) + filters = self._build_filters(windows, window_size) + # Ignoring type error because DocumentStore protocol doesn't define filter_documents_async + fetched_documents = ( + await self.document_store.filter_documents_async(filters) # type: ignore[attr-defined] + if filters + else [] + ) - return {"context_windows": context_text, "context_documents": context_documents} + return self._assemble_context(retrieved_documents, windows, window_size, fetched_documents) @staticmethod def _validate_window_size(window_size: int) -> None: @@ -262,50 +261,82 @@ def _raise_if_documents_do_not_have_expected_metadata(self, retrieved_documents: ): raise ValueError(f"The retrieved documents must have '{self.source_id_meta_field}' in their metadata.") - def _retrieve_context_for_document(self, doc: Document, window_size: int) -> tuple[str, list[Document]]: - source_ids = [doc.meta.get(field) for field in self._source_id_meta_fields] - split_id = doc.meta.get(self.split_id_meta_field) - - if any(source_id is None for source_id in source_ids) or split_id is None: - logger.warning( - "Document {doc_id} is missing required metadata fields to be used with " - "SentenceWindowRetriever: {source_id} or {split_id}. Skipping context retrieval for this document.", - doc_id=doc.id, - source_id=self._source_id_meta_fields, - split_id=self.split_id_meta_field, - ) - return doc.content or "", [doc] + def _get_windows(self, retrieved_documents: list[Document]) -> list[tuple[list[Any], int] | None]: + """ + Extract the source IDs and split ID of each retrieved document, or `None` if the metadata is missing. + """ + windows: list[tuple[list[Any], int] | None] = [] + for doc in retrieved_documents: + source_ids = [doc.meta.get(field) for field in self._source_id_meta_fields] + split_id = doc.meta.get(self.split_id_meta_field) + if any(source_id is None for source_id in source_ids) or split_id is None: + logger.warning( + "Document {doc_id} is missing required metadata fields to be used with " + "SentenceWindowRetriever: {source_id} or {split_id}. Skipping context retrieval for this document.", + doc_id=doc.id, + source_id=self._source_id_meta_fields, + split_id=self.split_id_meta_field, + ) + windows.append(None) + else: + windows.append((source_ids, split_id)) + return windows + + def _build_filters(self, windows: list[tuple[list[Any], int] | None], window_size: int) -> dict[str, Any] | None: + """ + Combine the windows of all retrieved documents into one filter, so the Document Store is queried only once. + """ + conditions: list[dict[str, Any]] = [] + for window in windows: + if window is None: + continue + source_ids, split_id = window + condition = self._build_filter_conditions(split_id, window_size, source_ids) + if condition not in conditions: + conditions.append(condition) + if not conditions: + return None + return {"operator": "OR", "conditions": conditions} + + def _assemble_context( + self, + retrieved_documents: list[Document], + windows: list[tuple[list[Any], int] | None], + window_size: int, + fetched_documents: list[Document], + ) -> dict[str, Any]: + """ + Split the documents fetched with the combined filter back into the context window of each retrieved document. + """ + context_text = [] + context_documents = [] + for doc, window in zip(retrieved_documents, windows, strict=True): + if window is None: + context_text.append(doc.content or "") + context_documents.append(doc) + continue - assert split_id is not None - filter_conditions = self._build_filter_conditions(split_id, window_size, source_ids) - context_docs = self.document_store.filter_documents(filter_conditions) - context_text = self.merge_documents_text(context_docs) - context_docs_sorted = sorted(context_docs, key=lambda doc: doc.meta[self.split_id_meta_field]) + source_ids, split_id = window + context_docs = [ + fetched + for fetched in fetched_documents + if self._is_in_window(fetched, source_ids, split_id - window_size, split_id + window_size) + ] + context_text.append(self.merge_documents_text(context_docs)) + context_documents.extend(sorted(context_docs, key=lambda d: d.meta[self.split_id_meta_field])) - return context_text, context_docs_sorted + return {"context_windows": context_text, "context_documents": context_documents} - async def _retrieve_context_for_document_async(self, doc: Document, window_size: int) -> tuple[str, list[Document]]: - source_ids = [doc.meta.get(field) for field in self._source_id_meta_fields] + def _is_in_window(self, doc: Document, source_ids: list[Any], min_split_id: int, max_split_id: int) -> bool: split_id = doc.meta.get(self.split_id_meta_field) - - if any(source_id is None for source_id in source_ids) or split_id is None: - logger.warning( - "Document {doc_id} is missing required metadata fields to be used with " - "SentenceWindowRetriever: {source_id} or {split_id}. Skipping context retrieval for this document.", - doc_id=doc.id, - source_id=self._source_id_meta_fields, - split_id=self.split_id_meta_field, + return ( + split_id is not None + and min_split_id <= split_id <= max_split_id + and all( + doc.meta.get(field) == source_id + for field, source_id in zip(self._source_id_meta_fields, source_ids, strict=True) ) - return doc.content or "", [doc] - - assert split_id is not None - filter_conditions = self._build_filter_conditions(split_id, window_size, source_ids) - # Ignoring type error because DocumentStore protocol doesn't define filter_documents_async - context_docs = await self.document_store.filter_documents_async(filter_conditions) # type: ignore[attr-defined] - context_text = self.merge_documents_text(context_docs) - context_docs_sorted = sorted(context_docs, key=lambda doc: doc.meta[self.split_id_meta_field]) - - return context_text, context_docs_sorted + ) def _build_filter_conditions(self, split_id: int, window_size: int, source_ids: list[Any]) -> dict[str, Any]: min_before = split_id - window_size diff --git a/releasenotes/notes/sentence-window-retriever-single-filter-call-f4bf129e4fd0055a.yaml b/releasenotes/notes/sentence-window-retriever-single-filter-call-f4bf129e4fd0055a.yaml new file mode 100644 index 0000000000..c7f4233f8a --- /dev/null +++ b/releasenotes/notes/sentence-window-retriever-single-filter-call-f4bf129e4fd0055a.yaml @@ -0,0 +1,8 @@ +--- +enhancements: + - | + ``SentenceWindowRetriever`` now queries the Document Store once per ``run`` or ``run_async`` call instead of + once per retrieved document. It combines the windows of all retrieved documents into a single ``OR`` filter and + assigns the returned documents to each window in memory. This reduces the load on Document Stores such as + OpenSearch or Elasticsearch when many documents are retrieved. The ``context_windows`` and ``context_documents`` + outputs are unchanged. diff --git a/test/components/retrievers/test_sentence_window_retriever.py b/test/components/retrievers/test_sentence_window_retriever.py index 72e221835b..9cc05f3046 100644 --- a/test/components/retrievers/test_sentence_window_retriever.py +++ b/test/components/retrievers/test_sentence_window_retriever.py @@ -4,7 +4,7 @@ import random import re -from unittest.mock import ANY, Mock +from unittest.mock import ANY, Mock, patch import pytest @@ -311,6 +311,68 @@ def test_run_with_multiple_source_ids(self, in_memory_doc_store): assert len(result["context_documents"]) == 3 assert all(doc.meta["section"] == "1" for doc in result["context_documents"]) + def test_run_queries_document_store_once_for_all_retrieved_documents(self, in_memory_doc_store): + docs = [ + Document(content=f"{source}{split_id}.", meta={"source_id": source, "split_id": split_id}) + for source in ("a", "b") + for split_id in range(10) + ] + in_memory_doc_store.write_documents(docs) + retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, window_size=1) + + with patch.object( + in_memory_doc_store, "filter_documents", wraps=in_memory_doc_store.filter_documents + ) as filter_documents: + # docs[2] is a2 and docs[15] is b5; the duplicate a2 must not add a second condition to the filter + result = retriever.run(retrieved_documents=[docs[2], docs[15], docs[2]]) + + filter_documents.assert_called_once() + assert len(filter_documents.call_args.args[0]["conditions"]) == 2 + assert result["context_windows"] == ["a1.a2.a3.", "b4.b5.b6.", "a1.a2.a3."] + assert [doc.content for doc in result["context_documents"]] == [ + "a1.", + "a2.", + "a3.", + "b4.", + "b5.", + "b6.", + "a1.", + "a2.", + "a3.", + ] + + def test_run_with_documents_missing_metadata_queries_document_store_once(self, in_memory_doc_store): + docs = [ + Document(content=f"{split_id}.", meta={"source_id": "a", "split_id": split_id}) for split_id in range(5) + ] + in_memory_doc_store.write_documents(docs) + retriever = SentenceWindowRetriever( + document_store=in_memory_doc_store, window_size=1, raise_on_missing_meta_fields=False + ) + doc_without_meta = Document(content="No metadata.") + + with patch.object( + in_memory_doc_store, "filter_documents", wraps=in_memory_doc_store.filter_documents + ) as filter_documents: + result = retriever.run(retrieved_documents=[doc_without_meta, docs[2]]) + + filter_documents.assert_called_once() + assert result["context_windows"] == ["No metadata.", "1.2.3."] + assert result["context_documents"] == [doc_without_meta, docs[1], docs[2], docs[3]] + + def test_run_does_not_query_document_store_without_documents_to_expand(self, in_memory_doc_store): + retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, raise_on_missing_meta_fields=False) + doc_without_meta = Document(content="No metadata.") + + with patch.object(in_memory_doc_store, "filter_documents") as filter_documents: + assert retriever.run(retrieved_documents=[]) == {"context_windows": [], "context_documents": []} + assert retriever.run(retrieved_documents=[doc_without_meta]) == { + "context_windows": ["No metadata."], + "context_documents": [doc_without_meta], + } + + filter_documents.assert_not_called() + @pytest.mark.integration def test_run_with_pipeline(self, in_memory_doc_store): splitter = DocumentSplitter(split_length=1, split_overlap=0, split_by="period") diff --git a/test/components/retrievers/test_sentence_window_retriever_async.py b/test/components/retrievers/test_sentence_window_retriever_async.py index de85b08e60..4d3c728eeb 100644 --- a/test/components/retrievers/test_sentence_window_retriever_async.py +++ b/test/components/retrievers/test_sentence_window_retriever_async.py @@ -4,7 +4,7 @@ import random import re -from unittest.mock import AsyncMock, Mock +from unittest.mock import AsyncMock, Mock, patch import pytest @@ -204,6 +204,71 @@ async def test_run_async_with_multiple_source_ids(self, in_memory_doc_store): assert len(result["context_documents"]) == 3 assert all(doc.meta["section"] == "1" for doc in result["context_documents"]) + @pytest.mark.asyncio + async def test_run_async_queries_document_store_once_for_all_retrieved_documents(self, in_memory_doc_store): + docs = [ + Document(content=f"{source}{split_id}.", meta={"source_id": source, "split_id": split_id}) + for source in ("a", "b") + for split_id in range(10) + ] + in_memory_doc_store.write_documents(docs) + retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, window_size=1) + + with patch.object( + in_memory_doc_store, "filter_documents_async", wraps=in_memory_doc_store.filter_documents_async + ) as filter_documents_async: + # docs[2] is a2 and docs[15] is b5; the duplicate a2 must not add a second condition to the filter + result = await retriever.run_async(retrieved_documents=[docs[2], docs[15], docs[2]]) + + filter_documents_async.assert_awaited_once() + assert len(filter_documents_async.call_args.args[0]["conditions"]) == 2 + assert result["context_windows"] == ["a1.a2.a3.", "b4.b5.b6.", "a1.a2.a3."] + assert [doc.content for doc in result["context_documents"]] == [ + "a1.", + "a2.", + "a3.", + "b4.", + "b5.", + "b6.", + "a1.", + "a2.", + "a3.", + ] + + @pytest.mark.asyncio + async def test_run_async_with_documents_missing_metadata_queries_document_store_once(self, in_memory_doc_store): + docs = [ + Document(content=f"{split_id}.", meta={"source_id": "a", "split_id": split_id}) for split_id in range(5) + ] + in_memory_doc_store.write_documents(docs) + retriever = SentenceWindowRetriever( + document_store=in_memory_doc_store, window_size=1, raise_on_missing_meta_fields=False + ) + doc_without_meta = Document(content="No metadata.") + + with patch.object( + in_memory_doc_store, "filter_documents_async", wraps=in_memory_doc_store.filter_documents_async + ) as filter_documents_async: + result = await retriever.run_async(retrieved_documents=[doc_without_meta, docs[2]]) + + filter_documents_async.assert_awaited_once() + assert result["context_windows"] == ["No metadata.", "1.2.3."] + assert result["context_documents"] == [doc_without_meta, docs[1], docs[2], docs[3]] + + @pytest.mark.asyncio + async def test_run_async_does_not_query_document_store_without_documents_to_expand(self, in_memory_doc_store): + retriever = SentenceWindowRetriever(document_store=in_memory_doc_store, raise_on_missing_meta_fields=False) + doc_without_meta = Document(content="No metadata.") + + with patch.object(in_memory_doc_store, "filter_documents_async") as filter_documents_async: + assert await retriever.run_async(retrieved_documents=[]) == {"context_windows": [], "context_documents": []} + assert await retriever.run_async(retrieved_documents=[doc_without_meta]) == { + "context_windows": ["No metadata."], + "context_documents": [doc_without_meta], + } + + filter_documents_async.assert_not_awaited() + @pytest.mark.asyncio @pytest.mark.integration async def test_run_async_with_pipeline(self, in_memory_doc_store):