diff --git a/src/retrievers/csv_chroma.py b/src/retrievers/csv_chroma.py index 8e5fdc00..30333d2d 100644 --- a/src/retrievers/csv_chroma.py +++ b/src/retrievers/csv_chroma.py @@ -1,11 +1,10 @@ import asyncio from collections.abc import Coroutine from pathlib import Path -from typing import Annotated, Any, TypedDict +from typing import Any, TypedDict import chromadb.config from langchain.chains.query_constructor.schema import AttributeInfo -from langchain.retrievers import EnsembleRetriever, MultiQueryRetriever from langchain.retrievers.self_query.base import SelfQueryRetriever from langchain_chroma.vectorstores import Chroma from langchain_community.document_loaders.csv_loader import CSVLoader @@ -17,10 +16,12 @@ from langchain_core.documents import Document from langchain_core.embeddings import Embeddings from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.output_parsers import BaseOutputParser from langchain_core.prompts.prompt import PromptTemplate +from langchain_core.retrievers import BaseRetriever +from langchain_core.runnables import Runnable from nltk.tokenize import word_tokenize -from pydantic import AfterValidator, Field -from pydantic.json_schema import SkipJsonSchema +from pydantic import ConfigDict chroma_settings = chromadb.config.Settings(anonymized_telemetry=False) @@ -56,11 +57,6 @@ ) -ExcludedField = SkipJsonSchema[ - Annotated[Any, Field(default=None, exclude=True), AfterValidator(lambda x: None)] -] - - RESULTS_PER_RETRIEVER = 10 # The vector store is asked for more than we intend to keep, because one # 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]: return kept +RRF_K = 60 +"""Reciprocal Rank Fusion constant, from Cormack et al. (SIGIR 2009). + +Copied from the value LangChain's EnsembleRetriever uses, so that vendoring the +maths here changes nothing. It damps the difference between adjacent ranks: at +k=60 a first-place hit scores 1/61 and an eleventh-place hit 1/71. +""" + + +def reciprocal_rank_fusion( + doc_lists: list[list[Document]], weights: list[float] | None = None +) -> list[Document]: + """Fuse ranked lists by Reciprocal Rank Fusion. + + Vendored rather than borrowed. Reaching LangChain's implementation required + constructing `EnsembleRetriever(retrievers=[])` -- an ensemble with no + retrievers -- purely to call one method, on every fusion. Ranking is the + product's core retrieval quality, so a library upgrade should not be able to + reorder results without anyone noticing. + + Behaviour is identical to `EnsembleRetriever.weighted_reciprocal_rank`, which + the tests in tests/retrievers/ pin: + + - a document's score is the sum over lists of `weight / (rank + RRF_K)`, + with rank counted from 1 + - documents are identified by `page_content`, so the same text found by two + retrievers accumulates both scores + - ties keep the order of first appearance across the concatenated lists, + because the sort is stable + """ + if weights is None: + weights = [1 / len(doc_lists)] * len(doc_lists) + if len(doc_lists) != len(weights): + raise ValueError( + f"Got {len(doc_lists)} document lists and {len(weights)} weights; " + "they must correspond one to one." + ) + + scores: dict[str, float] = {} + for docs, weight in zip(doc_lists, weights, strict=True): + for rank, doc in enumerate(docs, start=1): + scores[doc.page_content] = scores.get(doc.page_content, 0.0) + weight / ( + rank + RRF_K + ) + + seen: set[str] = set() + unique: list[Document] = [] + for docs in doc_lists: + for doc in docs: + if doc.page_content not in seen: + seen.add(doc.page_content) + unique.append(doc) + return sorted(unique, key=lambda d: scores[d.page_content], reverse=True) + + +def unique_documents(documents: list[Document]) -> list[Document]: + """Drop repeats, keeping first appearance. + + Matches MultiQueryRetriever's final `unique_union`, which compares whole + Document objects rather than just their text. + """ + seen: list[Document] = [] + for doc in documents: + if doc not in seen: + seen.append(doc) + return seen + + +class LineListOutputParser(BaseOutputParser[list[str]]): + """Split an LLM response into non-empty lines. + + Same behaviour as LangChain's parser of the same name, which the query + expansion prompt was written against. + """ + + def parse(self, text: str) -> list[str]: + return [line for line in text.strip().split("\n") if line] + + def list_chroma_subdirectories(directory: Path) -> list[str]: return [ chroma_file.parent.name for chroma_file in directory.glob("*/chroma.sqlite3") @@ -140,9 +215,29 @@ class RetrieverDict(TypedDict): vector: SelfQueryRetriever -class HybridRetriever(MultiQueryRetriever): - retriever: ExcludedField = None - _retrievers: dict[str, RetrieverDict] +class HybridRetriever(BaseRetriever): + """BM25 and vector search over each Chroma collection, fused by RRF. + + A plain BaseRetriever. It previously subclassed MultiQueryRetriever and + reached into LangChain internals in five places -- overriding a required + field to None through SkipJsonSchema, building a throwaway + MultiQueryRetriever to steal its llm_chain, assigning a private attribute + outside pydantic, instantiating an empty EnsembleRetriever to borrow one + method, and overriding two internal methods. None of those are public API, + which is what blocked the LangChain upgrade. + + Query expansion is now done here rather than inherited: one LLM call turns + the question into alternates, and `include_original` appends the original + LAST, matching the order MultiQueryRetriever used. Order matters, because + RRF resolves ties by first appearance. + """ + + query_expander: Runnable[dict[str, str], list[str]] + include_original: bool = False + collection_retrievers: dict[str, RetrieverDict] + + # BM25Retriever and SelfQueryRetriever are not pydantic models. + model_config = ConfigDict(arbitrary_types_allowed=True) @classmethod def from_subdirectory( @@ -189,29 +284,53 @@ def from_subdirectory( "bm25": bm25_retriever, "vector": selfq_retriever, } - llm_chain = MultiQueryRetriever.from_llm( - bm25_retriever, llm, multi_query_prompt, None, include_original - ).llm_chain - hybrid_retriever = cls( - llm_chain=llm_chain, + # The expansion chain, built directly. This used to be extracted from a + # throwaway MultiQueryRetriever constructed only to reach its .llm_chain. + return cls( + query_expander=multi_query_prompt | llm | LineListOutputParser(), include_original=include_original, - _retrievers={}, + collection_retrievers=_retrievers, + ) + + def _get_relevant_documents( + self, query: str, *, run_manager: CallbackManagerForRetrieverRun + ) -> list[Document]: + """Expand the question, retrieve for every variant, fuse, de-duplicate. + + Reproduces what MultiQueryRetriever._get_relevant_documents did, so this + stage changes structure without changing results. Note the original query + is appended AFTER the generated ones: RRF breaks ties by first + appearance, so reordering here would silently change the ranking. + """ + queries = self.query_expander.invoke( + {"question": query}, config={"callbacks": run_manager.get_child()} + ) + if self.include_original: + queries.append(query) + return unique_documents(self.retrieve_documents(queries, 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.query_expander.ainvoke( + {"question": query}, config={"callbacks": run_manager.get_child()} ) - hybrid_retriever._retrievers = _retrievers - return hybrid_retriever + if self.include_original: + queries.append(query) + return unique_documents(await self.aretrieve_documents(queries, run_manager)) def weighted_reciprocal_rank( self, doc_lists: list[list[Document]] ) -> list[Document]: - return EnsembleRetriever( - retrievers=[], weights=[1 / len(doc_lists)] * len(doc_lists) - ).weighted_reciprocal_rank(doc_lists) + """Kept as a method so the existing tests address it unchanged.""" + return reciprocal_rank_fusion(doc_lists) def retrieve_documents( self, queries: list[str], run_manager: CallbackManagerForRetrieverRun ) -> list[Document]: subdirectory_docs: list[Document] = [] - for subdirectory, retrievers in self._retrievers.items(): + for subdirectory, retrievers in self.collection_retrievers.items(): bm25_retriever = retrievers["bm25"] vector_retriever = retrievers["vector"] doc_lists: list[list[Document]] = [] @@ -250,7 +369,7 @@ async def aretrieve_documents( run_manager: AsyncCallbackManagerForRetrieverRun, ) -> list[Document]: subdirectory_results: dict[str, list[Coroutine[Any, Any, list[Document]]]] = {} - for subdirectory, retrievers in self._retrievers.items(): + for subdirectory, retrievers in self.collection_retrievers.items(): bm25_retriever = retrievers["bm25"] vector_retriever = retrievers["vector"] subdirectory_results[subdirectory] = [] diff --git a/tests/retrievers/test_rrf_equivalence.py b/tests/retrievers/test_rrf_equivalence.py new file mode 100644 index 00000000..0b5ebc75 --- /dev/null +++ b/tests/retrievers/test_rrf_equivalence.py @@ -0,0 +1,72 @@ +"""The vendored RRF must behave exactly like the LangChain method it replaces. + +`reciprocal_rank_fusion` was copied out of `EnsembleRetriever` so that ranking -- +the core of retrieval quality -- is owned here and cannot be reordered by a +library upgrade without anyone noticing. That is only safe if the copy is +faithful, so this compares the two implementations directly rather than trusting +that the maths was transcribed correctly. + +This test is the reason Stage 1 of the retriever rewrite can claim "no behaviour +change". An end-to-end diff cannot show it: SelfQueryRetriever makes an LLM call, +so the same code run twice already differs on roughly a quarter of questions. +""" + +import random + +import pytest +from langchain_core.documents import Document + +from retrievers.csv_chroma import RRF_K, reciprocal_rank_fusion + +pytest.importorskip("langchain", reason="retrieval stack not installed") + +from langchain.retrievers import EnsembleRetriever # noqa: E402 + +pytestmark = pytest.mark.requires_retrieval_stack + + +def _langchain_rrf(doc_lists: list[list[Document]]) -> list[str]: + n = len(doc_lists) + ensemble = EnsembleRetriever(retrievers=[], weights=[1 / n] * n) + return [d.page_content for d in ensemble.weighted_reciprocal_rank(doc_lists)] + + +def _ours(doc_lists: list[list[Document]]) -> list[str]: + return [d.page_content for d in reciprocal_rank_fusion(doc_lists)] + + +def test_the_constant_matches_the_source_it_was_copied_from() -> None: + assert EnsembleRetriever(retrievers=[], weights=[1.0]).c == RRF_K + + +@pytest.mark.parametrize("seed", range(25)) +def test_equivalent_on_random_inputs(seed: int) -> None: + """Random shapes: varying list counts, overlaps, orderings and empties.""" + # S311: generating test data, not keys. + rng = random.Random(seed) # noqa: S311 + for _ in range(80): + n_lists = rng.randint(1, 5) + pool = [f"doc-{i}" for i in range(rng.randint(1, 12))] + lists = [ + [ + Document(page_content=c) + for c in rng.sample(pool, rng.randint(0, len(pool))) + ] + for _ in range(n_lists) + ] + assert _ours(lists) == _langchain_rrf(lists), f"diverged on {lists}" + + +def test_equivalent_when_every_list_is_empty() -> None: + assert _ours([[], []]) == _langchain_rrf([[], []]) == [] + + +def test_equivalent_when_one_retriever_returns_nothing() -> None: + """A retriever finding nothing must not drop the other's results.""" + found = [Document(page_content="a"), Document(page_content="b")] + assert _ours([found, []]) == _langchain_rrf([found, []]) == ["a", "b"] + + +def test_mismatched_weights_are_rejected() -> None: + with pytest.raises(ValueError, match="one to one"): + reciprocal_rank_fusion([[Document(page_content="a")]], weights=[0.5, 0.5])