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
68 changes: 44 additions & 24 deletions src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -465,43 +465,63 @@ def _get_relevant_documents(
is appended AFTER the generated ones: RRF breaks ties by first
appearance, so reordering here would silently change the ranking.
"""
wanted = expansion_alternates()
queries: list[str] = []
if wanted:
queries = self.query_expander.invoke(
expanded = (
self.query_expander.invoke(
{"question": query}, config={"callbacks": run_manager.get_child()}
)[:wanted]
if self.include_original:
queries.append(query)
return unique_documents(self.retrieve_documents(queries, run_manager))
)
if expansion_alternates()
else None
)
return unique_documents(
self.retrieve_documents(self._queries(query, expanded), run_manager)
)

async def _aget_relevant_documents(
self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
) -> list[Document]:
"""Async twin of the above; must agree with it document for document."""
queries = await self._expand(
query,
lambda: self.query_expander.ainvoke(
expanded = (
await self.query_expander.ainvoke(
{"question": query}, config={"callbacks": run_manager.get_child()}
),
)
if expansion_alternates()
else None
)
return unique_documents(
await self.aretrieve_documents(self._queries(query, expanded), run_manager)
)
return unique_documents(await self.aretrieve_documents(queries, run_manager))

async def _expand(self, query: str, call: Any) -> list[str]:
"""The queries to retrieve for, honouring the configured count.
def _queries(self, query: str, expanded: list[str] | None) -> list[str]:
"""The queries to retrieve for, from whatever the expander returned.

Shared by the sync and async paths deliberately. They are separate
implementations of the same retrieval and this file already carries a
note about them drifting -- but `test_sync_async_equivalence` drives
`retrieve_documents` directly, *below* this step, so a divergence here
would be invisible to the test written to catch exactly that.

At zero the expansion call is skipped entirely rather than made and
discarded -- that call is the larger half of the cost, and making it
anyway would keep the expense while losing the benefit.
At zero alternates the caller skips the expansion call rather than
making it and discarding the result: it is the larger half of the
cost, and discarding would keep the expense, lose the benefit, and
look identical in every other measurement.

The original is appended LAST, as it always was: RRF breaks ties by
first appearance, so reordering silently changes the ranking.
**Never returns an empty list.** With `include_original=False` -- the
default on `from_subdirectory` -- and no alternates, the obvious
assembly yields no queries and retrieval silently returns nothing. A
retriever that finds nothing reads exactly like a question with no
answer, which is the worst available way for this to fail.

The original goes LAST, as it always did: RRF breaks ties by first
appearance, so reordering silently changes the ranking.
"""
wanted = expansion_alternates()
queries: list[str] = []
if wanted:
queries = (await call())[:wanted]
if self.include_original:
queries = list(expanded[:wanted]) if (expanded and wanted) else []
if not queries and not self.include_original:
logger.warning(
"no expanded queries and include_original is off; retrieving "
"for the original question rather than for nothing"
)
if self.include_original or not queries:
queries.append(query)
return queries

Expand Down
51 changes: 44 additions & 7 deletions tests/retrievers/test_query_expansion.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
"""

import asyncio
from pathlib import Path

import pytest

Expand Down Expand Up @@ -40,8 +41,15 @@ def test_a_bad_value_falls_back_loudly_rather_than_disabling_recall(
) -> None:
# The dangerous direction: a typo must not silently turn expansion off,
# because nothing downstream would look any different.
#
# "Loudly" is asserted, not just implied. The first version of this test
# checked only the fallback value and would have passed against a silent
# one -- the same vacuous shape that bit the collection guards.
monkeypatch.setenv(ALTERNATES_ENV, raw)
assert expansion_alternates() == DEFAULT_ALTERNATES
with caplog.at_level("WARNING"):
assert expansion_alternates() == DEFAULT_ALTERNATES
if raw.strip():
assert ALTERNATES_ENV in caplog.text, "fell back without saying so"


class _Expander:
Expand All @@ -53,20 +61,24 @@ async def ainvoke(self, _inputs: object, config: object = None) -> list[str]:
return [f"alt {i}" for i in range(4)]


def _retriever(include_original: bool = True) -> HybridRetriever:
retriever = HybridRetriever.__new__(HybridRetriever)
object.__setattr__(retriever, "include_original", include_original)
return retriever


def _expand(
count: str | None, monkeypatch: pytest.MonkeyPatch
) -> tuple[list[str], int]:
if count is None:
monkeypatch.delenv(ALTERNATES_ENV, raising=False)
else:
monkeypatch.setenv(ALTERNATES_ENV, count)
retriever = HybridRetriever.__new__(HybridRetriever)
object.__setattr__(retriever, "include_original", True)
expander = _Expander()
queries = asyncio.run(
retriever._expand("the question", lambda: expander.ainvoke(None))
)
return queries, expander.calls
expanded = None
if expansion_alternates():
expanded = asyncio.run(expander.ainvoke(None))
return _retriever()._queries("the question", expanded), expander.calls


def test_the_original_question_is_always_asked_and_always_last(
Expand Down Expand Up @@ -99,3 +111,28 @@ def test_a_lower_count_truncates_rather_than_trusting_the_prompt(
queries, calls = _expand("2", monkeypatch)
assert queries == ["alt 0", "alt 1", "the question"]
assert calls == 1


def test_no_expansion_and_no_original_still_asks_something(
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
) -> None:
# `include_original=False` is `from_subdirectory`'s default, so this
# combination is reachable. The obvious assembly produces no queries at
# all and retrieval silently returns nothing -- which reads exactly like
# a question that has no answer, the worst available way to fail.
monkeypatch.setenv(ALTERNATES_ENV, "0")
with caplog.at_level("WARNING"):
queries = _retriever(include_original=False)._queries("the question", None)
assert queries == ["the question"]
assert "rather than for nothing" in caplog.text


def test_the_sync_and_async_paths_cannot_disagree_about_queries() -> None:
# They are separate implementations of the same retrieval, and this is
# the one step `test_sync_async_equivalence` structurally cannot cover:
# it drives `retrieve_documents` directly, below expansion. Asserting the
# source rather than the behaviour, because the behaviour is only equal
# while the code is shared -- which is the thing worth pinning.
source = Path("src/retrievers/csv_chroma.py").read_text()
assert source.count("self._queries(query, expanded)") == 2
assert source.count("def _queries") == 1
Loading