mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-06 01:45:27 +03:00
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:
1 parent
0b0de29557
commit
e7668b2c5f
2 files changed
+43
-9
No files matched your search
@@ -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.
|
||||
|
||||
|
||||
Reference in new issue
Block a user