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
7 changes: 6 additions & 1 deletion src/agent/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
28 changes: 20 additions & 8 deletions src/agent/profiles/base.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import asyncio
import logging
from typing import Annotated, Literal, TypedDict

from langchain_core.embeddings import Embeddings
Expand Down Expand Up @@ -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)
)
21 changes: 20 additions & 1 deletion src/agent/profiles/react_to_me.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
),
Expand All @@ -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:
Expand Down
30 changes: 27 additions & 3 deletions src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
6 changes: 6 additions & 0 deletions src/retrievers/userguide/prompt.py
Original file line number Diff line number Diff line change
@@ -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.
Expand Down Expand Up @@ -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}"),
]
)
47 changes: 47 additions & 0 deletions tests/agent/test_history_and_resources.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
14 changes: 10 additions & 4 deletions tests/agent/test_intent_classifier_sources.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
Loading