fix(mcp): preserve elicitation context and report relay failures (#45255)

* fix(mcp): correlate elicitation relays and report failures

* fix(mcp): preserve elicitation timeout errors on Python 3.10

* fix(mcp): keep request ID alias in type-checking imports

* fix(mcp): preserve elicitation type narrowing and clarify HTTP limits

* ci(mcp): publish elicitation coverage from GitHub Actions

---------

Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
joshua-berri 2026-10-08 10:18:25 -07:00 • committed by GitHub
parent 9f150f74ed
commit 49266e3b17
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 249 additions and 197 deletions

View file

@ -346,6 +346,16 @@ jobs:
timeout-minutes: 60
job-timeout-minutes: 100
- shard: mcp-elicitation
artifact-name: mcp-elicitation
test-path: >-
tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py
tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py
workers: 2
reruns: 0
timeout-minutes: 20
job-timeout-minutes: 60
- shard: proxy-infra
artifact-name: proxy-infra
test-path: >-

View file

@ -13,3 +13,20 @@ Listing failures retain per-server outcome metadata. An incomplete upstream cata
For a scoped rollout, keep the previous source build serving the control pool and send only selected clients to a separate candidate pool. All candidate replicas must share configuration and salt. Verify page one on one candidate replica and continuation on another, plus a fresh listing and tool call on the control pool. Do not mirror tool calls between pools
For rollback, stop sending new requests to the candidate pool, drain its in-flight operations, and return selected clients to the control pool. Clients must discard candidate cursors and start a fresh listing when crossing versions; older gateways do not validate these cursors. Keep the registry-revision migration installed when rolling back pagination. Verify a fresh listing and a tool call after switching pools
## Elicitation
Elicitation is disabled by default. Enable it per upstream in `config.yaml`; `allow_elicitation` is currently YAML-only and is not editable through the Admin UI or database-backed server API
```yaml
mcp_servers:
interactive:
url: https://mcp.example.com/mcp
transport: http
allow_elicitation: true
timeout: 60
```
The downstream MCP client must advertise the requested form or URL capability during initialization. Use the gateway's legacy SSE endpoint (`/mcp/sse`) for the verified interactive form and URL relay path. The current Streamable HTTP endpoint (`/mcp`) can lose initialization state before a tool call, so even a client that advertised support receives an explicit elicitation error instead of an input request. Successful interactive relay over Streamable HTTP is not currently verified. Stateless calls and LLM tool bridges with no downstream MCP client also receive explicit errors
The relay wait uses the upstream server's existing `timeout` setting, or `LITELLM_MCP_CLIENT_TIMEOUT` (60 seconds by default). The enclosing tool call also retains its existing timeout. Unsupported modes, disconnects and relay failures return errors, never a fabricated user decline. Actual user accept, decline and cancel responses are preserved; cancellation stops the relay

View file

@ -1,30 +1,35 @@
"""
MCP Elicitation Handler
Handles `elicitation/create` requests from upstream MCP servers by either:
1. Relaying them to the connected downstream MCP client (if it supports elicitation)
2. Returning a decline/error response (if no downstream client or unsupported)
Supports both Form mode (structured data collection) and URL mode (external URL
navigation for sensitive interactions like OAuth).
MCP Spec Reference:
https://modelcontextprotocol.io/specification/2025-11-25/client/elicitation
"""
"""Relay upstream elicitation through the initiating downstream MCP request."""
from typing import TYPE_CHECKING, Final, Protocol, Union
from __future__ import annotations
import asyncio
from typing import TYPE_CHECKING, Final, Protocol
from litellm._logging import verbose_logger
from litellm.constants import MCP_CLIENT_TIMEOUT
if TYPE_CHECKING:
from mcp.types import (
INTERNAL_ERROR,
INVALID_REQUEST,
REQUEST_TIMEOUT,
ClientCapabilities,
ElicitRequestedSchema,
ElicitRequestFormParams,
ElicitRequestParams,
ElicitRequestURLParams,
ElicitResult,
ErrorData,
RequestId,
)
# Guard imports that require the mcp package
try:
from mcp.types import (
INTERNAL_ERROR,
INVALID_REQUEST,
REQUEST_TIMEOUT,
ClientCapabilities,
ElicitRequestedSchema,
ElicitRequestFormParams,
ElicitRequestParams,
ElicitRequestURLParams,
@ -38,136 +43,65 @@ except ImportError:
class _DownstreamElicitSession(Protocol):
"""The downstream MCP client session methods this module relays elicitation requests through."""
async def elicit_url(
self, message: str, url: str, elicitation_id: str, related_request_id: RequestId | None = None
) -> ElicitResult: ...
async def elicit_url(self, message: str, url: str, elicitation_id: str) -> "ElicitResult": ...
async def elicit_form(self, message: str, requested_schema: dict[str, object]) -> "ElicitResult": ...
async def elicit(self, message: str, requested_schema: dict[str, object]) -> "ElicitResult": ...
async def elicit_form(
self, message: str, requested_schema: ElicitRequestedSchema, related_request_id: RequestId | None = None
) -> ElicitResult: ...
async def handle_elicitation_request(
context: object,
params: "ElicitRequestParams",
params: ElicitRequestParams,
downstream_session: _DownstreamElicitSession | None = None,
downstream_capabilities: object = None,
) -> Union["ElicitResult", "ErrorData"]:
"""
Handle an MCP elicitation/create request from an upstream MCP server.
In Gateway mode (Mode A), we relay the elicitation request to the
connected downstream client if they declared elicitation capabilities.
In Tool Bridge mode (Mode B), there's no persistent downstream MCP
client, so we return a decline response.
Args:
context: MCP RequestContext from the upstream server connection.
params: The ElicitRequestParams (either form or URL mode).
downstream_session: The ServerSession to the downstream client,
if available (for relaying).
downstream_capabilities: The downstream client's declared
capabilities, used to check elicitation support.
Returns:
ElicitResult with the user's response, or ErrorData on failure.
"""
downstream_capabilities: ClientCapabilities | None = None,
related_request_id: RequestId | None = None,
timeout: float = MCP_CLIENT_TIMEOUT,
) -> ElicitResult | ErrorData:
if not MCP_ELICITATION_AVAILABLE:
return ErrorData(
code=-1,
message="MCP elicitation is not available (mcp package not installed)",
)
return ErrorData(code=INTERNAL_ERROR, message="MCP elicitation is not available")
if downstream_session is None:
return ErrorData(code=INVALID_REQUEST, message="MCP elicitation requires a connected downstream MCP client")
try:
mode: Final = getattr(params, "mode", "form")
verbose_logger.info(
"MCP elicitation: received request mode=%s, message=%s",
mode,
getattr(params, "message", ""),
return await asyncio.wait_for(
_relay_elicitation_to_downstream(params, downstream_session, downstream_capabilities, related_request_id),
timeout=timeout,
)
# Check if we have a downstream session to relay to
if downstream_session is not None:
return await _relay_elicitation_to_downstream(
params=params,
downstream_session=downstream_session,
downstream_capabilities=downstream_capabilities,
)
# No downstream session — we're in Tool Bridge mode
# or the client doesn't support elicitation
verbose_logger.info("MCP elicitation: no downstream session available, declining")
return ElicitResult(
action="decline",
)
except Exception as e:
verbose_logger.exception("MCP elicitation handler failed: %s", e)
except asyncio.TimeoutError:
return ErrorData(code=REQUEST_TIMEOUT, message="MCP elicitation timed out waiting for the downstream client")
except Exception:
verbose_logger.warning("MCP elicitation: downstream relay failed")
return ErrorData(
code=-1,
message=f"Elicitation failed: {e}",
code=INTERNAL_ERROR, message="MCP elicitation failed while communicating with the downstream client"
)
async def _relay_elicitation_to_downstream(
params: "ElicitRequestParams",
params: ElicitRequestParams,
downstream_session: _DownstreamElicitSession,
downstream_capabilities: object = None,
) -> Union["ElicitResult", "ErrorData"]:
"""
Relay an elicitation request to the downstream MCP client.
Uses the ServerSession's elicit_form() or elicit_url() methods to
send the elicitation request back to the connected client.
Args:
params: The elicitation request parameters.
downstream_session: The ServerSession connected to the downstream client.
downstream_capabilities: Client capabilities to check support.
Returns:
ElicitResult from the downstream client.
"""
mode: Final = getattr(params, "mode", "form")
# Check if the downstream client supports the requested mode
if downstream_capabilities is not None:
elicit_caps: Final[object] = getattr(downstream_capabilities, "elicitation", None)
if elicit_caps is None:
verbose_logger.info("MCP elicitation: downstream client does not support elicitation")
return ElicitResult(action="decline")
if mode == "url":
url_cap: Final[object] = getattr(elicit_caps, "url", None)
if url_cap is None:
verbose_logger.info("MCP elicitation: downstream client does not support URL mode")
return ElicitResult(action="decline")
if mode == "form":
form_cap: Final[object] = getattr(elicit_caps, "form", None)
if form_cap is None:
verbose_logger.info("MCP elicitation: downstream client does not support form mode")
return ElicitResult(action="decline")
try:
if mode == "url" and isinstance(params, ElicitRequestURLParams):
# URL mode: relay URL to client for external navigation
verbose_logger.info(
"MCP elicitation: relaying URL mode to downstream, url=%s",
getattr(params, "url", ""),
)
result = await downstream_session.elicit_url(
message=params.message,
url=params.url,
elicitation_id=params.elicitation_id,
)
elif isinstance(params, ElicitRequestFormParams):
# Form mode: relay structured form to client
verbose_logger.info("MCP elicitation: relaying form mode to downstream")
result = await downstream_session.elicit_form(
message=params.message,
requested_schema=params.requested_schema,
)
else:
# Fallback for generic ElicitRequestParams — pass an empty schema
# since elicit() requires requested_schema as a positional arg.
verbose_logger.info("MCP elicitation: relaying generic elicitation to downstream")
result = await downstream_session.elicit(
message=getattr(params, "message", ""),
requested_schema=getattr(params, "requested_schema", {}),
)
verbose_logger.info(
"MCP elicitation: downstream responded with action=%s",
getattr(result, "action", "unknown"),
downstream_capabilities: ClientCapabilities | None = None,
related_request_id: RequestId | None = None,
) -> ElicitResult | ErrorData:
capabilities: Final = downstream_capabilities.elicitation if downstream_capabilities is not None else None
if capabilities is None:
return ErrorData(code=INVALID_REQUEST, message="Downstream client has not advertised elicitation support")
if isinstance(params, ElicitRequestURLParams):
if capabilities.url is None:
return ErrorData(code=INVALID_REQUEST, message="Downstream client does not support URL elicitation")
return await downstream_session.elicit_url(
message=params.message,
url=params.url,
elicitation_id=params.elicitation_id,
related_request_id=related_request_id,
)
return result
except Exception as e:
verbose_logger.warning("MCP elicitation: failed to relay to downstream: %s", e)
# If relay fails, decline gracefully
return ElicitResult(action="decline")
if capabilities.form is None and capabilities.url is not None:
return ErrorData(code=INVALID_REQUEST, message="Downstream client does not support form elicitation")
if not isinstance(params, ElicitRequestFormParams):
return ErrorData(code=INVALID_REQUEST, message="Unsupported MCP elicitation parameters")
return await downstream_session.elicit_form(
message=params.message,
requested_schema=params.requested_schema,
related_request_id=related_request_id,
)

View file

@ -11,7 +11,9 @@ from mcp.types import (
ErrorData,
)
from litellm.constants import MCP_CLIENT_TIMEOUT
from litellm.proxy._experimental.mcp_server.contracts import OperationContext
from litellm.proxy._experimental.mcp_server.mcp_context import get_active_mcp_request_ctx
from litellm.proxy._types import UserAPIKeyAuth
@ -62,11 +64,19 @@ def create_sampling_callback(
return callback
def create_elicitation_callback() -> ElicitationCallback:
def create_elicitation_callback(timeout: float | None = None) -> ElicitationCallback:
from litellm.proxy._experimental.mcp_server.server import get_active_mcp_session
downstream_session: Final = get_active_mcp_session()
downstream_capabilities: Final = getattr(downstream_session, "capabilities", None)
request: Final = get_active_mcp_request_ctx()
client_params: Final = downstream_session.client_params if downstream_session is not None else None
downstream_capabilities: Final = (
client_params.capabilities.model_copy(deep=True) if client_params is not None else None
)
related_request_id: Final = (
request.request_id if request is not None and request.session is downstream_session else None
)
relay_timeout: Final = timeout if timeout is not None else MCP_CLIENT_TIMEOUT
async def callback(context: object, params: ElicitRequestParams) -> ElicitResult | ErrorData:
from litellm.proxy._experimental.mcp_server.elicitation_handler import handle_elicitation_request
@ -76,6 +86,8 @@ def create_elicitation_callback() -> ElicitationCallback:
params=params,
downstream_session=downstream_session,
downstream_capabilities=downstream_capabilities,
related_request_id=related_request_id,
timeout=relay_timeout,
)
return callback

View file

@ -1735,12 +1735,12 @@ def _create_sampling_callback(
return create_sampling_callback(user_api_key_auth, raw_headers, client_ip, operation_context)
def _create_elicitation_callback():
def _create_elicitation_callback(timeout: float | None = None):
if not MCP_ELICITATION_AVAILABLE:
return None
from litellm.proxy._experimental.mcp_server.legacy_callbacks import create_elicitation_callback
return create_elicitation_callback()
return create_elicitation_callback(timeout=timeout)
def _record_mcp_guardrail_evaluations(
@ -4250,7 +4250,11 @@ class MCPServerManager:
if resolved_server.allow_sampling
else None
),
elicitation_callback=(_create_elicitation_callback() if resolved_server.allow_elicitation else None),
elicitation_callback=(
_create_elicitation_callback(timeout=resolved_server.timeout)
if resolved_server.allow_elicitation
else None
),
)
_create_mcp_client = create_mcp_client

View file

@ -7,12 +7,18 @@ well as the decline paths used in tool-bridge mode or when the downstream
client lacks the requested elicitation capability.
"""
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from mcp.types import (
ClientCapabilities,
ElicitationCapability,
FormElicitationCapability,
UrlElicitationCapability,
REQUEST_TIMEOUT,
ElicitRequestFormParams,
ElicitRequestURLParams,
ElicitResult,
@ -43,23 +49,21 @@ def _url_params(message: str = "please authorize") -> ElicitRequestURLParams:
)
def _caps(*, url=True, form=True) -> SimpleNamespace:
elicit = SimpleNamespace(
url=object() if url else None,
form=object() if form else None,
)
return SimpleNamespace(elicitation=elicit)
def _caps(*, url=True, form=True) -> ClientCapabilities:
return ClientCapabilities(elicitation=ElicitationCapability(
url=UrlElicitationCapability() if url else None,
form=FormElicitationCapability() if form else None,
))
class TestHandleElicitationRequest:
async def test_should_decline_when_no_downstream_session(self):
async def test_should_error_when_no_downstream_session(self):
result = await handle_elicitation_request(
context=SimpleNamespace(),
params=_form_params(),
downstream_session=None,
)
assert isinstance(result, ElicitResult)
assert result.action == "decline"
assert isinstance(result, ErrorData)
async def test_should_relay_to_downstream_when_session_present(self):
accepted = ElicitResult(action="accept", content={"name": "ada"})
@ -69,7 +73,7 @@ class TestHandleElicitationRequest:
context=SimpleNamespace(),
params=_form_params(),
downstream_session=session,
downstream_capabilities=None,
downstream_capabilities=_caps(),
)
assert result is accepted
@ -85,21 +89,6 @@ class TestHandleElicitationRequest:
assert isinstance(result, ErrorData)
assert "not available" in result.message
async def test_should_return_error_data_on_unexpected_failure(self):
class _ExplodingParams:
mode = "form"
@property
def message(self):
raise RuntimeError("boom")
result = await handle_elicitation_request(
context=SimpleNamespace(),
params=_ExplodingParams(),
downstream_session=None,
)
assert isinstance(result, ErrorData)
assert "boom" in result.message
class TestRelayElicitationToDownstream:
@ -136,25 +125,17 @@ class TestRelayElicitationToDownstream:
assert kwargs["url"] == "https://example.com/oauth"
assert kwargs["elicitation_id"] == "elc-1"
async def test_should_use_generic_elicit_for_unknown_param_type(self):
accepted = ElicitResult(action="accept")
session = SimpleNamespace(elicit=AsyncMock(return_value=accepted))
# A bare params object that is neither Form nor URL params triggers
# the generic fallback path.
params = SimpleNamespace(mode="form", message="hi", requested_schema={})
result = await _relay_elicitation_to_downstream(
params=params,
downstream_session=session,
downstream_capabilities=None,
)
assert result is accepted
session.elicit.assert_awaited_once()
async def test_should_decline_when_elicitation_unsupported(self):
async def test_should_reject_invalid_elicitation_parameters(self):
session = SimpleNamespace(elicit_form=AsyncMock())
caps = SimpleNamespace(elicitation=None)
result = await _relay_elicitation_to_downstream(
params=SimpleNamespace(mode="form"), downstream_session=session, downstream_capabilities=_caps(),
)
assert isinstance(result, ErrorData)
session.elicit_form.assert_not_awaited()
async def test_should_error_when_elicitation_unsupported(self):
session = SimpleNamespace(elicit_form=AsyncMock())
caps = ClientCapabilities()
result = await _relay_elicitation_to_downstream(
params=_form_params(),
@ -162,11 +143,10 @@ class TestRelayElicitationToDownstream:
downstream_capabilities=caps,
)
assert isinstance(result, ElicitResult)
assert result.action == "decline"
assert isinstance(result, ErrorData)
session.elicit_form.assert_not_awaited()
async def test_should_decline_url_mode_when_url_unsupported(self):
async def test_should_error_url_mode_when_url_unsupported(self):
session = SimpleNamespace(elicit_url=AsyncMock())
result = await _relay_elicitation_to_downstream(
@ -175,11 +155,10 @@ class TestRelayElicitationToDownstream:
downstream_capabilities=_caps(url=False, form=True),
)
assert isinstance(result, ElicitResult)
assert result.action == "decline"
assert isinstance(result, ErrorData)
session.elicit_url.assert_not_awaited()
async def test_should_decline_form_mode_when_form_unsupported(self):
async def test_should_error_form_mode_when_form_unsupported(self):
session = SimpleNamespace(elicit_form=AsyncMock())
result = await _relay_elicitation_to_downstream(
@ -188,24 +167,112 @@ class TestRelayElicitationToDownstream:
downstream_capabilities=_caps(url=True, form=False),
)
assert isinstance(result, ElicitResult)
assert result.action == "decline"
assert isinstance(result, ErrorData)
session.elicit_form.assert_not_awaited()
async def test_should_decline_when_downstream_relay_raises(self):
async def test_should_error_when_downstream_relay_raises(self):
session = SimpleNamespace(
elicit_form=AsyncMock(side_effect=RuntimeError("transport closed"))
)
result = await _relay_elicitation_to_downstream(
result = await handle_elicitation_request(
context=None,
params=_form_params(),
downstream_session=session,
downstream_capabilities=_caps(form=True),
)
assert isinstance(result, ElicitResult)
assert result.action == "decline"
assert isinstance(result, ErrorData)
if __name__ == "__main__":
pytest.main([__file__, "-v"])
@pytest.mark.asyncio
async def test_relay_failure_is_not_user_decline():
session = SimpleNamespace(elicit_form=AsyncMock(side_effect=RuntimeError("private upstream credential")))
result = await handle_elicitation_request(
context=None, params=_form_params(), downstream_session=session, downstream_capabilities=_caps(),
)
assert isinstance(result, ErrorData), "A failed relay must not claim that the user declined"
assert "private upstream credential" not in result.message
@pytest.mark.asyncio
async def test_unknown_client_capabilities_prevent_relay():
session = SimpleNamespace(elicit_form=AsyncMock(return_value=ElicitResult(action="accept")))
result = await handle_elicitation_request(
context=None, params=_form_params(), downstream_session=session, downstream_capabilities=None,
)
assert isinstance(result, ErrorData), "Unknown capabilities must fail explicitly"
session.elicit_form.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["form", "url"])
@pytest.mark.parametrize("action", ["accept", "decline", "cancel"])
async def test_relay_preserves_response_and_request_correlation(mode, action):
response = ElicitResult(action=action)
request = AsyncMock(return_value=response)
session = SimpleNamespace(elicit_form=request, elicit_url=request)
params = _form_params() if mode == "form" else _url_params()
result = await handle_elicitation_request(
context=SimpleNamespace(request_id="upstream-id"), params=params,
downstream_session=session, downstream_capabilities=_caps(), related_request_id=0,
)
assert result is response
assert request.await_args.kwargs["related_request_id"] == 0
if mode == "url":
assert request.await_args.kwargs["url"] == params.url
assert request.await_args.kwargs["elicitation_id"] == params.elicitation_id
else:
assert request.await_args.kwargs["requested_schema"] == params.requested_schema
@pytest.mark.asyncio
async def test_expired_relay_deadline_returns_explicit_timeout():
request = AsyncMock()
result = await handle_elicitation_request(
context=None, params=_form_params(), downstream_session=SimpleNamespace(elicit_form=request),
downstream_capabilities=_caps(), timeout=0,
)
assert isinstance(result, ErrorData)
assert result.code == REQUEST_TIMEOUT
assert "timed out" in result.message
request.assert_not_awaited()
@pytest.mark.asyncio
async def test_relay_cancellation_releases_downstream_waiter():
started = asyncio.Event()
finished = asyncio.Event()
async def wait_for_user(**kwargs):
started.set()
try:
await asyncio.Event().wait()
finally:
finished.set()
task = asyncio.create_task(handle_elicitation_request(
context=None, params=_form_params(), downstream_session=SimpleNamespace(elicit_form=wait_for_user),
downstream_capabilities=_caps(), related_request_id="tool-call",
))
await started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert finished.is_set()
@pytest.mark.asyncio
async def test_legacy_empty_elicitation_capability_supports_form():
accepted = ElicitResult(action="accept")
session = SimpleNamespace(elicit_form=AsyncMock(return_value=accepted))
result = await handle_elicitation_request(
context=None, params=_form_params(), downstream_session=session,
downstream_capabilities=ClientCapabilities(elicitation=ElicitationCapability()),
related_request_id="legacy-call",
)
assert result is accepted

View file

@ -155,18 +155,26 @@ async def test_elicitation_callback_keeps_initiating_session():
from litellm.proxy._experimental.mcp_server import server as legacy_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_elicitation_callback
initiating = MagicMock()
replacement = MagicMock()
recorder = AsyncMock()
from types import SimpleNamespace
from mcp.types import ClientCapabilities, ElicitationCapability, FormElicitationCapability, ElicitRequestFormParams, ElicitResult
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
capabilities = ClientCapabilities(elicitation=ElicitationCapability(form=FormElicitationCapability()))
accepted = ElicitResult(action="accept")
request = AsyncMock(return_value=accepted)
initiating = SimpleNamespace(client_params=SimpleNamespace(capabilities=capabilities), elicit_form=request)
token = legacy_server.active_mcp_session_var.set(initiating)
request_token = active_mcp_request_ctx_var.set(SimpleNamespace(session=initiating, request_id="initiating-call"))
try:
callback = _create_elicitation_callback()
legacy_server.active_mcp_session_var.set(replacement)
with patch("litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request", recorder):
await callback(None, None)
assert recorder.await_args.kwargs["downstream_session"] is initiating
assert recorder.await_args.kwargs["downstream_capabilities"] is initiating.capabilities
legacy_server.active_mcp_session_var.set(SimpleNamespace())
active_mcp_request_ctx_var.set(None)
capabilities.elicitation = None
result = await callback(None, ElicitRequestFormParams(message="Confirm", requested_schema={"type":"object"}))
assert result is accepted
assert request.await_args.kwargs["related_request_id"] == "initiating-call"
finally:
active_mcp_request_ctx_var.reset(request_token)
legacy_server.active_mcp_session_var.reset(token)