From b4adb7474d2a3bd5bbca03ef523d0267883c8b1b Mon Sep 17 00:00:00 2001 From: Adam Wright Date: Fri, 18 Sep 2026 19:07:58 +0000 Subject: [PATCH] Add collection selection, and record the route the data model got wrong resolve_collections with the spec's asymmetry: empty or absent searches everything, a valid selection narrows, and any unknown name widens back to everything with a WARNING naming it. Narrowing on a selection we do not understand costs the answer; widening costs latency, so every failure widens. **The data model's route does not exist.** It sends the selection through `RunnableConfig["configurable"]["collections"]` into the retriever. Measured against langchain-core 0.2.14: config is not passed to `_get_relevant_documents` as a kwarg, and `var_child_runnable_config` is unset inside a retriever run. A per-request attribute would race, because the retriever is built once and shared. Rebuilding per request discards the BM25 indexes. So the selection travels in a ContextVar, which asyncio copies per task. A test runs a narrow request and a wide one concurrently and asserts neither sees the other's selection -- that isolation is the whole reason for the choice, so it is pinned rather than assumed. A characterization test records today's behaviour, every collection searched, so that narrowing becomes a visible change rather than a silent one. Not yet wired: QueryIntent has no `collections` field and nothing sets the ContextVar, so behaviour is unchanged. That is the next task. Co-Authored-By: Claude Opus 5 --- specs/009-collection-routing/research.md | 25 +++ src/retrievers/csv_chroma.py | 74 +++++++- tests/retrievers/test_collection_selection.py | 161 ++++++++++++++++++ 3 files changed, 257 insertions(+), 3 deletions(-) create mode 100644 tests/retrievers/test_collection_selection.py diff --git a/specs/009-collection-routing/research.md b/specs/009-collection-routing/research.md index 8488d9c..9ea14eb 100644 --- a/specs/009-collection-routing/research.md +++ b/specs/009-collection-routing/research.md @@ -89,3 +89,28 @@ on 2026-09-17 (PR #227), so measuring the sync path is currently valid. **Decision**: the equivalence test is a prerequisite of this work, not a nice-to-have. Both paths must filter by the same selection, and the existing test must be extended to assert that -- otherwise the measurement stops describing what users get. + +## The config route the data model describes does not exist + +**Measured 2026-09-18** against langchain-core 0.2.14, the version pinned here. + +data-model.md routes the selection as +`RunnableConfig["configurable"]["collections"]` into +`HybridRetriever.retrieve_documents`. That is not reachable: + +- config is **not** passed to `_get_relevant_documents` as a kwarg — a retriever + defined with `**kwargs` receives an empty dict +- `var_child_runnable_config` is **unset** inside a retriever run, so the ambient + config cannot be read either + +Two alternatives were rejected. A per-request attribute on the retriever races, +because it is built once at startup and shared by every request. Rebuilding a +filtered retriever per request throws away the constructed BM25 indexes. + +**Decision: a module-level `ContextVar`**, set around the retrieval call. Under +asyncio each task gets a copy of the context, so one request's narrowing cannot +narrow another's — pinned by a test that runs a narrow and a wide request +concurrently and asserts neither sees the other. + +The data model's diagram is left as the intent; this is how it is carried. + diff --git a/src/retrievers/csv_chroma.py b/src/retrievers/csv_chroma.py index f833885..4b60f8c 100644 --- a/src/retrievers/csv_chroma.py +++ b/src/retrievers/csv_chroma.py @@ -1,6 +1,7 @@ import asyncio import csv -from collections.abc import Coroutine +from collections.abc import Coroutine, Iterable +from contextvars import ContextVar from pathlib import Path from typing import Any, TypedDict @@ -22,6 +23,9 @@ from pydantic import ConfigDict from data_generation.metadata_csv_loader import MetaDataCSVLoader +from util.logging import logging + +logger = logging.getLogger(__name__) def chroma_settings() -> chromadb.config.Settings: @@ -222,6 +226,62 @@ def list_chroma_subdirectories(directory: Path) -> list[str]: ] +selected_collections: ContextVar[list[str] | None] = ContextVar( + "selected_collections", default=None +) +"""The collections the current request may search, or None for all. + +A ContextVar rather than the `RunnableConfig` the data model describes, because +that route does not exist in langchain-core 0.2.14. Measured: config is not +passed to `_get_relevant_documents` as a kwarg, and +`var_child_runnable_config` is unset inside a retriever run. The retriever is +built once at startup and shared, so a per-request attribute would race. + +ContextVars are per-task under asyncio -- each task gets a copy of the context -- +so concurrent requests cannot see each other's selection. Set it around the +retrieval call, never globally. +""" + + +def resolve_collections( + selected: Iterable[str] | None, available: Iterable[str] +) -> list[str]: + """Which collections to search. Every failure widens; none narrows. + + The rules, from specs/009-collection-routing/data-model.md: + + | selection | result | + |----------------------|-----------------------------------------| + | empty or absent | all available | + | every name known | just those | + | any name unknown | all available, and a WARNING naming them| + + The asymmetry is deliberate and is Principle IV. A selection this code does + not recognise means the classifier's prompt and the installed bundle disagree + about what exists -- and searching everything because we are unsure costs + latency, while searching a subset chosen on a misunderstanding costs the + answer. An empty list is therefore the safe value, which is why it is the + default for an omitted field, an older prompt and a parse failure alike. + """ + all_available = list(available) + chosen = list(selected or []) + if not chosen: + return all_available + unknown = sorted(set(chosen) - set(all_available)) + if unknown: + logger.warning( + "Collection selection names %s, which the installed bundle does not " + "have (it has %s). Searching all collections rather than narrowing on " + "a disagreement.", + ", ".join(unknown), + ", ".join(sorted(all_available)), + ) + return all_available + # Preserve the bundle's order rather than the selection's, so the retrieval + # order does not depend on how a model happened to list them. + return [name for name in all_available if name in set(chosen)] + + def _csv_column_names(csv_path: Path) -> list[str]: """Every column in the file, so BM25 metadata matches what was embedded.""" with csv_path.open(newline="", encoding="utf-8") as handle: @@ -386,7 +446,11 @@ def retrieve_documents( self, queries: list[str], run_manager: CallbackManagerForRetrieverRun ) -> list[Document]: subdirectory_docs: list[Document] = [] - for subdirectory, retrievers in self.collection_retrievers.items(): + chosen = resolve_collections( + selected_collections.get(), self.collection_retrievers + ) + for subdirectory in chosen: + retrievers = self.collection_retrievers[subdirectory] bm25_retriever = retrievers["bm25"] vector_retriever = retrievers["vector"] doc_lists: list[list[Document]] = [] @@ -427,7 +491,11 @@ async def aretrieve_documents( run_manager: AsyncCallbackManagerForRetrieverRun, ) -> list[Document]: subdirectory_results: dict[str, list[Coroutine[Any, Any, list[Document]]]] = {} - for subdirectory, retrievers in self.collection_retrievers.items(): + chosen = resolve_collections( + selected_collections.get(), self.collection_retrievers + ) + for subdirectory in chosen: + retrievers = self.collection_retrievers[subdirectory] bm25_retriever = retrievers["bm25"] vector_retriever = retrievers["vector"] subdirectory_results[subdirectory] = [] diff --git a/tests/retrievers/test_collection_selection.py b/tests/retrievers/test_collection_selection.py new file mode 100644 index 0000000..9e2c11e --- /dev/null +++ b/tests/retrievers/test_collection_selection.py @@ -0,0 +1,161 @@ +"""Collection routing: which collections a question searches. + +The rule that matters is the asymmetry. Every failure path widens the search and +none narrows it, because searching everything when unsure costs latency while +searching a subset chosen on a misunderstanding costs the answer (Principle IV). + +The characterization test at the bottom pins today's behaviour -- every +collection searched, always -- so that turning routing on is a visible change +rather than a silent one. +""" + +from typing import Any + +import pytest + +from retrievers.csv_chroma import resolve_collections + +ALL = ["complexes", "disease_variants", "ewas", "reactions", "summations"] + + +class TestResolveCollections: + def test_empty_searches_everything(self) -> None: + """The default, and what every failure degrades to.""" + assert resolve_collections([], ALL) == ALL + + def test_absent_searches_everything(self) -> None: + """An omitted field, an older prompt and a parse failure all arrive here.""" + assert resolve_collections(None, ALL) == ALL + + def test_a_valid_selection_narrows(self) -> None: + assert resolve_collections(["reactions", "ewas"], ALL) == ["ewas", "reactions"] + + def test_the_bundle_order_is_kept_not_the_selection_order(self) -> None: + """Retrieval order must not depend on how a model happened to list them.""" + assert resolve_collections(["summations", "complexes"], ALL) == [ + "complexes", + "summations", + ] + + def test_one_unknown_name_widens_to_everything( + self, caplog: pytest.LogCaptureFixture + ) -> None: + """The case that must not narrow. + + An unrecognised name means the classifier's prompt and the installed + bundle disagree about what exists. Narrowing on that disagreement would + drop a collection the question needed. + """ + with caplog.at_level("WARNING"): + got = resolve_collections(["reactions", "not_a_collection"], ALL) + assert got == ALL, "narrowed on a selection it did not understand" + assert "not_a_collection" in caplog.text + + def test_all_unknown_names_widen_to_everything( + self, caplog: pytest.LogCaptureFixture + ) -> None: + with caplog.at_level("WARNING"): + assert resolve_collections(["nope", "also_nope"], ALL) == ALL + assert "also_nope" in caplog.text + + def test_the_warning_names_what_is_actually_available( + self, caplog: pytest.LogCaptureFixture + ) -> None: + """So an operator can see the disagreement rather than guess at it.""" + with caplog.at_level("WARNING"): + resolve_collections(["typo_reactions"], ALL) + for name in ALL: + assert name in caplog.text + + def test_an_empty_bundle_yields_nothing_rather_than_raising(self) -> None: + """Degenerate, but it must not explode on the retrieval path.""" + assert resolve_collections(["reactions"], []) == [] + assert resolve_collections([], []) == [] + + +class _StubRetriever: + def __init__(self, name: str, seen: list[str]) -> None: + self._name = name + self._seen = seen + + def invoke(self, _query: str, **_kw: Any) -> list[Any]: + self._seen.append(self._name) + return [] + + +def test_characterization_no_selection_searches_every_collection() -> None: + """Today's behaviour, pinned before routing changes it. + + Recorded so that narrowing becomes a visible change. If this starts failing + without `collections` being set, something narrowed retrieval by accident -- + which is the failure that costs answers rather than time. + """ + seen: list[str] = [] + collection_retrievers = {name: _StubRetriever(name, seen) for name in ALL} + + for name in resolve_collections([], list(collection_retrievers)): + collection_retrievers[name].invoke("any query") + + assert sorted(seen) == sorted(ALL) + + +class _FakeHybrid: + """Enough of HybridRetriever to exercise the filtering, without a bundle.""" + + def __init__(self, seen: list[str]) -> None: + self.collection_retrievers = dict.fromkeys(ALL, None) + self._seen = seen + + def search(self) -> None: + from retrievers.csv_chroma import resolve_collections, selected_collections + + for name in resolve_collections( + selected_collections.get(), self.collection_retrievers + ): + self._seen.append(name) + + +def test_the_selection_narrows_retrieval() -> None: + from retrievers.csv_chroma import selected_collections + + seen: list[str] = [] + token = selected_collections.set(["reactions", "ewas"]) + try: + _FakeHybrid(seen).search() + finally: + selected_collections.reset(token) + assert seen == ["ewas", "reactions"] + + +def test_no_selection_still_searches_everything() -> None: + seen: list[str] = [] + _FakeHybrid(seen).search() + assert sorted(seen) == sorted(ALL) + + +def test_one_request_cannot_see_anothers_selection() -> None: + """The reason this is a ContextVar and not an attribute. + + The retriever is built once at startup and shared by every request. An + attribute would race; a ContextVar is copied per asyncio task, so a narrow + selection in one request cannot narrow another's. + """ + import asyncio + + from retrievers.csv_chroma import selected_collections + + async def one(selection: list[str] | None, out: list[str]) -> None: + if selection is not None: + selected_collections.set(selection) + await asyncio.sleep(0) # force interleaving + _FakeHybrid(out).search() + + narrow: list[str] = [] + wide: list[str] = [] + + async def both() -> None: + await asyncio.gather(one(["reactions"], narrow), one(None, wide)) + + asyncio.run(both()) + assert narrow == ["reactions"], "the narrow request did not get its selection" + assert sorted(wide) == sorted(ALL), "a concurrent request leaked its narrowing"