From 49266e3b17afa89477557c46ebd86a6a3e56da38 Mon Sep 17 00:00:00 2001 From: joshua-berri Date: Thu, 8 Oct 2026 10:18:25 -0700 Subject: [PATCH] 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> --- .github/workflows/test-unit.yml | 10 + .../proxy/_experimental/mcp_server/README.md | 17 ++ .../mcp_server/elicitation_handler.py | 192 ++++++------------ .../mcp_server/legacy_callbacks.py | 16 +- .../mcp_server/mcp_server_manager.py | 10 +- .../test_mcp_elicitation_handler.py | 177 +++++++++++----- .../mcp_server/test_mcp_server_manager.py | 24 ++- 7 files changed, 249 insertions(+), 197 deletions(-) diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 54fdc6b43a2..ce4f5b52f4e 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -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: >- diff --git a/litellm/proxy/_experimental/mcp_server/README.md b/litellm/proxy/_experimental/mcp_server/README.md index 1123a9779ac..7d1f58965b0 100644 --- a/litellm/proxy/_experimental/mcp_server/README.md +++ b/litellm/proxy/_experimental/mcp_server/README.md @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py index 57d2d86d506..245a0612435 100644 --- a/litellm/proxy/_experimental/mcp_server/elicitation_handler.py +++ b/litellm/proxy/_experimental/mcp_server/elicitation_handler.py @@ -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, + ) diff --git a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py index f52d2a006d2..13c3d265429 100644 --- a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py +++ b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c42ab60c48b..217e373030d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py index a59b02ec01d..dd8dddf106e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_elicitation_handler.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 5731738f137..c7b477a6cc6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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)