diff --git a/libs/langchain_v1/langchain/mcp/elicitation.py b/libs/langchain_v1/langchain/mcp/elicitation.py index f0056be8d9..77955a5c30 100644 --- a/libs/langchain_v1/langchain/mcp/elicitation.py +++ b/libs/langchain_v1/langchain/mcp/elicitation.py @@ -6,10 +6,14 @@ call to be retried with the answers. This module drives that loop and sources each answer from `interrupt()`, so the human already reviewing an agent's work answers the server's question too. -The loop is driven here rather than through FastMCP's own handler because -FastMCP converts any exception a handler raises into an MCP error, which would -swallow the `GraphInterrupt` that suspends the graph. Calling `interrupt()` from -this module's own frame lets it propagate. +The loop is driven here rather than through the SDK's own +`run_input_required_driver` because that driver answers each request from a +callback, run concurrently in a task group. LangGraph matches resume values to +`interrupt()` calls by their order in the node, so firing them concurrently +would scramble that matching — and FastMCP converts any exception a callback +raises into an MCP error, swallowing the `GraphInterrupt` that suspends the +graph. Calling `interrupt()` from this module's own frame keeps one interrupt +per round and lets it propagate. The retry bounds mirror the SDK driver's. """ from __future__ import annotations @@ -18,6 +22,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar import anyio from langgraph.types import interrupt +from mcp import InputRequiredRoundsExceededError from mcp.types import ( CallToolResult, ElicitRequest, @@ -38,8 +43,9 @@ if TYPE_CHECKING: _ResultT = TypeVar("_ResultT") -_STILL_WORKING_SLEEP_SECONDS = 0.05 -"""Pause before retrying a round that asked nothing, so the retry is not a spin.""" +_STATE_ONLY_BACKOFF_INITIAL_SECONDS = 0.05 +_STATE_ONLY_BACKOFF_CAP_SECONDS = 0.25 +"""Backoff for rounds that ask nothing, matching the SDK's own driver.""" ELICITATION_INTERRUPT_TYPE: Final = "mcp_elicitation" """Discriminator on the interrupt payload, so a handler can recognize it.""" @@ -246,17 +252,28 @@ async def _call_tool_with_interrupts( Raises: GraphInterrupt: Every time the server asks something that has not been answered yet. This is the mechanism, not a failure. + InputRequiredRoundsExceededError: If the server keeps asking past + `client.input_required_max_rounds`, so a server that never reaches + a terminal result cannot loop forever. NotImplementedError: If the server asks for sampling or roots. ValueError: If a resumed answer is missing or malformed. """ session = client.session + max_rounds = client.input_required_max_rounds result = await _await_monitored( client, session.call_tool(tool_name, arguments, allow_input_required=True) ) + rounds = 0 + state_only_delay = _STATE_ONLY_BACKOFF_INITIAL_SECONDS while isinstance(result, InputRequiredResult): + rounds += 1 + if rounds > max_rounds: + raise InputRequiredRoundsExceededError(max_rounds) + responses: InputResponses | None = None if result.input_requests: + state_only_delay = _STATE_ONLY_BACKOFF_INITIAL_SECONDS requests = _elicit_requests(result.input_requests, tool_name) request: MCPElicitationInterrupt = { "type": ELICITATION_INTERRUPT_TYPE, @@ -268,8 +285,9 @@ async def _call_tool_with_interrupts( else: # A round carrying only `request_state` means the server is still # working and wants to be asked again. Nobody needs to be - # interrupted for that; just pause so the retry is not a spin. - await anyio.sleep(_STILL_WORKING_SLEEP_SECONDS) + # interrupted for that; just back off so the retry is not a spin. + await anyio.sleep(state_only_delay) + state_only_delay = min(state_only_delay * 2, _STATE_ONLY_BACKOFF_CAP_SECONDS) result = await _await_monitored( client, diff --git a/libs/langchain_v1/tests/unit_tests/mcp/test_elicitation.py b/libs/langchain_v1/tests/unit_tests/mcp/test_elicitation.py index 587f3b65d2..264c8ffdf6 100644 --- a/libs/langchain_v1/tests/unit_tests/mcp/test_elicitation.py +++ b/libs/langchain_v1/tests/unit_tests/mcp/test_elicitation.py @@ -12,6 +12,7 @@ import pytest from fastmcp import Context, FastMCP from langgraph.checkpoint.memory import InMemorySaver from langgraph.types import Command +from mcp import InputRequiredRoundsExceededError from mcp.server.mcpserver import Elicit, MCPServer, Resolve from mcp.shared.exceptions import MCPError from mcp.types import ( @@ -173,8 +174,9 @@ class _FakeSession: class _FakeClient: - def __init__(self, session: _FakeSession) -> None: + def __init__(self, session: _FakeSession, max_rounds: int = 10) -> None: self.session = session + self.input_required_max_rounds = max_rounds def _elicit_round(key: str = "ask") -> InputRequiredResult: @@ -232,6 +234,20 @@ async def test_a_state_only_round_is_retried_without_asking_anyone() -> None: assert session.calls[1]["input_responses"] is None +@pytest.mark.asyncio +async def test_a_server_that_never_finishes_is_cut_off() -> None: + """State-only rounds are bounded, so a stalled server cannot spin forever.""" + working = InputRequiredResult(input_requests=None, request_state="state-1") + session = _FakeSession([working]) + client = _FakeClient(session, max_rounds=3) + + with pytest.raises(InputRequiredRoundsExceededError): + await _call_tool_with_interrupts(client, "stalled", {}) # type: ignore[arg-type] + + # The initial call plus one retry per permitted round, and no more. + assert len(session.calls) == 4 + + def _multi_question_server(calls: dict[str, int]) -> FastMCP[None]: """A server that asks two things at once, then a third on the next round.