Skip to content
Merged
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
137 changes: 84 additions & 53 deletions haystack/components/retrievers/sentence_window_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand All @@ -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:
Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -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.
64 changes: 63 additions & 1 deletion test/components/retrievers/test_sentence_window_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@

import random
import re
from unittest.mock import ANY, Mock
from unittest.mock import ANY, Mock, patch

import pytest

Expand Down Expand Up @@ -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")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@

import random
import re
from unittest.mock import AsyncMock, Mock
from unittest.mock import AsyncMock, Mock, patch

import pytest

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