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
12 changes: 0 additions & 12 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -176,15 +176,3 @@ disable_error_code = ["typeddict-item", "no-any-return"]
module = ["agent.profiles.react_to_me"]
disable_error_code = ["typeddict-item", "override"]

[[tool.mypy.overrides]]
# TODO(phase-2): `EmbeddingEnvironment.get_dir()` returns `Path | None` but the
# parameter is typed `Path`, so an uninstalled bundle passes None straight through.
# Fixed properly by moving the lookup out of the default argument (see ruff B008).
module = [
"retrievers.reactome.rag",
"retrievers.uniprot.rag",
"retrievers.plantreactome.rag",
"retrievers.userguide.rag",
]
disable_error_code = ["assignment"]

9 changes: 7 additions & 2 deletions src/agent/profiles/cross_database.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
)
from retrievers.reactome.rag import create_reactome_rag
from retrievers.uniprot.rag import create_uniprot_rag
from util.embedding_environment import EmbeddingEnvironment


class CrossDatabaseState(BaseState):
Expand All @@ -43,8 +44,12 @@ def __init__(
super().__init__(llm, embedding)

# Create runnables (tasks & tools)
self.reactome_rag: Runnable = create_reactome_rag(llm, embedding)
self.uniprot_rag: Runnable = create_uniprot_rag(llm, embedding)
self.reactome_rag: Runnable = create_reactome_rag(
llm, embedding, EmbeddingEnvironment.require_dir("reactome")
)
self.uniprot_rag: Runnable = create_uniprot_rag(
llm, embedding, EmbeddingEnvironment.require_dir("uniprot")
)

self.completeness_checker = create_completeness_grader(llm)
self.write_reactome_query = create_reactome_rewriter_w_uniprot(llm)
Expand Down
6 changes: 5 additions & 1 deletion src/agent/profiles/plantreactome.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from agent.profiles.base import BaseGraphBuilder, BaseState
from agent.tasks.unsafe_question import create_unsafe_answer_generator
from retrievers.plantreactome.rag import create_plantreactome_rag
from util.embedding_environment import EmbeddingEnvironment


class PlantReactomeState(BaseState):
Expand All @@ -28,7 +29,10 @@ def __init__(
llm, streaming=True
)
self.plantreactome_rag: Runnable = create_plantreactome_rag(
llm, embedding, streaming=True
llm,
embedding,
EmbeddingEnvironment.require_dir("plantreactome"),
streaming=True,
)

# Create graph
Expand Down
7 changes: 6 additions & 1 deletion src/agent/profiles/react_to_me.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,12 @@ def __init__(
)

self.rags: dict[SourceName, Runnable] = {
"reactome": create_reactome_rag(llm, embedding, streaming=True),
"reactome": create_reactome_rag(
llm,
embedding,
EmbeddingEnvironment.require_dir("reactome"),
streaming=True,
),
}
self._available_sources: frozenset[SourceName] = frozenset({"reactome"})
self._register_userguide_rag(llm, embedding)
Expand Down
22 changes: 19 additions & 3 deletions src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,12 @@
# retriever returns. Changing it trades recall against the model's difficulty
# attending to the middle of a long context; the right number should come from an
# answer-quality evaluation rather than from taste.
MAX_DOCUMENTS_PER_COLLECTION = RESULTS_PER_RETRIEVER
#
# It is the DEFAULT, not the budget: a caller passes its own to
# HybridRetriever.from_subdirectory. That is what makes the evaluation above
# possible -- sweeping this value used to mean editing the module and restarting,
# so two budgets could not be compared within one process.
DEFAULT_MAX_DOCUMENTS_PER_COLLECTION = RESULTS_PER_RETRIEVER


def dedupe_by_entity(docs: list[Document], limit: int) -> list[Document]:
Expand Down Expand Up @@ -194,12 +199,15 @@ def create_bm25_chroma_ensemble_retriever(
llm: BaseChatModel,
embedding: Embeddings,
embeddings_directory: Path,
*,
max_documents_per_collection: int = DEFAULT_MAX_DOCUMENTS_PER_COLLECTION,
) -> "HybridRetriever":
return HybridRetriever.from_subdirectory(
llm,
embedding,
embeddings_directory,
include_original=True,
max_documents_per_collection=max_documents_per_collection,
)


Expand Down Expand Up @@ -228,6 +236,8 @@ class HybridRetriever(BaseRetriever):
query_expander: Runnable[dict[str, str], list[str]]
include_original: bool = False
collection_retrievers: dict[str, RetrieverDict]
# How many fused documents this instance contributes per collection.
max_documents_per_collection: int = DEFAULT_MAX_DOCUMENTS_PER_COLLECTION

# BM25Retriever and the Chroma retriever are not pydantic models.
model_config = ConfigDict(arbitrary_types_allowed=True)
Expand All @@ -240,6 +250,7 @@ def from_subdirectory(
embeddings_directory: Path,
*,
include_original: bool = False,
max_documents_per_collection: int = DEFAULT_MAX_DOCUMENTS_PER_COLLECTION,
) -> "HybridRetriever":
_retrievers: dict[str, RetrieverDict] = {}
for subdirectory in list_chroma_subdirectories(embeddings_directory):
Expand Down Expand Up @@ -279,6 +290,7 @@ def from_subdirectory(
query_expander=multi_query_prompt | llm | LineListOutputParser(),
include_original=include_original,
collection_retrievers=_retrievers,
max_documents_per_collection=max_documents_per_collection,
)

def _get_relevant_documents(
Expand Down Expand Up @@ -348,7 +360,9 @@ def retrieve_documents(
doc_lists.append(dedupe_by_entity(bm25_docs, RESULTS_PER_RETRIEVER))
doc_lists.append(dedupe_by_entity(vector_docs, RESULTS_PER_RETRIEVER))
subdirectory_docs.extend(
self.weighted_reciprocal_rank(doc_lists)[:MAX_DOCUMENTS_PER_COLLECTION]
self.weighted_reciprocal_rank(doc_lists)[
: self.max_documents_per_collection
]
)
return subdirectory_docs

Expand Down Expand Up @@ -393,6 +407,8 @@ async def aretrieve_documents(
for docs in await asyncio.gather(*subdir_results)
]
subdirectory_docs.extend(
self.weighted_reciprocal_rank(doc_lists)[:MAX_DOCUMENTS_PER_COLLECTION]
self.weighted_reciprocal_rank(doc_lists)[
: self.max_documents_per_collection
]
)
return subdirectory_docs
5 changes: 1 addition & 4 deletions src/retrievers/plantreactome/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,15 +7,12 @@
from retrievers.csv_chroma import create_bm25_chroma_ensemble_retriever
from retrievers.plantreactome.prompt import plantreactome_qa_prompt
from retrievers.rag_chain import create_rag_chain
from util.embedding_environment import EmbeddingEnvironment


def create_plantreactome_rag(
llm: BaseChatModel,
embedding: Embeddings,
# TODO(phase-2): resolved at import time, so importing this module requires an
# installed embeddings bundle. Blocks unit-testing; fix with the agent-API refactor.
embeddings_directory: Path = EmbeddingEnvironment.get_dir("plantreactome"), # noqa: B008
embeddings_directory: Path,
*,
streaming: bool = False,
) -> Runnable:
Expand Down
5 changes: 1 addition & 4 deletions src/retrievers/reactome/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,15 +7,12 @@
from retrievers.csv_chroma import create_bm25_chroma_ensemble_retriever
from retrievers.rag_chain import create_rag_chain
from retrievers.reactome.prompt import reactome_qa_prompt
from util.embedding_environment import EmbeddingEnvironment


def create_reactome_rag(
llm: BaseChatModel,
embedding: Embeddings,
# TODO(phase-2): resolved at import time, so importing this module requires an
# installed embeddings bundle. Blocks unit-testing; fix with the agent-API refactor.
embeddings_directory: Path = EmbeddingEnvironment.get_dir("reactome"), # noqa: B008
embeddings_directory: Path,
*,
streaming: bool = False,
) -> Runnable:
Expand Down
5 changes: 1 addition & 4 deletions src/retrievers/uniprot/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,15 +7,12 @@
from retrievers.csv_chroma import create_bm25_chroma_ensemble_retriever
from retrievers.rag_chain import create_rag_chain
from retrievers.uniprot.prompt import uniprot_qa_prompt
from util.embedding_environment import EmbeddingEnvironment


def create_uniprot_rag(
llm: BaseChatModel,
embedding: Embeddings,
# TODO(phase-2): resolved at import time, so importing this module requires an
# installed embeddings bundle. Blocks unit-testing; fix with the agent-API refactor.
embeddings_directory: Path = EmbeddingEnvironment.get_dir("uniprot"), # noqa: B008
embeddings_directory: Path,
*,
streaming: bool = False,
) -> Runnable:
Expand Down
5 changes: 1 addition & 4 deletions src/retrievers/userguide/rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,15 +7,12 @@
from retrievers.rag_chain import create_rag_chain
from retrievers.userguide.prompt import userguide_qa_prompt
from retrievers.userguide.retriever import create_userguide_retriever
from util.embedding_environment import EmbeddingEnvironment


def create_userguide_rag(
llm: BaseChatModel,
embedding: Embeddings,
# TODO(phase-2): resolved at import time, so importing this module requires an
# installed embeddings bundle. Blocks unit-testing; fix with the agent-API refactor.
embeddings_directory: Path = EmbeddingEnvironment.get_dir("userguide"), # noqa: B008
embeddings_directory: Path,
*,
streaming: bool = False,
) -> Runnable:
Expand Down
31 changes: 31 additions & 0 deletions src/util/embedding_environment.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,37 @@ def get_dir(cls, key: str) -> Path | None:
return EM_ARCHIVE / cls._get().embeddings[key]
return None

@classmethod
def require_dir(cls, key: str) -> Path:
"""`get_dir`, but for bundles the caller cannot run without.

`get_dir` returns None for an unknown key, and returns a path for a
known one without checking that the path exists. Both failures used to
travel: None was passed into a parameter typed `Path` (four mypy
baseline entries suppressed the error), and a stale `embeddings/current`
pointing at a deleted bundle handed Chroma a missing directory, which
Chroma creates -- so the chatbot answered every question from an empty
collection instead of refusing to start.

Raising FileNotFoundError rather than SystemExit: it is the accurate
exception, nothing catches it for a required bundle so the process still
stops, and the optional userguide path in react_to_me.py already catches
exactly this to degrade deliberately.
"""
directory = cls.get_dir(key)
if directory is None:
available = ", ".join(sorted(cls.get_dict())) or "none"
raise FileNotFoundError(
f"No embeddings bundle installed for {key!r} (installed: {available}). "
f"Run ./bin/embeddings_manager install <embedding-id>."
)
if not directory.is_dir():
raise FileNotFoundError(
f"{key!r} points at {directory}, which does not exist. "
f"{EM_CURRENT} is stale; re-run ./bin/embeddings_manager install."
)
return directory

@classmethod
def get_model(cls, key: str) -> str:
return str(cls._get().embeddings[key].parent.parent)
Expand Down
93 changes: 91 additions & 2 deletions tests/retrievers/test_hybrid_retriever.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@

from retrievers import csv_chroma # noqa: E402
from retrievers.csv_chroma import ( # noqa: E402
MAX_DOCUMENTS_PER_COLLECTION,
DEFAULT_MAX_DOCUMENTS_PER_COLLECTION,
HybridRetriever,
RetrieverDict,
dedupe_by_entity,
Expand Down Expand Up @@ -201,7 +201,7 @@ def test_fused_results_are_capped_per_collection() -> None:
many = [[_doc(f"doc {i}") for i in range(50)]]
assert len(_fuse(many)) == 50, "the fusion itself still returns everything"
assert (
MAX_DOCUMENTS_PER_COLLECTION < 50
DEFAULT_MAX_DOCUMENTS_PER_COLLECTION < 50
), "the cap, applied by the caller, is what bounds the prompt"


Expand Down Expand Up @@ -232,3 +232,92 @@ def test_retriever_dict_accepts_any_base_retriever() -> None:
"""
hints = get_type_hints(RetrieverDict)
assert hints["vector"] is BaseRetriever


class _StubVectorRetriever(BaseRetriever):
"""Returns a fixed list, so a budget test needs no Chroma store and no LLM."""

docs: list[Document]

def _get_relevant_documents(self, query: str, **kwargs: Any) -> list[Document]:
return self.docs


def _retriever_with_budget(budget: int | None) -> HybridRetriever:
from langchain_community.retrievers import BM25Retriever
from langchain_core.runnables import RunnableLambda

corpus = [_doc(f"doc {i}") for i in range(30)]
bm25 = BM25Retriever.from_documents(corpus) # default whitespace tokenizer, no nltk
bm25.k = 30
collection: RetrieverDict = {
"bm25": bm25,
"vector": _StubVectorRetriever(docs=corpus),
}
kwargs: dict[str, Any] = (
{} if budget is None else {"max_documents_per_collection": budget}
)
return HybridRetriever(
# No expansion: one query in, one query out, so the only thing varying
# between the two instances below is the budget.
query_expander=RunnableLambda(lambda inputs: [inputs["question"]]),
include_original=False,
collection_retrievers={"reactions": collection},
**kwargs,
)


def test_two_budgets_can_coexist_in_one_process() -> None:
"""The point of Stage 3: the budget is an argument, not a module constant.

While it was a module constant, comparing two budgets meant editing
csv_chroma.py and restarting -- so the answer-quality evaluation that is
supposed to settle the number could not be run. Asserting both instances in
the same test is the whole claim: not that the value can be changed, but that
two values can be live at once.
"""
small = _retriever_with_budget(3)
large = _retriever_with_budget(7)

assert len(small.invoke("anything")) == 3
assert len(large.invoke("anything")) == 7
# and the smaller is a prefix of the larger: same ranking, less of it
assert [d.page_content for d in small.invoke("anything")] == [
d.page_content for d in large.invoke("anything")
][:3]


def test_omitting_the_budget_behaves_as_the_old_constant_did() -> None:
"""Making it an argument must not quietly change production's context size.

Production passes no budget, so the default is what actually ships.
"""
assert (
len(_retriever_with_budget(None).invoke("x"))
== DEFAULT_MAX_DOCUMENTS_PER_COLLECTION
)


@pytest.mark.parametrize(
"module",
["reactome", "uniprot", "plantreactome", "userguide"],
)
def test_rag_factories_take_the_bundle_rather_than_resolving_it(module: str) -> None:
"""FR-006: no default argument may call EmbeddingEnvironment at import time.

The default used to be `EmbeddingEnvironment.get_dir(...)`, evaluated once
when the module was imported. That made importing a rag module require an
installed bundle, and it fed `Path | None` into a parameter typed `Path` --
suppressed by four mypy baseline entries, now deleted.
"""
import importlib
import inspect

mod = importlib.import_module(f"retrievers.{module}.rag")
fn = getattr(mod, f"create_{module}_rag")
parameter = inspect.signature(fn).parameters["embeddings_directory"]

assert (
parameter.default is inspect.Parameter.empty
), "a default here is evaluated at import time; the caller must pass the bundle"
assert parameter.annotation is Path, "and it is a Path, never Path | None"
Loading
Loading