diff --git a/src/agent/graph.py b/src/agent/graph.py index 2203afd..28fcaa1 100644 --- a/src/agent/graph.py +++ b/src/agent/graph.py @@ -237,7 +237,12 @@ def _citation_for(document: "Document") -> "AnswerEvent | None": return AnswerEvent( kind="citation", st_id=str(stable_id), - display_name=str(metadata.get("display_name") or ""), + # Disease-variant documents have no display_name; their name is + # the variant ("ABCA1 W590S [plasma membrane]"). Every one of + # their citations went out unlabelled (review, area 3). + display_name=str( + metadata.get("display_name") or metadata.get("variant") or "" + ), ) source = str(metadata.get("source") or "") # Only a real web URL. `source` is a generic LangChain field, and the CSV diff --git a/src/agent/profiles/base.py b/src/agent/profiles/base.py index ff13fed..0e09992 100644 --- a/src/agent/profiles/base.py +++ b/src/agent/profiles/base.py @@ -1,4 +1,5 @@ import asyncio +import logging from typing import Annotated, Literal, TypedDict from langchain_core.embeddings import Embeddings @@ -85,14 +86,25 @@ async def postprocess(self, state: BaseState, config: RunnableConfig) -> BaseSta config["configurable"].get("enable_postprocess") and state["safety"] == "true" ): - result: SearchState = await self.search_workflow.ainvoke( - SearchState( - input=state["rephrased_input"], - generation=state["answer"], - ), - config=RunnableConfig(callbacks=config["callbacks"]), - ) - search_results = result["search_results"] + try: + result: SearchState = await self.search_workflow.ainvoke( + SearchState( + input=state["rephrased_input"], + generation=state["answer"], + ), + config=RunnableConfig(callbacks=config["callbacks"]), + ) + search_results = result["search_results"] + except Exception: + # Optional, and after the answer has already streamed and been + # saved. A failed grader or search used to fail the whole turn: + # the reader saw "something went wrong" under a complete + # answer, and a retry re-sent a question already in the + # history (review, area 3). + logging.getLogger(__name__).warning( + "post-answer web search failed; answering without it", + exc_info=True, + ) return BaseState( additional_content=AdditionalContent(search_results=search_results) ) diff --git a/src/agent/profiles/react_to_me.py b/src/agent/profiles/react_to_me.py index a38ca9d..4e8732b 100644 --- a/src/agent/profiles/react_to_me.py +++ b/src/agent/profiles/react_to_me.py @@ -6,6 +6,7 @@ from langchain_core.language_models.chat_models import BaseChatModel from langchain_core.messages import AIMessage, HumanMessage from langchain_core.runnables import Runnable, RunnableConfig +from langchain_openai.chat_models.base import OpenAIRefusalError from langgraph.graph.state import StateGraph from agent.history import recent @@ -164,7 +165,7 @@ async def preprocess( safety_check: SafetyCheck intent: QueryIntent safety_check, intent = await asyncio.gather( - self.safety_checker.ainvoke({"rephrased_input": rephrased_input}, config), + self._check_safety(rephrased_input, config), self.intent_classifier.ainvoke( {"rephrased_input": rephrased_input}, config ), @@ -186,6 +187,24 @@ async def preprocess( collections=intent.collections, ) + async def _check_safety( + self, rephrased_input: str, config: RunnableConfig + ) -> SafetyCheck: + """The safety verdict -- and a refusal to give one is a "no". + + With json_schema output, a model that refuses to classify a question + raises OpenAIRefusalError rather than answering "false"; the turn then + failed outright, and a harmful question got an error instead of the + polite refusal (review, area 3). + """ + try: + result: SafetyCheck = await self.safety_checker.ainvoke( + {"rephrased_input": rephrased_input}, config + ) + except OpenAIRefusalError: + return SafetyCheck(safety="false", reason_unsafe="the request was declined") + return result + async def _answer_from_live_services( self, state: ReactToMeState, config: RunnableConfig ) -> ReactToMeState: diff --git a/src/retrievers/csv_chroma.py b/src/retrievers/csv_chroma.py index 176a8d1..b2bd8b4 100644 --- a/src/retrievers/csv_chroma.py +++ b/src/retrievers/csv_chroma.py @@ -23,12 +23,21 @@ from nltk.tokenize import word_tokenize from pydantic import ConfigDict +from data_generation.disease_variant import ( + CONTENT_COLUMNS as DISEASE_VARIANT_CONTENT_COLUMNS, +) from data_generation.metadata_csv_loader import MetaDataCSVLoader from util.logging import logging logger = logging.getLogger(__name__) +#: Collections whose vectors were built from a subset of the CSV's columns; +#: the keyword index must see the same text. Others embed every column. +EMBEDDED_CONTENT_COLUMNS: dict[str, list[str]] = { + "disease_variants": DISEASE_VARIANT_CONTENT_COLUMNS, +} + #: Distinct query tokens BM25 scores. rank_bm25 scans every document once per #: query token, repeats included: ~45 ms each over Release 97, so a 2,000- #: character question cost ~16 s of CPU and an 8,000-character chat message @@ -433,10 +442,25 @@ def from_subdirectory( # fusion. # # Every column is promoted, from the header, so this cannot drift - # from whatever the CSV holds. Content is unchanged: the loader only - # trims content when `content_columns` is given, which it is not. + # from whatever the CSV holds. + # + # Content follows what was embedded. The disease-variant store was + # built from 12 of its 21 columns; indexing all 21 here meant the + # two retrievers never returned the same text for a variant, so RRF + # -- which keys on the text -- never fused them: 0 of 300 matched, + # and "Which ABCA1 variants are there?" filled its 10 slots with 5 + # variants, each twice (review, area 3). loader = MetaDataCSVLoader( - file_path=str(csv_path), metadata_columns=_csv_column_names(csv_path) + file_path=str(csv_path), + metadata_columns=_csv_column_names(csv_path), + # Only columns the file has: none at all would leave every + # document empty, and BM25 divides by their mean length. + content_columns=[ + column + for column in EMBEDDED_CONTENT_COLUMNS.get(subdirectory, []) + if column in _csv_column_names(csv_path) + ] + or None, ) data = loader.load() bm25_retriever = BoundedBM25Retriever.from_documents( diff --git a/src/retrievers/userguide/prompt.py b/src/retrievers/userguide/prompt.py index a4d5e20..e57ff82 100644 --- a/src/retrievers/userguide/prompt.py +++ b/src/retrievers/userguide/prompt.py @@ -1,5 +1,7 @@ from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder +from agent.tasks.language_instruction import LANGUAGE_INSTRUCTION + userguide_system_prompt = """ You are a helpful guide to the **Reactome website** and its tools. Your primary responsibility is to answer questions about **how to use Reactome** — the Pathway Browser, search, analysis tools, Details Panel, and related features — using only the user guide excerpts provided in the context. @@ -38,10 +40,14 @@ - The Sources list is complete and de-duplicated. """ +# The same language instruction, in the same place, as the Reactome prompt. +# The user guide's answers ignored the reader's language: the detected +# language was passed and no prompt variable took it (review, area 3). userguide_qa_prompt = ChatPromptTemplate.from_messages( [ ("system", userguide_system_prompt), MessagesPlaceholder(variable_name="chat_history"), + ("system", LANGUAGE_INSTRUCTION), ("user", "Context:\n{context}\n\nQuestion: {input}"), ] ) diff --git a/tests/agent/test_history_and_resources.py b/tests/agent/test_history_and_resources.py index 9990fad..10458d7 100644 --- a/tests/agent/test_history_and_resources.py +++ b/tests/agent/test_history_and_resources.py @@ -89,3 +89,50 @@ async def many() -> None: asyncio.run(many()) assert calls == [1] + + +# --- answer quality, review area 3 ------------------------------------------------ + + +def test_a_disease_variant_citation_is_labelled_with_its_variant() -> None: + # Disease-variant documents carry no display_name; every one of their + # citations went out with an empty label. + from agent.graph import _citation_for + + doc = Document( + page_content="...", + metadata={"st_id": "R-HSA-5682201", "variant": "PTEN R130G [cytosol]"}, + ) + citation = _citation_for(doc) + assert citation is not None + assert citation.display_name == "PTEN R130G [cytosol]" + + +def test_a_failed_post_answer_step_does_not_fail_the_turn() -> None: + from agent.profiles.base import BaseGraphBuilder + + class _Failing: + async def ainvoke(self, *_a: Any, **_k: Any) -> Any: + raise RuntimeError("429 from the grader") + + builder = BaseGraphBuilder.__new__(BaseGraphBuilder) + builder.search_workflow = _Failing() # type: ignore[assignment] + state: Any = {"safety": "true", "rephrased_input": "q", "answer": "a"} + config: Any = {"configurable": {"enable_postprocess": True}, "callbacks": None} + result = asyncio.run(builder.postprocess(state, config)) + assert result["additional_content"]["search_results"] == [] + + +def test_a_refused_safety_check_is_an_unsafe_verdict() -> None: + from langchain_openai.chat_models.base import OpenAIRefusalError + + from agent.profiles.react_to_me import ReactToMeGraphBuilder + + class _Refusing: + async def ainvoke(self, *_a: Any, **_k: Any) -> Any: + raise OpenAIRefusalError("I can't help with that.") + + builder = ReactToMeGraphBuilder.__new__(ReactToMeGraphBuilder) + builder.safety_checker = _Refusing() # type: ignore[assignment] + verdict = asyncio.run(builder._check_safety("something", {})) + assert verdict.safety == "false" diff --git a/tests/agent/test_intent_classifier_sources.py b/tests/agent/test_intent_classifier_sources.py index 40f4973..b3fd119 100644 --- a/tests/agent/test_intent_classifier_sources.py +++ b/tests/agent/test_intent_classifier_sources.py @@ -18,12 +18,18 @@ ONE: frozenset[SourceName] = frozenset({"reactome"}) -def test_without_the_mcp_the_prompt_is_byte_for_byte_what_it_was() -> None: - """FR-007: with the new step disabled, behaviour is exactly as before. +def test_without_the_mcp_the_prompt_never_mentions_live() -> None: + """Without live services the classifier must not be told they exist. - Not "similar" -- identical. A deployment with no MCP server must classify - the same questions the same way, and pay the same tokens doing it. + This replaced a test that compared the prompt with a constant built by + the same function at import, so it could never fail -- and its claim, + "byte-for-byte what it was", had stopped being true when the routing + rules were extended (review, area 3). """ + import re + + message = build_classifier_message(TWO) + assert re.search(r"\blive\b", message, re.IGNORECASE) is None assert build_classifier_message(TWO) == intent_classifier_message