Skip to content

Commit cf31e1f

Browse files
Stage 1: HybridRetriever becomes a plain BaseRetriever
Structural only. All five reaches into LangChain internals are gone: subclassing MultiQueryRetriever -> BaseRetriever overriding a required field via SkipJsonSchema -> field removed stealing .llm_chain from a throwaway instance -> the chain is built directly assigning _retrievers outside pydantic -> a declared field EnsembleRetriever(retrievers=[]) for one method-> RRF vendored here Query expansion, unique_union and the LineListOutputParser are reproduced rather than inherited. The original query is still appended AFTER the generated ones, which matters: RRF resolves ties by first appearance, so reordering there would change the ranking silently. The plan's exit criterion -- zero difference from an end-to-end capture -- turned out to be unachievable, and that is worth recording. SelfQueryRetriever makes an LLM call, so the SAME code run twice already differs on 2 of 8 questions. The old-vs-new difference was also 2 of 8, i.e. entirely within that noise floor, so the comparison could not have shown anything either way. Replaced with a decisive test instead: the vendored RRF is compared directly against EnsembleRetriever.weighted_reciprocal_rank over 2000 random inputs -- varying list counts, overlaps, orderings and empty lists -- and agrees on all of them, including the constant. Perturbing RRF_K to 61 makes that test fail, so it is not vacuous. Also verified through the full chain: sync and async return identical documents in identical order for a fixed query set, and invoke() still yields 40 documents. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
1 parent c6fd825 commit cf31e1f

2 files changed

Lines changed: 216 additions & 25 deletions

File tree

‎src/retrievers/csv_chroma.py‎

Lines changed: 144 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,10 @@
11
import asyncio
22
from collections.abc import Coroutine
33
from pathlib import Path
4-
from typing import Annotated, Any, TypedDict
4+
from typing import Any, TypedDict
55

66
import chromadb.config
77
from langchain.chains.query_constructor.schema import AttributeInfo
8-
from langchain.retrievers import EnsembleRetriever, MultiQueryRetriever
98
from langchain.retrievers.self_query.base import SelfQueryRetriever
109
from langchain_chroma.vectorstores import Chroma
1110
from langchain_community.document_loaders.csv_loader import CSVLoader
@@ -17,10 +16,12 @@
1716
from langchain_core.documents import Document
1817
from langchain_core.embeddings import Embeddings
1918
from langchain_core.language_models.chat_models import BaseChatModel
19+
from langchain_core.output_parsers import BaseOutputParser
2020
from langchain_core.prompts.prompt import PromptTemplate
21+
from langchain_core.retrievers import BaseRetriever
22+
from langchain_core.runnables import Runnable
2123
from nltk.tokenize import word_tokenize
22-
from pydantic import AfterValidator, Field
23-
from pydantic.json_schema import SkipJsonSchema
24+
from pydantic import ConfigDict
2425

2526
chroma_settings = chromadb.config.Settings(anonymized_telemetry=False)
2627

@@ -56,11 +57,6 @@
5657
)
5758

5859

59-
ExcludedField = SkipJsonSchema[
60-
Annotated[Any, Field(default=None, exclude=True), AfterValidator(lambda x: None)]
61-
]
62-
63-
6460
RESULTS_PER_RETRIEVER = 10
6561
# The vector store is asked for more than we intend to keep, because one
6662
# Reactome entity can occupy several rows -- a reaction appears once per
@@ -111,6 +107,85 @@ def dedupe_by_entity(docs: list[Document], limit: int) -> list[Document]:
111107
return kept
112108

113109

