fix(langchain): bound the MCP input-required retry loop

A server that needs input can answer `tools/call` with an `InputRequiredResult`
carrying only `request_state` and no requests, meaning "still working, ask me
again". Retrying is the protocol's own continuation mechanism, but the retry
was unbounded: a fixed 50ms sleep and no cap, so a server that never reached a
terminal result pinned the agent in a tool call at 20 requests a second
indefinitely.

The loop now honors the client's `input_required_max_rounds` (10 by default,
and settable on `fastmcp.Client`), raising `InputRequiredRoundsExceededError`
once exhausted, and backs off exponentially from 50ms to a 250ms cap — the same
bounds the MCP SDK's own `run_input_required_driver` applies, so a stalled
server is cut off the same way it would be through `client.call_tool`.

Delegating to that driver outright is not an option: it answers each request
from a callback run concurrently in a task group, while LangGraph matches
resume values to `interrupt()` calls by their order in the node. Firing them
concurrently would scramble that matching, so the loop stays here and borrows
only the bounds. That reasoning is now recorded in the module docstring.
This commit is contained in:
Sydney Runkle committed 2026-08-27 10:09:21 -07:00
1 parent 0b0de29557
commit e7668b2c5f
2 files changed
+43 -9

No files matched your search

+26 -8
View File
@@ -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,
@@ -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.