diff --git a/src/agent/graph.py b/src/agent/graph.py index bf4cfb5..6d0aa94 100644 --- a/src/agent/graph.py +++ b/src/agent/graph.py @@ -1,7 +1,7 @@ import asyncio import os import re -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Mapping from dataclasses import dataclass from typing import Any, Literal, cast @@ -198,6 +198,22 @@ def resolve_llm_model(llm_config: "LLMConfig | None") -> tuple[str, str, str | N MAX_CITATIONS = 12 +def _is_sub_retriever(event: Mapping[str, Any]) -> bool: + """True for the per-collection searches inside the hybrid retriever. + + Identified by the tags this repository sets itself in + `HybridRetriever.retrieve_documents` -- `{collection}-bm25-{i}` and + `{collection}-vector-{i}` -- rather than by retriever class name. The name + is not usable: the user guide's own chain-level retriever is a + `VectorStoreRetriever` too, and filtering on that would silently drop every + user-guide citation. + """ + return any( + "-bm25-" in str(tag) or "-vector-" in str(tag) + for tag in (event.get("tags") or []) + ) + + def _citation_for(document: "Document") -> "AnswerEvent | None": """One source, from whichever identifier the collection actually carries. @@ -451,7 +467,38 @@ async def astream_answer( retrieval_done = True continue + if kind == "on_retriever_end" and _is_sub_retriever(event): + # Skip. The hybrid retriever runs one BM25 and one vector search + # per query variant per collection, and each fires its own + # on_retriever_end -- 51 events for one question. Taking + # citations from those meant taking them from whichever fired + # first, which is always the first collection's BM25. + # + # Measured: 97 of 97 cited documents came from reactions.csv, + # including for "which diseases involve variants of the PTEN + # gene", where the disease_variants documents were retrieved and + # could never be cited. The contract calls citations "the most + # relevant few"; they were one sub-retriever's few. + continue + if kind == "on_retriever_end": + # Skip the sub-retrievers. The Reactome ensemble fires 51 of + # these -- one per collection per query variant per retriever -- + # and only the last is the fused result. Citations were taken + # from the first events to arrive, so every one came from + # whichever collection happened to be iterated first: measured, + # 97 of 97 cited documents were from reactions.csv, including for + # "which diseases involve variants of the PTEN gene", whose + # disease_variants documents were retrieved and never cited. + # + # Keyed on the child tags this repo creates itself + # ({collection}-bm25-{i}, {collection}-vector-{i}) rather than on + # the retriever class: the user guide's top-level retriever is a + # VectorStoreRetriever, so filtering by class would silently drop + # its citations instead. + tags = event.get("tags") or [] + if any("-bm25-" in tag or "-vector-" in tag for tag in tags): + continue retrieval_done = True for document in event["data"].get("output") or []: if len(seen_citations) >= MAX_CITATIONS: diff --git a/src/retrievers/csv_chroma.py b/src/retrievers/csv_chroma.py index 4b60f8c..8966ce5 100644 --- a/src/retrievers/csv_chroma.py +++ b/src/retrievers/csv_chroma.py @@ -445,7 +445,7 @@ def weighted_reciprocal_rank( def retrieve_documents( self, queries: list[str], run_manager: CallbackManagerForRetrieverRun ) -> list[Document]: - subdirectory_docs: list[Document] = [] + collection_lists: list[list[Document]] = [] chosen = resolve_collections( selected_collections.get(), self.collection_retrievers ) @@ -478,12 +478,28 @@ def retrieve_documents( # only across query variants. See issue #170. 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( + collection_lists.append( self.weighted_reciprocal_rank(doc_lists)[ : self.max_documents_per_collection ] ) - return subdirectory_docs + # Fused across collections, not concatenated. This is issue #170 one + # level up: there, bm25 and vector results were joined end to end so the + # two were never scored against each other. Here the per-collection lists + # were joined the same way, so the output was grouped by collection + # rather than ranked -- ten reactions, then ten complexes, and so on. + # + # That made the citation cap arbitrary. Measured before this change: 97 + # of 97 cited documents came from reactions.csv, including for "which + # diseases involve variants of the PTEN gene", where the disease_variants + # documents were retrieved but sat at positions 41-50 and could never be + # cited. The contract tells the website to treat citations as "the most + # relevant few", which was not true of a list ordered by iteration. + # + # The count is unchanged, so this reorders rather than re-scopes. + return self.weighted_reciprocal_rank(collection_lists)[ + : self.max_documents_per_collection * len(collection_lists) + ] async def aretrieve_documents( self, @@ -521,7 +537,7 @@ async def aretrieve_documents( subdirectory_results[subdirectory].extend( (bm25_results, vector_results) ) - subdirectory_docs: list[Document] = [] + collection_lists: list[list[Document]] = [] for subdir_results in subdirectory_results.values(): # Separate lists, de-duplicated per entity, matching the synchronous # path above. See issues #169 and #170. @@ -529,9 +545,13 @@ async def aretrieve_documents( dedupe_by_entity(docs, RESULTS_PER_RETRIEVER) for docs in await asyncio.gather(*subdir_results) ] - subdirectory_docs.extend( + collection_lists.append( self.weighted_reciprocal_rank(doc_lists)[ : self.max_documents_per_collection ] ) - return subdirectory_docs + # Fused across collections, matching the synchronous path. The two must + # agree document for document, which is what tests/retrievers pins. + return self.weighted_reciprocal_rank(collection_lists)[ + : self.max_documents_per_collection * len(collection_lists) + ]