110+
RRF_K = 60
111+
"""Reciprocal Rank Fusion constant, from Cormack et al. (SIGIR 2009).
112+
113+
Copied from the value LangChain's EnsembleRetriever uses, so that vendoring the
114+
maths here changes nothing. It damps the difference between adjacent ranks: at
115+
k=60 a first-place hit scores 1/61 and an eleventh-place hit 1/71.
116+
"""
117+
118+
119+
def reciprocal_rank_fusion(
120+
doc_lists: list[list[Document]], weights: list[float] | None = None
121+
) -> list[Document]:
122+
"""Fuse ranked lists by Reciprocal Rank Fusion.
123+
124+
Vendored rather than borrowed. Reaching LangChain's implementation required
125+
constructing `EnsembleRetriever(retrievers=[])` -- an ensemble with no
126+
retrievers -- purely to call one method, on every fusion. Ranking is the
127+
product's core retrieval quality, so a library upgrade should not be able to
128+
reorder results without anyone noticing.
129+
130+
Behaviour is identical to `EnsembleRetriever.weighted_reciprocal_rank`, which
131+
the tests in tests/retrievers/ pin:
132+
133+
- a document's score is the sum over lists of `weight / (rank + RRF_K)`,
134+
with rank counted from 1
135+
- documents are identified by `page_content`, so the same text found by two
136+
retrievers accumulates both scores
137+
- ties keep the order of first appearance across the concatenated lists,
138+
because the sort is stable
139+
"""
140+
if weights is None:
141+
weights = [1 / len(doc_lists)] * len(doc_lists)
142+
if len(doc_lists) != len(weights):
143+
raise ValueError(
144+
f"Got {len(doc_lists)} document lists and {len(weights)} weights; "
145+
"they must correspond one to one."
146+
)
147+
148+
scores: dict[str, float] = {}
149+
for docs, weight in zip(doc_lists, weights, strict=True):
150+
for rank, doc in enumerate(docs, start=1):
151+
scores[doc.page_content] = scores.get(doc.page_content, 0.0) + weight / (
152+
rank + RRF_K
153+
)
154+
155+
seen: set[str] = set()
156+
unique: list[Document] = []
157+
for docs in doc_lists:
158+
for doc in docs:
159+
if doc.page_content not in seen:
160+
seen.add(doc.page_content)
161+
unique.append(doc)
162+
return sorted(unique, key=lambda d: scores[d.page_content], reverse=True)
163+
164+
165+
def unique_documents(documents: list[Document]) -> list[Document]:
166+
"""Drop repeats, keeping first appearance.
167+
168+
Matches MultiQueryRetriever's final `unique_union`, which compares whole
169+
Document objects rather than just their text.
170+
"""
171+
seen: list[Document] = []
172+
for doc in documents:
173+
if doc not in seen:
174+
seen.append(doc)
175+
return seen
176+
177+
178+
class LineListOutputParser(BaseOutputParser[list[str]]):
179+
"""Split an LLM response into non-empty lines.
180+
181+
Same behaviour as LangChain's parser of the same name, which the query
182+
expansion prompt was written against.
183+
"""
184+
185+
def parse(self, text: str) -> list[str]:
186+
return [line for line in text.strip().split("\n") if line]
187+
188+
114189
def list_chroma_subdirectories(directory: Path) -> list[str]:
115190
return [
116191
chroma_file.parent.name for chroma_file in directory.glob("*/chroma.sqlite3")
@@ -140,9 +215,29 @@ class RetrieverDict(TypedDict):
140215
vector: SelfQueryRetriever
141216

142217

143-
class HybridRetriever(MultiQueryRetriever):
144-
retriever: ExcludedField = None
145-
_retrievers: dict[str, RetrieverDict]
218+
class HybridRetriever(BaseRetriever):
219+
"""BM25 and vector search over each Chroma collection, fused by RRF.
220+
221+
A plain BaseRetriever. It previously subclassed MultiQueryRetriever and
222+
reached into LangChain internals in five places -- overriding a required
223+
field to None through SkipJsonSchema, building a throwaway
224+
MultiQueryRetriever to steal its llm_chain, assigning a private attribute
225+
outside pydantic, instantiating an empty EnsembleRetriever to borrow one
226+
method, and overriding two internal methods. None of those are public API,
227+
which is what blocked the LangChain upgrade.
228+
229+
Query expansion is now done here rather than inherited: one LLM call turns
230+
the question into alternates, and `include_original` appends the original
231+
LAST, matching the order MultiQueryRetriever used. Order matters, because
232+
RRF resolves ties by first appearance.
233+
"""
234+
235+
query_expander: Runnable[dict[str, str], list[str]]
236+
include_original: bool = False
237+
collection_retrievers: dict[str, RetrieverDict]
238+
239+
# BM25Retriever and SelfQueryRetriever are not pydantic models.
240+
model_config = ConfigDict(arbitrary_types_allowed=True)
146241

147242
@classmethod
148243
def from_subdirectory(
@@ -189,29 +284,53 @@ def from_subdirectory(
189284
"bm25": bm25_retriever,
190285
"vector": selfq_retriever,
191286
}
192-
llm_chain = MultiQueryRetriever.from_llm(
193-
bm25_retriever, llm, multi_query_prompt, None, include_original
194-
).llm_chain
195-
hybrid_retriever = cls(
196-
llm_chain=llm_chain,
287+
# The expansion chain, built directly. This used to be extracted from a
288+
# throwaway MultiQueryRetriever constructed only to reach its .llm_chain.
289+
return cls(
290+
query_expander=multi_query_prompt | llm | LineListOutputParser(),
197291
include_original=include_original,
198-
_retrievers={},
292+
collection_retrievers=_retrievers,
293+
)
294+
295+
def _get_relevant_documents(
296+
self, query: str, *, run_manager: CallbackManagerForRetrieverRun
297+
) -> list[Document]:
298+
"""Expand the question, retrieve for every variant, fuse, de-duplicate.
299+
300+
Reproduces what MultiQueryRetriever._get_relevant_documents did, so this
301+
stage changes structure without changing results. Note the original query
302+
is appended AFTER the generated ones: RRF breaks ties by first
303+
appearance, so reordering here would silently change the ranking.
304+
"""
305+
queries = self.query_expander.invoke(
306+
{"question": query}, config={"callbacks": run_manager.get_child()}
307+
)
308+
if self.include_original:
309+
queries.append(query)
310+
return unique_documents(self.retrieve_documents(queries, run_manager))
311+
312+
async def _aget_relevant_documents(
313+
self, query: str, *, run_manager: AsyncCallbackManagerForRetrieverRun
314+
) -> list[Document]:
315+
"""Async twin of the above; must agree with it document for document."""
316+
queries = await self.query_expander.ainvoke(
317+
{"question": query}, config={"callbacks": run_manager.get_child()}
199318
)
200-
hybrid_retriever._retrievers = _retrievers
201-
return hybrid_retriever
319+
if self.include_original:
320+
queries.append(query)
321+
return unique_documents(await self.aretrieve_documents(queries, run_manager))
202322

203323
def weighted_reciprocal_rank(
204324
self, doc_lists: list[list[Document]]
205325
) -> list[Document]:
206-
return EnsembleRetriever(
207-
retrievers=[], weights=[1 / len(doc_lists)] * len(doc_lists)
208-
).weighted_reciprocal_rank(doc_lists)
326+
"""Kept as a method so the existing tests address it unchanged."""
327+
return reciprocal_rank_fusion(doc_lists)
209328

210329
def retrieve_documents(
211330
self, queries: list[str], run_manager: CallbackManagerForRetrieverRun
212331
) -> list[Document]:
213332
subdirectory_docs: list[Document] = []
214-
for subdirectory, retrievers in self._retrievers.items():
333+
for subdirectory, retrievers in self.collection_retrievers.items():
215334
bm25_retriever = retrievers["bm25"]
216335
vector_retriever = retrievers["vector"]
217336
doc_lists: list[list[Document]] = []
@@ -250,7 +369,7 @@ async def aretrieve_documents(
250369
run_manager: AsyncCallbackManagerForRetrieverRun,
251370
) -> list[Document]:
252371
subdirectory_results: dict[str, list[Coroutine[Any, Any, list[Document]]]] = {}
253-
for subdirectory, retrievers in self._retrievers.items():
372+
for subdirectory, retrievers in self.collection_retrievers.items():
254373
bm25_retriever = retrievers["bm25"]
255374
vector_retriever = retrievers["vector"]
256375
subdirectory_results[subdirectory] = []
Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,72 @@
1+
"""The vendored RRF must behave exactly like the LangChain method it replaces.
2+
3+
`reciprocal_rank_fusion` was copied out of `EnsembleRetriever` so that ranking --
4+
the core of retrieval quality -- is owned here and cannot be reordered by a
5+
library upgrade without anyone noticing. That is only safe if the copy is
6+
faithful, so this compares the two implementations directly rather than trusting
7+
that the maths was transcribed correctly.
8+
9+
This test is the reason Stage 1 of the retriever rewrite can claim "no behaviour
10+
change". An end-to-end diff cannot show it: SelfQueryRetriever makes an LLM call,
11+
so the same code run twice already differs on roughly a quarter of questions.
12+
"""
13+
14+
import random
15+
16+
import pytest
17+
from langchain_core.documents import Document
18+
19+
from retrievers.csv_chroma import RRF_K, reciprocal_rank_fusion
20+
21+
pytest.importorskip("langchain", reason="retrieval stack not installed")
22+
23+
from langchain.retrievers import EnsembleRetriever # noqa: E402
24+
25+
pytestmark = pytest.mark.requires_retrieval_stack
26+
27+
28+
def _langchain_rrf(doc_lists: list[list[Document]]) -> list[str]:
29+
n = len(doc_lists)
30+
ensemble = EnsembleRetriever(retrievers=[], weights=[1 / n] * n)
31+
return [d.page_content for d in ensemble.weighted_reciprocal_rank(doc_lists)]
32+
33+
34+
def _ours(doc_lists: list[list[Document]]) -> list[str]:
35+
return [d.page_content for d in reciprocal_rank_fusion(doc_lists)]
36+
37+
38+
def test_the_constant_matches_the_source_it_was_copied_from() -> None:
39+
assert EnsembleRetriever(retrievers=[], weights=[1.0]).c == RRF_K
40+
41+
42+
@pytest.mark.parametrize("seed", range(25))
43+
def test_equivalent_on_random_inputs(seed: int) -> None:
44+
"""Random shapes: varying list counts, overlaps, orderings and empties."""
45+
# S311: generating test data, not keys.
46+
rng = random.Random(seed) # noqa: S311
47+
for _ in range(80):
48+
n_lists = rng.randint(1, 5)
49+
pool = [f"doc-{i}" for i in range(rng.randint(1, 12))]
50+
lists = [
51+
[
52+
Document(page_content=c)
53+
for c in rng.sample(pool, rng.randint(0, len(pool)))
54+
]
55+
for _ in range(n_lists)
56+
]
57+
assert _ours(lists) == _langchain_rrf(lists), f"diverged on {lists}"
58+
59+
60+
def test_equivalent_when_every_list_is_empty() -> None:
61+
assert _ours([[], []]) == _langchain_rrf([[], []]) == []
62+
63+
64+
def test_equivalent_when_one_retriever_returns_nothing() -> None:
65+
"""A retriever finding nothing must not drop the other's results."""
66+
found = [Document(page_content="a"), Document(page_content="b")]
67+
assert _ours([found, []]) == _langchain_rrf([found, []]) == ["a", "b"]
68+
69+
70+
def test_mismatched_weights_are_rejected() -> None:
71+
with pytest.raises(ValueError, match="one to one"):
72+
reciprocal_rank_fusion([[Document(page_content="a")]], weights=[0.5, 0.5])

0 commit comments

Comments
 (0)