diff --git a/src/agent/graph.py b/src/agent/graph.py index 8209df9..5d24c9c 100644 --- a/src/agent/graph.py +++ b/src/agent/graph.py @@ -22,6 +22,7 @@ from agent.profile_names import ProfileName from agent.profiles import create_profile_graphs from agent.profiles.base import InputState, OutputState +from reactome_mcp.session import get_mcp_tools from util.config_yml.models import LLMConfig from util.embedding_environment import EmbeddingEnvironment from util.logging import logging @@ -492,8 +493,14 @@ async def astream_answer( # retriever, so nothing else is streaming at the answer node and # there is nothing to tell apart. output = event["data"].get("output") - if isinstance(output, dict) and "live" in ( - output.get("active_sources") or [] + if ( + isinstance(output, dict) + and "live" in (output.get("active_sources") or []) + # Only if live tools will really answer. Without them the + # same node falls back to retrieval, whose query expander + # streams first -- and the search page showed its + # alternate questions as the answer (review, area 3). + and await get_mcp_tools() is not None ): retrieval_done = True continue diff --git a/src/reactome_mcp/answer.py b/src/reactome_mcp/answer.py index 9f471f4..6c1d053 100644 --- a/src/reactome_mcp/answer.py +++ b/src/reactome_mcp/answer.py @@ -27,6 +27,8 @@ logger = logging.getLogger(__name__) +#: Tool calls run per round; the rest are answered without running. +MAX_CALLS_PER_ROUND = 6 MAX_TOOL_ROUNDS = 2 @@ -118,7 +120,19 @@ async def answer_from_live_services( if not calls: break - for call in calls: + for index, call in enumerate(calls): + if index >= MAX_CALLS_PER_ROUND: + # Answered, not dropped: every tool call needs a reply for the + # next request to be valid. One question had 150 calls in a + # round, 300 upstream requests, run one after another with no + # overall deadline (review, area 3). + messages.append( + ToolMessage( + content="Not run: too many lookups at once. Use the results above.", + tool_call_id=call["id"], + ) + ) + continue tool = by_name.get(call["name"]) if tool is None: # The model invented a tool name. Tell it so, rather than @@ -138,7 +152,9 @@ async def answer_from_live_services( # A failed lookup is information, not a crash. The model needs # to say it could not find out, rather than invent. logger.warning("live tool %s failed: %s", call["name"], exc) - result = f"This lookup failed: {exc}" + # The type, not the message: it carried the MCP's internal URL + # ("... for url 'http://reactome_mcp:4320/mcp'") to OpenAI. + result = f"This lookup failed ({type(exc).__name__})." # Recorded before it is stringified into the model's context, # which is the last point at which it is still a fact rather # than a paraphrase. diff --git a/src/reactome_mcp/http_client.py b/src/reactome_mcp/http_client.py index 7d15c1e..ae7702c 100644 --- a/src/reactome_mcp/http_client.py +++ b/src/reactome_mcp/http_client.py @@ -21,6 +21,8 @@ text that is perfectly valid. """ +import asyncio +import itertools import json import logging from typing import Any @@ -43,6 +45,13 @@ def __init__(self, base_url: str, timeout: float = 30.0) -> None: self.timeout = timeout self._session_id: str | None = None self._client = httpx.AsyncClient(timeout=timeout) + # A fresh id for every request. Every call used to send id 2 on the + # one session the whole process shares, and the server matches + # replies by id: two users' overlapping lookups got each other's + # results -- one of them another reader's identifiers and analysis + # token -- and the other call hung (review, area 3). + self._ids = itertools.count(2) + self._init_lock = asyncio.Lock() def _headers(self) -> dict[str, str]: headers = {"Content-Type": "application/json", "Accept": ACCEPT} @@ -51,7 +60,9 @@ def _headers(self) -> dict[str, str]: return headers @staticmethod - def _parse(response: httpx.Response) -> dict[str, Any]: + def _parse( + response: httpx.Response, expected_id: int | None = None + ) -> dict[str, Any]: """Read one JSON-RPC message out of a JSON or SSE body.""" body = response.text if "text/event-stream" in response.headers.get("content-type", ""): @@ -76,6 +87,12 @@ def _parse(response: httpx.Response) -> dict[str, Any]: f"MCP returned a {type(message).__name__}, expected an object" ) + if expected_id is not None and message.get("id") != expected_id: + # Never someone else's answer: refuse rather than hand it on. + raise MCPToolError( + f"MCP replied to request {message.get('id')!r}, not {expected_id}" + ) + if "error" in message: error = message["error"] raise MCPToolError(f"MCP error {error.get('code')}: {error.get('message')}") @@ -128,23 +145,50 @@ async def initialize(self, client_name: str = "reactome-chatbot") -> dict[str, A ) return result + async def _ensure_session(self) -> None: + async with self._init_lock: + if self._session_id is None: + await self.initialize() + + @staticmethod + def _session_lost(response: httpx.Response) -> bool: + """The server no longer knows our session -- it restarted, or evicted + it. Answered 400 "No valid session" by reactome-mcp, 404 by the spec.""" + if response.status_code == 404: + return True + return response.status_code == 400 and "session" in response.text.lower() + async def call( self, method: str, params: dict[str, Any] | None = None ) -> dict[str, Any]: - response = await self._client.post( - self.endpoint, - headers=self._headers(), - json={"jsonrpc": "2.0", "id": 2, "method": method, "params": params or {}}, - ) - response.raise_for_status() - return self._parse(response) + await self._ensure_session() + for attempt in (1, 2): + request_id = next(self._ids) + response = await self._client.post( + self.endpoint, + headers=self._headers(), + json={ + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params or {}, + }, + ) + if attempt == 1 and self._session_lost(response): + # Re-initialize once. Without this, one MCP restart turned + # live lookups off until the chatbot restarted (review, 3). + logger.info("MCP session lost; initializing a new one") + async with self._init_lock: + self._session_id = None + await self.initialize() + continue + response.raise_for_status() + return self._parse(response, request_id) + raise MCPToolError("unreachable") async def call_tool( self, name: str, arguments: dict[str, Any] | None = None ) -> str: - if self._session_id is None: - await self.initialize() - result = await self.call( "tools/call", {"name": name, "arguments": arguments or {}} ) @@ -161,8 +205,6 @@ async def call_tool( return text async def list_tools(self) -> list[dict[str, Any]]: - if self._session_id is None: - await self.initialize() tools = (await self.call("tools/list")).get("tools", []) if not isinstance(tools, list): raise MCPToolError(f"MCP returned a {type(tools).__name__} tool list") diff --git a/src/reactome_mcp/session.py b/src/reactome_mcp/session.py index 47f0e00..aceeaf6 100644 --- a/src/reactome_mcp/session.py +++ b/src/reactome_mcp/session.py @@ -20,6 +20,7 @@ import contextlib import logging import os +import time from pathlib import Path from langchain_core.tools import BaseTool @@ -35,7 +36,13 @@ _manager: MCPProcessManager | None = None _http: MCPHttpClient | None = None _tools: list[BaseTool] | None = None -_failed = False +#: When starting the MCP last failed. Retried after `RETRY_AFTER_SECONDS`: +#: remembered for the life of the process, one failed first connect -- the +#: MCP not up yet after a deploy -- turned live lookups off until the +#: chatbot restarted, while the router kept sending questions to them +#: (review, area 3). +_failed_at: float | None = None +RETRY_AFTER_SECONDS = 60.0 def mcp_server_path() -> Path | None: @@ -53,6 +60,12 @@ def is_configured() -> bool: return mcp_server_url() is not None or mcp_server_path() is not None +def _recently_failed() -> bool: + return ( + _failed_at is not None and time.monotonic() - _failed_at < RETRY_AFTER_SECONDS + ) + + async def get_mcp_tools() -> list[BaseTool] | None: """The MCP tools, starting the server if it is not already running. @@ -61,18 +74,18 @@ async def get_mcp_tools() -> list[BaseTool] | None: every question would turn one misconfiguration into a stall on every request. """ - global _manager, _tools, _failed + global _manager, _http, _tools, _failed_at if _tools is not None: return _tools - if _failed or not is_configured(): + if not is_configured() or _recently_failed(): return None async with _lock: # Another coroutine may have finished while this one waited. if _tools is not None: return _tools - if _failed: + if _recently_failed(): return None url = mcp_server_url() @@ -101,7 +114,7 @@ async def get_mcp_tools() -> list[BaseTool] | None: return None _tools = create_mcp_tools(client) except Exception as exc: - _failed = True + _failed_at = time.monotonic() stderr = "" with contextlib.suppress(Exception): # Best effort: this runs while reporting another failure and @@ -109,9 +122,10 @@ async def get_mcp_tools() -> list[BaseTool] | None: if _manager is not None: stderr = await _manager.stderr_tail() logger.warning( - "MCP server unavailable (%s); live Reactome tools are off for this " - "process. Checked %s.%s", + "MCP server unavailable (%s); live Reactome tools are off for " + "%.0fs, then retried. Checked %s.%s", exc, + RETRY_AFTER_SECONDS, where, f" Server said: {stderr}" if stderr else "", ) diff --git a/src/reactome_mcp/tools.py b/src/reactome_mcp/tools.py index 37aa896..1d035d1 100644 --- a/src/reactome_mcp/tools.py +++ b/src/reactome_mcp/tools.py @@ -16,6 +16,7 @@ """ import logging +import re from typing import Any, Protocol from langchain_core.tools import BaseTool, tool @@ -31,6 +32,21 @@ async def call_tool( ) -> str: ... +#: A line naming the analysis token, or a link that carries it. +_TOKEN_LINE = re.compile(r"token|ANALYSIS=|/AnalysisService/", re.IGNORECASE) + + +def without_token(text: str) -> str: + """An analysis result as the model may see it: no token, no token links. + + The token is a bearer capability for the full result, the reader's + identifiers included. reactome-mcp's reply opens with it, and the live + loop put the reply verbatim into the model's context (review, area 3). + The chat's own gene-list path already drops the same line. + """ + return "\n".join(line for line in text.splitlines() if not _TOKEN_LINE.search(line)) + + def create_mcp_tools(client: ToolCaller) -> list[BaseTool]: """Wrap the curated MCP tools as LangChain tools bound to one client.""" @@ -63,8 +79,10 @@ async def reactome_analyze_identifiers(identifiers: list[str]) -> str: run by Reactome, not a lookup: do not answer such a question from retrieved documents instead. """ - return await client.call_tool( - "reactome_analyze_identifiers", {"identifiers": identifiers} + return without_token( + await client.call_tool( + "reactome_analyze_identifiers", {"identifiers": identifiers} + ) ) @tool diff --git a/tests/agent/test_live_answers_stream.py b/tests/agent/test_live_answers_stream.py index c94c1e7..ac4a99d 100644 --- a/tests/agent/test_live_answers_stream.py +++ b/tests/agent/test_live_answers_stream.py @@ -17,9 +17,21 @@ import asyncio from typing import Any, cast +import pytest + from agent.graph import AgentGraph +@pytest.fixture(autouse=True) +def _live_tools_ready(monkeypatch: pytest.MonkeyPatch) -> None: + """The live path is answered from tools; by default they are up.""" + + async def tools() -> list[Any]: + return [object()] + + monkeypatch.setattr("agent.graph.get_mcp_tools", tools) + + def _chunk(text: str) -> Any: class Chunk: content = text @@ -120,3 +132,31 @@ def test_a_non_live_route_still_waits_for_retrieval() -> None: assert state == "answered" assert "".join(text) == "The actual answer." assert "expanded query" not in "".join(text) + + +def test_without_live_tools_the_fallback_still_waits_for_retrieval( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Routed to live, but the MCP is down: the same node falls back to + retrieval, whose query expander streams first. Opening the boundary at + preprocess put its alternate questions on the search page as the answer + (review, area 3).""" + + async def no_tools() -> None: + return None + + monkeypatch.setattr("agent.graph.get_mcp_tools", no_tools) + events = [ + _preprocess_end(["live"]), + _model_stream("Which organisms are represented in Reactome?"), # expander + { + "event": "on_retriever_end", + "name": "retriever", + "metadata": {}, + "data": {"output": []}, + }, + _model_stream("Reactome covers 15 species."), + ] + state, text = _drive(events) + assert state == "answered" + assert "".join(text) == "Reactome covers 15 species." diff --git a/tests/reactome_mcp/test_http_client.py b/tests/reactome_mcp/test_http_client.py index f95209a..48f5536 100644 --- a/tests/reactome_mcp/test_http_client.py +++ b/tests/reactome_mcp/test_http_client.py @@ -49,7 +49,7 @@ def handler(request: httpx.Request) -> httpx.Response: text=sse( { "jsonrpc": "2.0", - "id": 1, + "id": json.loads(request.content).get("id"), "result": {"serverInfo": {"name": "reactome"}}, } ), @@ -60,7 +60,7 @@ def handler(request: httpx.Request) -> httpx.Response: text=sse( { "jsonrpc": "2.0", - "id": 2, + "id": json.loads(request.content).get("id"), "result": {"content": [{"type": "text", "text": "96"}]}, } ), @@ -87,7 +87,13 @@ def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, headers={"content-type": "text/event-stream", "mcp-session-id": "s"}, - text=sse({"jsonrpc": "2.0", "id": 1, "result": result}), + text=sse( + { + "jsonrpc": "2.0", + "id": json.loads(request.content).get("id"), + "result": result, + } + ), ) assert asyncio.run(_client(handler).call_tool("x")) == "hello" @@ -104,7 +110,11 @@ def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, headers={"content-type": "application/json", "mcp-session-id": "s"}, - json={"jsonrpc": "2.0", "id": 1, "result": result}, + json={ + "jsonrpc": "2.0", + "id": json.loads(request.content).get("id"), + "result": result, + }, ) assert asyncio.run(_client(handler).call_tool("x")) == "json" @@ -117,7 +127,13 @@ def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, headers={"content-type": "text/event-stream"}, - text=sse({"jsonrpc": "2.0", "id": 1, "result": {}}), + text=sse( + { + "jsonrpc": "2.0", + "id": json.loads(request.content).get("id"), + "result": {}, + } + ), ) with pytest.raises(MCPToolError, match="no mcp-session-id"): @@ -131,7 +147,13 @@ def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, headers={"content-type": "text/event-stream", "mcp-session-id": "s"}, - text=sse({"jsonrpc": "2.0", "id": 1, "result": {}}), + text=sse( + { + "jsonrpc": "2.0", + "id": json.loads(request.content).get("id"), + "result": {}, + } + ), ) return httpx.Response( 200, @@ -139,7 +161,7 @@ def handler(request: httpx.Request) -> httpx.Response: text=sse( { "jsonrpc": "2.0", - "id": 2, + "id": json.loads(request.content).get("id"), "error": {"code": -32601, "message": "no such tool"}, } ), @@ -163,7 +185,13 @@ def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, headers={"content-type": "text/event-stream", "mcp-session-id": "s"}, - text=sse({"jsonrpc": "2.0", "id": 1, "result": result}), + text=sse( + { + "jsonrpc": "2.0", + "id": json.loads(request.content).get("id"), + "result": result, + } + ), ) with pytest.raises(MCPToolError, match="pathway not found"): @@ -187,12 +215,18 @@ def handler(request: httpx.Request) -> httpx.Response: return httpx.Response( 200, headers={"content-type": "text/event-stream", "mcp-session-id": "s"}, - text=sse({"jsonrpc": "2.0", "id": 1, "result": {}}), + text=sse( + { + "jsonrpc": "2.0", + "id": json.loads(request.content).get("id"), + "result": {}, + } + ), ) stream = sse({"jsonrpc": "2.0", "method": "notifications/progress"}) + sse( { "jsonrpc": "2.0", - "id": 2, + "id": json.loads(request.content).get("id"), "result": {"content": [{"type": "text", "text": "the answer"}]}, } ) @@ -201,3 +235,82 @@ def handler(request: httpx.Request) -> httpx.Response: ) assert asyncio.run(_client(handler).call_tool("x")) == "the answer" + + +# --- review, area 3 ----------------------------------------------------------- + + +def _echoing_server( + log: list[dict[str, Any]], lost_once: list[bool] | None = None +) -> Any: + """Answers each request under its own id; can forget the session once.""" + + def handler(request: httpx.Request) -> httpx.Response: + body = json.loads(request.content) + log.append(body) + if body.get("method") == "initialize": + result: dict[str, Any] = {"serverInfo": {"name": "fake"}} + elif lost_once and lost_once[0] and body.get("method") == "tools/call": + lost_once[0] = False + return httpx.Response( + 400, + json={ + "jsonrpc": "2.0", + "error": { + "code": -32000, + "message": "No valid session. Send initialize first.", + }, + "id": None, + }, + ) + else: + name = (body.get("params") or {}).get("name", "") + result = {"content": [{"type": "text", "text": f"answer to {name}"}]} + return httpx.Response( + 200, + headers={"content-type": "text/event-stream", "mcp-session-id": "s"}, + text=sse({"jsonrpc": "2.0", "id": body.get("id"), "result": result}), + ) + + return handler + + +def test_every_request_has_its_own_id() -> None: + # Every call sent id 2 on the shared session; the server matched replies + # by id, so overlapping users got each other's results (review, area 3). + log: list[dict[str, Any]] = [] + client = _client(_echoing_server(log)) + + async def both() -> list[str]: + return list( + await asyncio.gather( + client.call_tool("search"), client.call_tool("species") + ) + ) + + assert asyncio.run(both()) == ["answer to search", "answer to species"] + ids = [b["id"] for b in log if b.get("method") == "tools/call"] + assert len(ids) == len(set(ids)) == 2 + + +def test_a_reply_to_another_request_is_refused() -> None: + def handler(request: httpx.Request) -> httpx.Response: + body = json.loads(request.content) + reply_id = body.get("id") if body.get("method") == "initialize" else 999 + return httpx.Response( + 200, + headers={"content-type": "text/event-stream", "mcp-session-id": "s"}, + text=sse({"jsonrpc": "2.0", "id": reply_id, "result": {"content": []}}), + ) + + with pytest.raises(MCPToolError, match="not"): + asyncio.run(_client(handler).call_tool("x")) + + +def test_a_lost_session_is_renewed_once() -> None: + # reactome-mcp answers 400 "No valid session" after it restarts; the + # client never re-initialized, so live lookups stayed off. + log: list[dict[str, Any]] = [] + client = _client(_echoing_server(log, lost_once=[True])) + assert asyncio.run(client.call_tool("search")) == "answer to search" + assert [b.get("method") for b in log].count("initialize") == 2 diff --git a/tests/reactome_mcp/test_live_answer.py b/tests/reactome_mcp/test_live_answer.py index 2b9d5c9..d6a360e 100644 --- a/tests/reactome_mcp/test_live_answer.py +++ b/tests/reactome_mcp/test_live_answer.py @@ -98,7 +98,10 @@ def test_a_failing_tool_is_reported_to_the_model_not_raised() -> None: assert "could not" in answer.lower() tool_messages = [m for m in llm.seen[-1] if isinstance(m, ToolMessage)] - assert any("the service is down" in str(m.content) for m in tool_messages) + # Told that it failed, and of what type -- not the message, which carried + # the MCP's internal URL to OpenAI (review, area 3). + assert any("lookup failed" in str(m.content) for m in tool_messages) + assert not any("the service is down" in str(m.content) for m in tool_messages) def test_an_invented_tool_name_is_survivable() -> None: @@ -279,3 +282,30 @@ def test_the_report_is_optional_so_existing_callers_are_unaffected() -> None: answer_from_live_services(llm, [failing_tool], "anything") ).lower() ) + + +def test_the_analysis_token_never_reaches_the_model() -> None: + from reactome_mcp.tools import without_token + + raw = ( + "## Analysis\n**Token:** MjAyNjEwMDRfMTIz\n" + "View: https://reactome.org/PathwayBrowser/#/DTAB=AN&ANALYSIS=MjAyNjEwMDRfMTIz\n" + "1. Cell Cycle (R-HSA-1640170) p=1e-6 FDR=1e-4\n" + ) + shown = without_token(raw) + assert "MjAyNjEwMDRfMTIz" not in shown + assert "Cell Cycle (R-HSA-1640170)" in shown + + +def test_a_round_runs_at_most_a_few_tool_calls() -> None: + # 150 calls in one round ran 300 upstream requests, one after another + # (review, area 3). The rest are answered, so the next request is valid. + from reactome_mcp.answer import MAX_CALLS_PER_ROUND + + calls = [_call("reactome_species", str(i)) for i in range(20)] + llm = _FakeLLM([AIMessage("", tool_calls=calls), AIMessage("Done.")]) + asyncio.run(answer_from_live_services(llm, [reactome_species], "many")) + tool_messages = [m for m in llm.seen[-1] if isinstance(m, ToolMessage)] + assert len(tool_messages) == 20 # every call answered + ran = [m for m in tool_messages if "Not run" not in str(m.content)] + assert len(ran) == MAX_CALLS_PER_ROUND diff --git a/tests/reactome_mcp/test_session_retry.py b/tests/reactome_mcp/test_session_retry.py new file mode 100644 index 0000000..22c9f48 --- /dev/null +++ b/tests/reactome_mcp/test_session_retry.py @@ -0,0 +1,48 @@ +"""A failed MCP start is retried, not remembered for the life of the process.""" + +import asyncio +from typing import Any + +import pytest + +from reactome_mcp import session + + +@pytest.fixture(autouse=True) +def _fresh(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(session, "_tools", None) + monkeypatch.setattr(session, "_failed_at", None) + monkeypatch.setenv("REACTOME_MCP_URL", "http://mcp.test") + monkeypatch.delenv("REACTOME_MCP_SERVER", raising=False) + + +def test_a_failed_start_is_retried_after_a_while( + monkeypatch: pytest.MonkeyPatch, +) -> None: + attempts: list[int] = [] + up = {"now": False} + + class _Client: + def __init__(self, _url: str) -> None: + pass + + async def initialize(self) -> None: + attempts.append(1) + if not up["now"]: + raise ConnectionError("not up yet") + + monkeypatch.setattr(session, "MCPHttpClient", _Client) + monkeypatch.setattr(session, "create_mcp_tools", lambda _c: [object()]) + clock = {"t": 1000.0} + monkeypatch.setattr("reactome_mcp.session.time.monotonic", lambda: clock["t"]) + + async def get() -> Any: + return await session.get_mcp_tools() + + assert asyncio.run(get()) is None # the MCP is not up yet after a deploy + assert asyncio.run(get()) is None # not hammered on every question + assert len(attempts) == 1 + up["now"] = True + clock["t"] += session.RETRY_AFTER_SECONDS + 1 + assert asyncio.run(get()) is not None # and back, without a restart + assert len(attempts) == 2