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
25 changes: 25 additions & 0 deletions specs/009-collection-routing/research.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

74 changes: 71 additions & 3 deletions src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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]] = []
Expand Down Expand Up @@ -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] = []
Expand Down
161 changes: 161 additions & 0 deletions tests/retrievers/test_collection_selection.py
Original file line number Diff line number Diff line change
@@ -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"
Loading