Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import asyncio
from concurrent.futures import ThreadPoolExecutor
from math import inf
from typing import Any

from haystack import Document, component, default_from_dict, default_to_dict
Expand Down Expand Up @@ -148,7 +149,7 @@ def run(self, queries: list[str], retriever_kwargs: dict[str, Any] | None = None

# de-duplicate and sort
docs = _deduplicate_documents(docs)
docs.sort(key=lambda x: x.score or 0.0, reverse=True)
docs.sort(key=lambda doc: doc.score if doc.score is not None else -inf, reverse=True)
return {"documents": docs}

@component.output_types(documents=list[Document])
Expand Down Expand Up @@ -182,7 +183,7 @@ async def _bounded_run_one(query: str) -> list[Document] | None:
results = await _gather_tasks_with_cancel(tasks)
docs: list[Document] = [doc for result in results if result for doc in result]
docs = _deduplicate_documents(docs)
docs.sort(key=lambda x: x.score or 0.0, reverse=True)
docs.sort(key=lambda doc: doc.score if doc.score is not None else -inf, reverse=True)
return {"documents": docs}

def _run_on_thread(self, query: str, retriever_kwargs: dict[str, Any] | None = None) -> list[Document] | None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import asyncio
from concurrent.futures import ThreadPoolExecutor
from math import inf
from typing import Any

from haystack import Document, component, default_from_dict, default_to_dict
Expand Down Expand Up @@ -126,7 +127,7 @@ def run(self, queries: list[str], retriever_kwargs: dict[str, Any] | None = None

# de-duplicate and sort
docs = _deduplicate_documents(docs)
docs.sort(key=lambda x: x.score or 0.0, reverse=True)
docs.sort(key=lambda doc: doc.score if doc.score is not None else -inf, reverse=True)
return {"documents": docs}

@component.output_types(documents=list[Document])
Expand Down Expand Up @@ -160,7 +161,7 @@ async def _bounded_run_one(query: str) -> list[Document] | None:
results = await _gather_tasks_with_cancel(tasks)
docs: list[Document] = [doc for result in results if result for doc in result]
docs = _deduplicate_documents(docs)
docs.sort(key=lambda x: x.score or 0.0, reverse=True)
docs.sort(key=lambda doc: doc.score if doc.score is not None else -inf, reverse=True)
return {"documents": docs}

def _run_on_thread(self, query: str, retriever_kwargs: dict[str, Any] | None = None) -> list[Document] | None:
Expand Down
5 changes: 3 additions & 2 deletions haystack/components/retrievers/text_embedding_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
#
# SPDX-License-Identifier: Apache-2.0

from math import inf
from typing import Any

from haystack import Document, component, default_from_dict, default_to_dict
Expand Down Expand Up @@ -125,7 +126,7 @@ def run(
docs: list[Document] = result["documents"]

# sort
docs.sort(key=lambda x: x.score or 0.0, reverse=True)
docs.sort(key=lambda doc: doc.score if doc.score is not None else -inf, reverse=True)
return {"documents": docs}

@component.output_types(documents=list[Document])
Expand Down Expand Up @@ -153,7 +154,7 @@ async def run_async(
)

docs: list[Document] = result["documents"]
docs.sort(key=lambda x: x.score or 0.0, reverse=True)
docs.sort(key=lambda doc: doc.score if doc.score is not None else -inf, reverse=True)
return {"documents": docs}

def to_dict(self) -> dict[str, Any]:
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
---
fixes:
- |
Fix ``MultiQueryEmbeddingRetriever``, ``MultiQueryTextRetriever``, and
``TextEmbeddingRetriever`` sorting so documents without a score are placed after all
scored documents, including documents with negative scores.
Original file line number Diff line number Diff line change
Expand Up @@ -233,6 +233,30 @@ def run(
assert contents.count("Solar energy is renewable") == 1
assert contents.count("Wind energy is clean") == 1

def test_run_sorts_unscored_documents_after_negative_scores(self):
documents = [
Document(content="unscored", id="none"),
Document(content="negative", id="negative", score=-0.2),
Document(content="zero", id="zero", score=0.0),
Document(content="positive", id="positive", score=0.4),
]

@component
class MockRetriever:
@component.output_types(documents=list[Document])
def run(
self, query_embedding: list[float], filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

retriever = MultiQueryEmbeddingRetriever(
retriever=MockRetriever(), query_embedder=MockTextEmbedder(), max_workers=1
)

result = retriever.run(queries=["query"])

assert [doc.id for doc in result["documents"]] == ["positive", "zero", "negative", "none"]

@pytest.mark.skipif(os.environ.get("OPENAI_API_KEY", "") == "", reason="OPENAI_API_KEY is not set")
@pytest.mark.integration
def test_run_with_filters(self, document_store_with_embeddings):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,35 @@ async def run_async(
scores.append(doc.score)
assert scores == sorted(scores, reverse=True)

@pytest.mark.asyncio
async def test_run_async_sorts_unscored_documents_after_negative_scores(self):
documents = [
Document(content="unscored", id="none"),
Document(content="negative", id="negative", score=-0.2),
Document(content="zero", id="zero", score=0.0),
Document(content="positive", id="positive", score=0.4),
]

@component
class MockRetriever:
@component.output_types(documents=list[Document])
def run(
self, query_embedding: list[float], filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

@component.output_types(documents=list[Document])
async def run_async(
self, query_embedding: list[float], filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

retriever = MultiQueryEmbeddingRetriever(retriever=MockRetriever(), query_embedder=MockTextEmbedder())

result = await retriever.run_async(queries=["query"])

assert [doc.id for doc in result["documents"]] == ["positive", "zero", "negative", "none"]

@pytest.mark.asyncio
async def test_run_async_deduplication(self):
doc2 = Document(content="Wind energy is clean", id="doc2", score=0.8)
Expand Down
25 changes: 24 additions & 1 deletion test/components/retrievers/test_multi_query_text_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,12 @@
# SPDX-License-Identifier: Apache-2.0

import os
from typing import Any
from unittest.mock import ANY, AsyncMock, Mock

import pytest

from haystack import Document, Pipeline
from haystack import Document, Pipeline, component
from haystack.components.generators.chat import OpenAIChatGenerator
from haystack.components.query import QueryExpander
from haystack.components.retrievers import InMemoryBM25Retriever, MultiQueryTextRetriever
Expand Down Expand Up @@ -150,6 +151,28 @@ def test_run_with_multiple_queries(self, document_store_with_docs):
scores = [doc.score for doc in result["documents"] if doc.score is not None]
assert scores == sorted(scores, reverse=True)

def test_run_sorts_unscored_documents_after_negative_scores(self):
documents = [
Document(content="unscored", id="none"),
Document(content="negative", id="negative", score=-0.2),
Document(content="zero", id="zero", score=0.0),
Document(content="positive", id="positive", score=0.4),
]

@component
class MockRetriever:
@component.output_types(documents=list[Document])
def run(
self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

retriever = MultiQueryTextRetriever(retriever=MockRetriever(), max_workers=1)

result = retriever.run(queries=["query"])

assert [doc.id for doc in result["documents"]] == ["positive", "zero", "negative", "none"]

@pytest.mark.integration
def test_run_with_filters(self, document_store_with_docs):
in_memory_retriever = InMemoryBM25Retriever(document_store=document_store_with_docs)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,35 @@ async def test_run_async_with_multiple_queries(self, document_store_with_docs):
scores = [doc.score for doc in result["documents"] if doc.score is not None]
assert scores == sorted(scores, reverse=True)

@pytest.mark.asyncio
async def test_run_async_sorts_unscored_documents_after_negative_scores(self):
documents = [
Document(content="unscored", id="none"),
Document(content="negative", id="negative", score=-0.2),
Document(content="zero", id="zero", score=0.0),
Document(content="positive", id="positive", score=0.4),
]

@component
class MockRetriever:
@component.output_types(documents=list[Document])
def run(
self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

@component.output_types(documents=list[Document])
async def run_async(
self, query: str, filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

retriever = MultiQueryTextRetriever(retriever=MockRetriever())

result = await retriever.run_async(queries=["query"])

assert [doc.id for doc in result["documents"]] == ["positive", "zero", "negative", "none"]

@pytest.mark.asyncio
async def test_run_async_deduplication(self):
doc2 = Document(content="Wind energy is clean", id="doc2", score=0.8)
Expand Down
22 changes: 22 additions & 0 deletions test/components/retrievers/test_text_embedding_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,28 @@ def run(
scores.append(doc.score)
assert scores == sorted(scores, reverse=True)

def test_run_sorts_unscored_documents_after_negative_scores(self):
documents = [
Document(content="unscored", id="none"),
Document(content="negative", id="negative", score=-0.2),
Document(content="zero", id="zero", score=0.0),
Document(content="positive", id="positive", score=0.4),
]

@component
class MockRetriever:
@component.output_types(documents=list[Document])
def run(
self, query_embedding: list[float], filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

retriever = TextEmbeddingRetriever(retriever=MockRetriever(), text_embedder=MockTextEmbedder())

result = retriever.run(query="query")

assert [doc.id for doc in result["documents"]] == ["positive", "zero", "negative", "none"]

def test_to_dict(self):
retriever = TextEmbeddingRetriever(
retriever=InMemoryEmbeddingRetriever(document_store=InMemoryDocumentStore()),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,35 @@ async def run_async(
scores.append(doc.score)
assert scores == sorted(scores, reverse=True)

@pytest.mark.asyncio
async def test_run_async_sorts_unscored_documents_after_negative_scores(self):
documents = [
Document(content="unscored", id="none"),
Document(content="negative", id="negative", score=-0.2),
Document(content="zero", id="zero", score=0.0),
Document(content="positive", id="positive", score=0.4),
]

@component
class MockRetriever:
@component.output_types(documents=list[Document])
def run(
self, query_embedding: list[float], filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

@component.output_types(documents=list[Document])
async def run_async(
self, query_embedding: list[float], filters: dict[str, Any] | None = None, top_k: int | None = None
) -> dict[str, Any]:
return {"documents": documents}

retriever = TextEmbeddingRetriever(retriever=MockRetriever(), text_embedder=MockTextEmbedder())

result = await retriever.run_async(query="query")

assert [doc.id for doc in result["documents"]] == ["positive", "zero", "negative", "none"]

@pytest.mark.asyncio
async def test_run_async_falls_back_to_sync_when_no_run_async(self, document_store_with_categorized_docs):
@component
Expand Down
Loading