mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
feat(mcp): share compatibility-aware result conversion across tool surfaces (#43089)
* feat(mcp): share compatibility-aware result conversion across tool surfaces Adds one converter that turns text, JSON, SDK results, interim InputRequiredResult values and exceptions into a CallToolResult shaped for the negotiated MCP revision. Legacy revisions keep object-only structuredContent with a lossless text fallback for other JSON values, and reject interim results through failure accounting. Modern revisions pass arbitrary structuredContent and InputRequiredResult through without completed-success accounting or post-call hooks. OpenAPI tools keep the upstream body verbatim and gain structuredContent Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): accept the wire compat argument in local-registry fakes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(mcp): read tagged outcome fields directly in the result converter Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(mcp): keep the mutable-ok marker on the list literal it suppresses Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci(mcp): rerun the mcp-integration shard after a tcp cancellation timeout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): cover result conversion boundaries and explicit returns * fix(mcp): keep SSE connections on the legacy protocol --------- Co-authored-by: joshua <joshua@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
c19ce71bcd
commit
61953318bf
19 changed files with 985 additions and 122 deletions
|
|
@ -36,6 +36,7 @@ from mcp.types import (
|
|||
REQUEST_TIMEOUT,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
InputRequiredResult,
|
||||
ListPromptsResult,
|
||||
ListResourcesResult,
|
||||
ListResourceTemplatesResult,
|
||||
|
|
@ -44,7 +45,6 @@ from mcp.types import (
|
|||
Prompt,
|
||||
ResourceTemplate,
|
||||
ServerNotification,
|
||||
TextContent,
|
||||
)
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
|
|
@ -61,6 +61,7 @@ from litellm.constants import (
|
|||
from litellm.experimental_mcp_client.tools import list_tools_with_pagination
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCPAuth,
|
||||
|
|
@ -828,17 +829,15 @@ class MCPClient:
|
|||
@staticmethod
|
||||
def error_tool_result(exc: Exception) -> MCPCallToolResult:
|
||||
"""The error result ``call_tool`` returns when it swallows a failure (no re-execution)."""
|
||||
return MCPCallToolResult(
|
||||
content=[TextContent(type="text", text=f"{type(exc).__name__}: {exc}")],
|
||||
is_error=True,
|
||||
)
|
||||
return error_text_result(exc)
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
call_tool_request_params: MCPCallToolRequestParams,
|
||||
host_progress_callback: Callable | None = None,
|
||||
raise_on_error: bool = False,
|
||||
) -> MCPCallToolResult:
|
||||
allow_input_required: bool = False,
|
||||
) -> MCPCallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call an MCP Tool.
|
||||
|
||||
|
|
@ -847,6 +846,9 @@ class MCPClient:
|
|||
``isError=True`` result. The token-exchange (OBO) tool-call path uses this to detect
|
||||
an upstream 401 so it can re-mint the exchanged token and retry once; every other
|
||||
caller keeps the default and gets graceful ``isError`` degradation.
|
||||
allow_input_required: When True, a 2026-07-28 upstream may answer with an interim
|
||||
``InputRequiredResult`` and it is returned as is. The SDK rejects it otherwise, so
|
||||
callers only opt in when the downstream side can carry it.
|
||||
"""
|
||||
verbose_logger.info("MCP client calling tool '%s'", call_tool_request_params.name)
|
||||
|
||||
|
|
@ -869,6 +871,7 @@ class MCPClient:
|
|||
name=call_tool_request_params.name,
|
||||
arguments=call_tool_request_params.arguments,
|
||||
progress_callback=on_progress,
|
||||
allow_input_required=allow_input_required,
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from datetime import datetime
|
|||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
|
@ -26,6 +27,7 @@ class OperationContext:
|
|||
raw_headers: Mapping[str, str] | None = field(default=None, repr=False)
|
||||
client_ip: str | None = None
|
||||
mcp_proxy_mode: bool = False
|
||||
wire_compat: WireCompat = WireCompat.LEGACY
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "_caller", copy_caller(self._caller))
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from mcp.types import (
|
|||
CallToolResult,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
InputRequiredResult,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
)
|
||||
|
|
@ -133,6 +134,12 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
|||
ServerSpec,
|
||||
TokenExchangeConfig,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import (
|
||||
WireCompat,
|
||||
complete_call_tool_result,
|
||||
handler_outcome,
|
||||
to_gateway_tool,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.sampling_handler import (
|
||||
MCP_SAMPLING_AVAILABLE,
|
||||
)
|
||||
|
|
@ -5361,16 +5368,9 @@ class MCPServerManager:
|
|||
prefix: Final = get_server_prefix(server)
|
||||
|
||||
for tool in tools:
|
||||
tool_copy = tool.model_copy(deep=True)
|
||||
|
||||
original_name = tool_copy.name
|
||||
original_name = tool.name
|
||||
prefixed_name = add_server_prefix_to_name(original_name, prefix)
|
||||
|
||||
name_to_use = prefixed_name if add_prefix else original_name
|
||||
|
||||
# Preserve all tool fields including metadata/_meta by avoiding mutation
|
||||
tool_copy.name = name_to_use
|
||||
prefixed_tools.append(tool_copy)
|
||||
prefixed_tools.append(to_gateway_tool(tool, prefixed_name if add_prefix else original_name))
|
||||
|
||||
# Register every known prefix form (alias, server_name, server_id,
|
||||
# short ID) so call_tool can resolve regardless of which form a
|
||||
|
|
@ -5547,6 +5547,7 @@ class MCPServerManager:
|
|||
server: MCPServer,
|
||||
tool_name: str,
|
||||
arguments: _ToolArguments,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call an OpenAPI tool handler directly.
|
||||
|
|
@ -5586,14 +5587,7 @@ class MCPServerManager:
|
|||
# Call the tool handler with the arguments
|
||||
# The handler is an async function that makes the HTTP request
|
||||
handler_result: Final = await tool.handler(**arguments)
|
||||
|
||||
# Convert the handler result (string response) to CallToolResult format
|
||||
result: Final = CallToolResult(
|
||||
content=[TextContent(type="text", text=str(handler_result))],
|
||||
is_error=False,
|
||||
)
|
||||
|
||||
return result
|
||||
return complete_call_tool_result(handler_outcome(handler_result), wire_compat)
|
||||
|
||||
except MCPUpstreamAuthError:
|
||||
# The caller must re-authenticate upstream, so this keeps its type all the way to the
|
||||
|
|
@ -5820,7 +5814,8 @@ class MCPServerManager:
|
|||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
) -> CallToolResult:
|
||||
allow_input_required: bool = False,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry.
|
||||
|
||||
The exchanged token is baked into the client at build time, so the retry invalidates the
|
||||
|
|
@ -5830,7 +5825,10 @@ class MCPServerManager:
|
|||
"""
|
||||
try:
|
||||
return await client.call_tool(
|
||||
call_tool_params, host_progress_callback=host_progress_callback, raise_on_error=True
|
||||
call_tool_params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
raise_on_error=True,
|
||||
allow_input_required=allow_input_required,
|
||||
)
|
||||
except Exception as exc:
|
||||
if _extract_upstream_auth_failure(exc) is None:
|
||||
|
|
@ -5848,7 +5846,11 @@ class MCPServerManager:
|
|||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback)
|
||||
return await retry_client.call_tool(
|
||||
call_tool_params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
allow_input_required=allow_input_required,
|
||||
)
|
||||
|
||||
async def _call_regular_mcp_tool(
|
||||
self,
|
||||
|
|
@ -5865,7 +5867,8 @@ class MCPServerManager:
|
|||
hook_extra_headers: dict[str, str] | None = None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
client_ip: str | None = None,
|
||||
) -> CallToolResult:
|
||||
allow_input_required: bool = False,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call a regular MCP tool using the MCP client.
|
||||
|
||||
|
|
@ -6036,6 +6039,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
allow_input_required=allow_input_required,
|
||||
)
|
||||
|
||||
tool_call_coro = _obo_call_tool_limited()
|
||||
|
|
@ -6049,7 +6053,11 @@ class MCPServerManager:
|
|||
async def _call_tool_via_client(client, params):
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
if not relays_upstream_auth:
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
return await client.call_tool(
|
||||
params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
allow_input_required=allow_input_required,
|
||||
)
|
||||
# The client-forwarded modes carry the caller's own upstream token, so an upstream
|
||||
# 401 (expired/invalid token) is the caller's to resolve: relay it as
|
||||
# MCPUpstreamAuthError so single-server REST callers turn it into a 401 +
|
||||
|
|
@ -6061,7 +6069,10 @@ class MCPServerManager:
|
|||
# the same isError degradation the default path produces.
|
||||
try:
|
||||
return await client.call_tool(
|
||||
params, host_progress_callback=host_progress_callback, raise_on_error=True
|
||||
params,
|
||||
host_progress_callback=host_progress_callback,
|
||||
raise_on_error=True,
|
||||
allow_input_required=allow_input_required,
|
||||
)
|
||||
except Exception as e:
|
||||
auth_info: Final = _extract_upstream_auth_failure(e)
|
||||
|
|
@ -6114,7 +6125,7 @@ class MCPServerManager:
|
|||
result: Final = mcp_responses[result_index]
|
||||
self._remember_upstream_initialize_instructions(mcp_server, client)
|
||||
|
||||
return cast(CallToolResult, result)
|
||||
return cast("CallToolResult | InputRequiredResult", result)
|
||||
|
||||
def _resolve_mcp_server_for_tool_call(
|
||||
self,
|
||||
|
|
@ -6318,7 +6329,8 @@ class MCPServerManager:
|
|||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
) -> CallToolResult:
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
|
||||
|
|
@ -6427,7 +6439,7 @@ class MCPServerManager:
|
|||
resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments, wire_compat)
|
||||
finally:
|
||||
_request_auth_header.reset(auth_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
|
|
@ -6449,6 +6461,7 @@ class MCPServerManager:
|
|||
host_progress_callback=host_progress_callback,
|
||||
hook_extra_headers=hook_result.get("extra_headers"),
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
allow_input_required=wire_compat is WireCompat.MODERN,
|
||||
)
|
||||
|
||||
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)
|
||||
|
|
|
|||
|
|
@ -52,6 +52,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
header_value,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import JsonResult, TextResult, parse_http_body
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
|
@ -497,7 +498,7 @@ def create_tool_function(
|
|||
path_params, query_params, body_params = extract_parameters(operation)
|
||||
original_method: Final = method.lower()
|
||||
|
||||
async def tool_function(**kwargs: object) -> str:
|
||||
async def tool_function(**kwargs: object) -> TextResult | JsonResult:
|
||||
"""
|
||||
Dynamically generated tool function.
|
||||
|
||||
|
|
@ -531,7 +532,7 @@ def create_tool_function(
|
|||
# Sanitize and encode path parameter to prevent traversal attacks
|
||||
safe_value = _sanitize_path_parameter_value(param_value, param_name)
|
||||
except ValueError as exc:
|
||||
return "Invalid path parameter: " + str(exc)
|
||||
return TextResult("Invalid path parameter: " + str(exc))
|
||||
# Replace {param_name} or {{param_name}} in URL
|
||||
url = url.replace("{" + param_name + "}", safe_value)
|
||||
url = url.replace("{{" + param_name + "}}", safe_value)
|
||||
|
|
@ -580,7 +581,7 @@ def create_tool_function(
|
|||
elif original_method == "patch":
|
||||
response = await client.patch(url, params=params, json=json_body, headers=effective_headers)
|
||||
else:
|
||||
return f"Unsupported HTTP method: {original_method}"
|
||||
return TextResult(f"Unsupported HTTP method: {original_method}")
|
||||
except MaskedHTTPStatusError as e:
|
||||
_raise_for_upstream_failure(e.response, upstream, relays_upstream_auth)
|
||||
raise
|
||||
|
|
@ -588,7 +589,7 @@ def create_tool_function(
|
|||
_request_upstream_url.reset(url_token)
|
||||
|
||||
_raise_for_upstream_failure(response, upstream, relays_upstream_auth)
|
||||
return response.text
|
||||
return parse_http_body(response.text)
|
||||
|
||||
return tool_function
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from mcp.types import (
|
|||
GetPromptRequest,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
InputRequiredResult,
|
||||
ListPromptsRequest,
|
||||
ListPromptsResult,
|
||||
ListResourcesRequest,
|
||||
|
|
@ -85,6 +86,12 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
|||
_request_extra_headers,
|
||||
_request_resolved_auth_headers,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import (
|
||||
WireCompat,
|
||||
complete_call_tool_result,
|
||||
handler_outcome,
|
||||
to_call_tool_result,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
|
@ -1804,8 +1811,9 @@ async def execute_mcp_tool(
|
|||
host_progress_callback: ProgressCallback | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
**kwargs: object, # kwargs-ok: preserves the existing REST and decorated logging call contract
|
||||
) -> CallToolResult:
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
context: Final = prepare_context(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
@ -1813,6 +1821,7 @@ async def execute_mcp_tool(
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
wire_compat=wire_compat,
|
||||
)
|
||||
operation: Final = AuthorizedToolCall(
|
||||
name=name,
|
||||
|
|
@ -1839,8 +1848,9 @@ async def _execute_mcp_tool(
|
|||
host_progress_callback: ProgressCallback | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
**kwargs: Any,
|
||||
) -> CallToolResult:
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Execute MCP tool.
|
||||
|
||||
|
|
@ -2088,7 +2098,7 @@ async def _execute_mcp_tool(
|
|||
_extra_token: Final = _request_extra_headers.set(forwarded_headers)
|
||||
_resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers)
|
||||
try:
|
||||
response = await _handle_local_mcp_tool(name, arguments)
|
||||
response = await _handle_local_mcp_tool(name, arguments, wire_compat)
|
||||
finally:
|
||||
_request_auth_header.reset(_auth_token)
|
||||
_request_extra_headers.reset(_extra_token)
|
||||
|
|
@ -2112,6 +2122,7 @@ async def _execute_mcp_tool(
|
|||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
host_progress_callback=host_progress_callback,
|
||||
wire_compat=wire_compat,
|
||||
)
|
||||
|
||||
# Fall back to local tool registry with original name (legacy support)
|
||||
|
|
@ -2169,10 +2180,13 @@ async def _execute_mcp_tool(
|
|||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args
|
||||
|
||||
response = await _handle_local_mcp_tool(original_tool_name, arguments)
|
||||
response = await _handle_local_mcp_tool(original_tool_name, arguments, wire_compat)
|
||||
|
||||
converted: Final = to_call_tool_result(response, wire_compat)
|
||||
if isinstance(converted, InputRequiredResult):
|
||||
return converted
|
||||
return await _run_post_mcp_call_guardrails(
|
||||
result=response,
|
||||
result=converted,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
request_data=kwargs,
|
||||
|
|
@ -2206,6 +2220,13 @@ async def _run_post_mcp_call_guardrails(
|
|||
)
|
||||
|
||||
|
||||
def suppress_completed_success_logging(logging_obj: LiteLLMLoggingObj) -> None:
|
||||
"""An interim ``InputRequiredResult`` is not a completed call, so the ``@client`` wrapper
|
||||
on ``call_mcp_tool`` must not run the success handlers for it when the coroutine returns."""
|
||||
logging_obj.has_run_logging(event_type="sync_success")
|
||||
logging_obj.has_run_logging(event_type="async_success")
|
||||
|
||||
|
||||
async def _fire_mcp_tool_call_logging(
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
result: CallToolResult,
|
||||
|
|
@ -2322,10 +2343,14 @@ async def call_mcp_tool(
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
**kwargs: Any,
|
||||
) -> CallToolResult:
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""
|
||||
Call a specific tool with the provided arguments (handles prefixed tool names).
|
||||
|
||||
A modern ``InputRequiredResult`` is an interim answer, so it is returned as is and skips the
|
||||
completed-call logging below.
|
||||
"""
|
||||
start_time: Final = datetime.now()
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None)
|
||||
|
|
@ -2376,12 +2401,17 @@ async def call_mcp_tool(
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
wire_compat=wire_compat,
|
||||
**kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
await fire_mcp_tool_call_failure_logging(litellm_logging_obj, e, start_time, user_api_key_auth, kwargs)
|
||||
raise
|
||||
|
||||
if isinstance(response, InputRequiredResult):
|
||||
if litellm_logging_obj:
|
||||
suppress_completed_success_logging(litellm_logging_obj)
|
||||
return response
|
||||
if litellm_logging_obj:
|
||||
response = await _fire_mcp_tool_call_logging(
|
||||
logging_obj=litellm_logging_obj,
|
||||
|
|
@ -2547,7 +2577,8 @@ async def _handle_managed_mcp_tool(
|
|||
host_progress_callback: ProgressCallback | None = None,
|
||||
guardrail_context: Mapping[str, object] | None = None,
|
||||
client_ip: str | None = None,
|
||||
) -> CallToolResult:
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
"""Handle tool execution for managed server tools"""
|
||||
# Import here to avoid circular import
|
||||
from litellm.proxy.proxy_server import proxy_logging_obj
|
||||
|
|
@ -2566,12 +2597,15 @@ async def _handle_managed_mcp_tool(
|
|||
host_progress_callback=host_progress_callback,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
guardrail_context=guardrail_context,
|
||||
wire_compat=wire_compat,
|
||||
)
|
||||
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
|
||||
return call_tool_result
|
||||
|
||||
|
||||
async def _handle_local_mcp_tool(name: str, arguments: dict[str, object]) -> CallToolResult:
|
||||
async def _handle_local_mcp_tool(
|
||||
name: str, arguments: dict[str, object], wire_compat: WireCompat = WireCompat.LEGACY
|
||||
) -> CallToolResult:
|
||||
"""Execute a local-registry tool and report whether it succeeded.
|
||||
|
||||
Returns the result rather than bare content because the verdict is part of it: the content
|
||||
|
|
@ -2604,10 +2638,7 @@ async def _handle_local_mcp_tool(name: str, arguments: dict[str, object]) -> Cal
|
|||
content=[TextContent(text=f"Error: {e}", type="text")], # mutable-ok: MCP result content
|
||||
is_error=True,
|
||||
)
|
||||
return CallToolResult(
|
||||
content=[TextContent(text=str(result), type="text")], # mutable-ok: MCP result content
|
||||
is_error=False,
|
||||
)
|
||||
return complete_call_tool_result(handler_outcome(result), wire_compat)
|
||||
|
||||
|
||||
_MCP_CREDENTIAL_REQUEST_FIELDS: Final = frozenset(
|
||||
|
|
@ -2694,7 +2725,7 @@ async def _execute_handle_list_tools(
|
|||
|
||||
async def _execute_mcp_server_tool_call(
|
||||
context: OperationContext, params: CallToolRequestParams, host_progress_callback: ProgressCallback | None = None
|
||||
) -> CallToolResult:
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
from mcp.types import CallToolResult
|
||||
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
|
|
@ -2778,6 +2809,7 @@ async def _execute_mcp_server_tool_call(
|
|||
raw_headers=raw_headers,
|
||||
client_ip=_client_ip,
|
||||
host_progress_callback=host_progress_callback,
|
||||
wire_compat=context.wire_compat,
|
||||
**data, # for logging
|
||||
)
|
||||
except MCPMissingUserEnvVarsError as e:
|
||||
|
|
@ -3032,6 +3064,7 @@ def prepare_context(
|
|||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
) -> OperationContext:
|
||||
return OperationContext(
|
||||
_caller=user_api_key_auth,
|
||||
|
|
@ -3042,6 +3075,7 @@ def prepare_context(
|
|||
raw_headers=raw_headers,
|
||||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
wire_compat=wire_compat,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -3058,6 +3092,7 @@ GatewayOperation: TypeAlias = (
|
|||
GatewayResult: TypeAlias = (
|
||||
ListToolsResult
|
||||
| CallToolResult
|
||||
| InputRequiredResult
|
||||
| ListPromptsResult
|
||||
| GetPromptResult
|
||||
| ListResourcesResult
|
||||
|
|
@ -3071,13 +3106,17 @@ class GatewayOperations:
|
|||
self._host_progress_callback = host_progress_callback
|
||||
|
||||
@overload
|
||||
async def execute(self, operation: AuthorizedToolCall, context: OperationContext) -> CallToolResult: ...
|
||||
async def execute(
|
||||
self, operation: AuthorizedToolCall, context: OperationContext
|
||||
) -> CallToolResult | InputRequiredResult: ...
|
||||
|
||||
@overload
|
||||
async def execute(self, operation: ListToolsRequest, context: OperationContext) -> ListToolsResult: ...
|
||||
|
||||
@overload
|
||||
async def execute(self, operation: CallToolRequest, context: OperationContext) -> CallToolResult: ...
|
||||
async def execute(
|
||||
self, operation: CallToolRequest, context: OperationContext
|
||||
) -> CallToolResult | InputRequiredResult: ...
|
||||
|
||||
@overload
|
||||
async def execute(self, operation: ListPromptsRequest, context: OperationContext) -> ListPromptsResult: ...
|
||||
|
|
@ -3113,6 +3152,7 @@ class GatewayOperations:
|
|||
client_ip=_client_ip,
|
||||
host_progress_callback=operation.host_progress_callback,
|
||||
guardrail_context=operation.guardrail_context,
|
||||
wire_compat=context.wire_compat,
|
||||
**operation.logging_data,
|
||||
)
|
||||
case ListToolsRequest(params=params):
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat, complete_call_tool_result
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
acting_user_auth,
|
||||
build_effective_auth_contexts,
|
||||
|
|
@ -1197,7 +1198,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
# Call execute_mcp_tool directly (permission checks already done)
|
||||
_tool_start_time: Final = datetime.now()
|
||||
result: Final = await execute_mcp_tool(
|
||||
executed: Final = await execute_mcp_tool(
|
||||
name=tool_name,
|
||||
arguments=tool_arguments,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
|
|
@ -1212,6 +1213,7 @@ if MCP_AVAILABLE:
|
|||
guardrail_context=MCPRequestContext.resolve_guardrail_context(data),
|
||||
requested_server_id=canonical_server_id,
|
||||
)
|
||||
result: Final = complete_call_tool_result(executed, WireCompat.LEGACY)
|
||||
except Exception as e:
|
||||
request_data: Final = proxy_base_llm_response_processor.data
|
||||
await _safe_fire_mcp_tool_call_failure_logging(
|
||||
|
|
|
|||
120
litellm/proxy/_experimental/mcp_server/result_conversion.py
Normal file
120
litellm/proxy/_experimental/mcp_server/result_conversion.py
Normal file
|
|
@ -0,0 +1,120 @@
|
|||
"""Compatibility-aware conversion of upstream outcomes into MCP SDK results.
|
||||
|
||||
Every gateway surface that turns a tool outcome (text, JSON, an SDK result, an
|
||||
interim result, an exception) into the ``CallToolResult`` it sends downstream
|
||||
goes through ``to_call_tool_result`` so the per-revision wire rules live in one
|
||||
place. SDK 2.x serializes ``structuredContent`` as object-only on the handshake
|
||||
revisions (``2024-11-05`` .. ``2025-11-25``) and admits any JSON value, plus
|
||||
``input_required`` interim results, only on ``2026-07-28``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from mcp.types import CallToolResult, ContentBlock, InputRequiredResult, TextContent, Tool
|
||||
from typing_extensions import ReadOnly, TypedDict, assert_never
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import (
|
||||
JsonResult,
|
||||
TextResult,
|
||||
WireCompat,
|
||||
handler_outcome,
|
||||
parse_http_body,
|
||||
wire_compat_for,
|
||||
)
|
||||
|
||||
__all__ = (
|
||||
"INPUT_REQUIRED_UNSUPPORTED_MESSAGE",
|
||||
"JsonResult",
|
||||
"TextResult",
|
||||
"ToolOutcome",
|
||||
"WireCompat",
|
||||
"complete_call_tool_result",
|
||||
"error_text_result",
|
||||
"handler_outcome",
|
||||
"parse_http_body",
|
||||
"to_call_tool_result",
|
||||
"to_gateway_tool",
|
||||
"wire_compat_for",
|
||||
)
|
||||
|
||||
ToolOutcome: TypeAlias = TextResult | JsonResult | CallToolResult | InputRequiredResult | Exception
|
||||
|
||||
|
||||
class _Downgraded(TypedDict):
|
||||
structured_content: ReadOnly[None]
|
||||
content: ReadOnly[list[ContentBlock]] # mutable-ok: SDK list field
|
||||
|
||||
|
||||
class _Renamed(TypedDict):
|
||||
name: ReadOnly[str]
|
||||
|
||||
|
||||
INPUT_REQUIRED_UNSUPPORTED_MESSAGE: Final = (
|
||||
"Error: upstream tool returned an input_required interim result, which this MCP protocol revision cannot carry"
|
||||
)
|
||||
|
||||
|
||||
def error_text_result(exc: Exception) -> CallToolResult:
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text=f"{type(exc).__name__}: {exc}")], # mutable-ok: SDK list field
|
||||
is_error=True,
|
||||
)
|
||||
|
||||
|
||||
def to_call_tool_result(outcome: ToolOutcome, compat: WireCompat) -> CallToolResult | InputRequiredResult:
|
||||
match outcome:
|
||||
case TextResult():
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text=outcome.text)], # mutable-ok: SDK list field
|
||||
is_error=False,
|
||||
)
|
||||
case JsonResult():
|
||||
keep_structured: Final = compat is WireCompat.MODERN or isinstance(outcome.value, dict)
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text=outcome.original_text)], # mutable-ok: SDK list field
|
||||
is_error=False,
|
||||
structured_content=outcome.value if keep_structured else None,
|
||||
)
|
||||
case CallToolResult():
|
||||
return _downgrade_structured_content(outcome) if compat is WireCompat.LEGACY else outcome
|
||||
case InputRequiredResult():
|
||||
if compat is WireCompat.MODERN:
|
||||
return outcome
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text=INPUT_REQUIRED_UNSUPPORTED_MESSAGE)], # mutable-ok: SDK
|
||||
is_error=True,
|
||||
)
|
||||
case Exception():
|
||||
return error_text_result(outcome)
|
||||
return assert_never(outcome)
|
||||
|
||||
|
||||
def complete_call_tool_result(outcome: ToolOutcome, compat: WireCompat) -> CallToolResult:
|
||||
"""``to_call_tool_result`` for callers that can never carry an interim result."""
|
||||
converted: Final = to_call_tool_result(outcome, compat)
|
||||
if isinstance(converted, InputRequiredResult):
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text=INPUT_REQUIRED_UNSUPPORTED_MESSAGE)], # mutable-ok: SDK
|
||||
is_error=True,
|
||||
)
|
||||
return converted
|
||||
|
||||
|
||||
def _downgrade_structured_content(result: CallToolResult) -> CallToolResult:
|
||||
structured: Final = result.structured_content
|
||||
if structured is None or isinstance(structured, dict):
|
||||
return result
|
||||
fallback: Final = TextContent(type="text", text=json.dumps(structured))
|
||||
update: Final[_Downgraded] = {
|
||||
"structured_content": None,
|
||||
"content": [*result.content, fallback], # mutable-ok: SDK list field
|
||||
}
|
||||
return result.model_copy(update=update)
|
||||
|
||||
|
||||
def to_gateway_tool(tool: Tool, name: str) -> Tool:
|
||||
update: Final[_Renamed] = {"name": name}
|
||||
return tool.model_copy(deep=True, update=update)
|
||||
|
|
@ -145,6 +145,7 @@ try:
|
|||
|
||||
from mcp import ReadResourceResult, Resource
|
||||
from mcp.server import Server
|
||||
from mcp.server.runner import serve_loop
|
||||
from mcp.server.session import ServerSession as _McpServerSession
|
||||
from mcp.types import (
|
||||
BlobResourceContents,
|
||||
|
|
@ -504,6 +505,7 @@ if MCP_AVAILABLE:
|
|||
_invalidate_byok_cred_cache,
|
||||
_mcp_session_id_from_headers,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import wire_compat_for
|
||||
|
||||
try:
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
|
|
@ -516,6 +518,7 @@ if MCP_AVAILABLE:
|
|||
GetPromptRequestParams,
|
||||
Implementation,
|
||||
InitializeRequest,
|
||||
InputRequiredResult,
|
||||
ListPromptsResult,
|
||||
ListResourcesResult,
|
||||
ListResourceTemplatesResult,
|
||||
|
|
@ -818,7 +821,15 @@ if MCP_AVAILABLE:
|
|||
client_ip,
|
||||
) = await get_or_extract_auth_context()
|
||||
yield operations.prepare_context(
|
||||
auth, token, servers, server_headers, oauth_headers, headers, client_ip, _mcp_proxy_mode.get()
|
||||
auth,
|
||||
token,
|
||||
servers,
|
||||
server_headers,
|
||||
oauth_headers,
|
||||
headers,
|
||||
client_ip,
|
||||
_mcp_proxy_mode.get(),
|
||||
wire_compat_for(ctx.protocol_version),
|
||||
)
|
||||
|
||||
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
|
||||
|
|
@ -875,7 +886,9 @@ if MCP_AVAILABLE:
|
|||
_dispatch_virtual_mcp_tool,
|
||||
)
|
||||
|
||||
async def mcp_server_tool_call(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
|
||||
async def mcp_server_tool_call(
|
||||
ctx: ServerRequestContext, params: CallToolRequestParams
|
||||
) -> CallToolResult | InputRequiredResult:
|
||||
async with _legacy_operation_context(ctx, trace=True) as context:
|
||||
return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute(
|
||||
CallToolRequest(params=params), context
|
||||
|
|
@ -2384,8 +2397,17 @@ if MCP_AVAILABLE:
|
|||
scoped_server_endpoint=scoped_server_endpoint,
|
||||
is_initialize=scope.get("method") == "GET",
|
||||
):
|
||||
async with sse.connect_sse(transport_scope, receive, send) as (read_stream, write_stream):
|
||||
await server.run(read_stream, write_stream, server.create_initialization_options())
|
||||
async with (
|
||||
sse.connect_sse(transport_scope, receive, send) as (read_stream, write_stream),
|
||||
server.lifespan(server) as lifespan_state,
|
||||
):
|
||||
await serve_loop(
|
||||
server,
|
||||
read_stream,
|
||||
write_stream,
|
||||
lifespan_state=lifespan_state,
|
||||
init_options=server.create_initialization_options(),
|
||||
)
|
||||
except MCPUpstreamAuthError as e:
|
||||
# Upstream delegated auth returned 401; surface it to the client so
|
||||
# standards-compliant MCP clients trigger the upstream OAuth flow.
|
||||
|
|
|
|||
56
litellm/proxy/_experimental/mcp_server/tool_outcome.py
Normal file
56
litellm/proxy/_experimental/mcp_server/tool_outcome.py
Normal file
|
|
@ -0,0 +1,56 @@
|
|||
"""SDK-free half of the result conversion boundary.
|
||||
|
||||
``openapi_to_mcp_generator`` and ``contracts`` must import without the ``mcp``
|
||||
package installed, so the compatibility enum and the tagged handler outcomes
|
||||
live here; ``result_conversion`` turns them into SDK results.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Final
|
||||
|
||||
from mcp_types.version import MODERN_PROTOCOL_VERSIONS
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
_JSON_VALUE: Final = TypeAdapter(JsonValue)
|
||||
|
||||
|
||||
class WireCompat(str, Enum):
|
||||
LEGACY = "legacy"
|
||||
MODERN = "modern"
|
||||
|
||||
|
||||
def wire_compat_for(protocol_version: str) -> WireCompat:
|
||||
return WireCompat.MODERN if protocol_version in MODERN_PROTOCOL_VERSIONS else WireCompat.LEGACY
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TextResult:
|
||||
text: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class JsonResult:
|
||||
value: JsonValue
|
||||
original_text: str
|
||||
|
||||
|
||||
def parse_http_body(body: str) -> TextResult | JsonResult:
|
||||
if not body.strip():
|
||||
return TextResult(body)
|
||||
try:
|
||||
value: Final = _JSON_VALUE.validate_json(body)
|
||||
except ValidationError:
|
||||
return TextResult(body)
|
||||
if value is None:
|
||||
return TextResult(body)
|
||||
return JsonResult(value=value, original_text=body)
|
||||
|
||||
|
||||
def handler_outcome(value: object) -> TextResult | JsonResult:
|
||||
"""Normalize what a registered tool handler returned; OpenAPI handlers already return a tagged outcome."""
|
||||
if isinstance(value, (TextResult, JsonResult)):
|
||||
return value
|
||||
return TextResult(str(value))
|
||||
|
|
@ -13,6 +13,7 @@ from typing_extensions import ReadOnly, Required, assert_never
|
|||
|
||||
import litellm
|
||||
from litellm.llms.litellm_proxy.skills.skill_search import DEFAULT_SKILL_SEARCH_TOP_K
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat, complete_call_tool_result
|
||||
from litellm.proxy.agent_endpoints.agent_search import DEFAULT_AGENT_SEARCH_TOP_K
|
||||
from litellm.proxy.common_utils.semantic_text_index import (
|
||||
Embedder,
|
||||
|
|
@ -634,7 +635,7 @@ async def handle_mcp_tool_call(
|
|||
|
||||
raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
|
||||
|
||||
return await execute_mcp_tool(
|
||||
result: Final = await execute_mcp_tool(
|
||||
name=tool_name,
|
||||
arguments=arguments,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
|
|
@ -649,3 +650,4 @@ async def handle_mcp_tool_call(
|
|||
requested_server_id=requested_server_id,
|
||||
guardrail_context=guardrail_context,
|
||||
)
|
||||
return complete_call_tool_result(result, WireCompat.LEGACY)
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from typing import Any, Dict, Optional
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from mcp.types import CallToolResult
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -409,7 +410,7 @@ class TestCallToolFlowsHookHeaders:
|
|||
manager,
|
||||
"_call_openapi_tool_handler",
|
||||
new_callable=AsyncMock,
|
||||
return_value=MagicMock(),
|
||||
return_value=CallToolResult(content=[], isError=False),
|
||||
):
|
||||
import litellm.proxy._experimental.mcp_server.mcp_server_manager as mgr_mod
|
||||
|
||||
|
|
@ -456,7 +457,7 @@ class TestCallToolFlowsHookHeaders:
|
|||
manager,
|
||||
"_call_openapi_tool_handler",
|
||||
new_callable=AsyncMock,
|
||||
return_value=MagicMock(),
|
||||
return_value=CallToolResult(content=[], isError=False),
|
||||
):
|
||||
proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
|
||||
|
|
@ -1076,9 +1077,9 @@ class TestOpenApiByokCallTool:
|
|||
user_auth = UserAPIKeyAuth(user_id="default_user_id", api_key="sk-dashboard")
|
||||
captured_auth: dict[str, Optional[str]] = {}
|
||||
|
||||
async def fake_openapi_handler(_server, _name, _arguments):
|
||||
async def fake_openapi_handler(_server, _name, _arguments, _wire_compat):
|
||||
captured_auth["value"] = _request_auth_header.get()
|
||||
return MagicMock()
|
||||
return CallToolResult(content=[], isError=False)
|
||||
|
||||
with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server):
|
||||
with patch(
|
||||
|
|
@ -1316,9 +1317,9 @@ class TestOpenApiResolvedUpstreamAuth:
|
|||
user_auth = UserAPIKeyAuth(user_id="alice", api_key="sk-user")
|
||||
captured: Dict[str, Any] = {}
|
||||
|
||||
async def fake_openapi_handler(_server, _name, _arguments):
|
||||
async def fake_openapi_handler(_server, _name, _arguments, _wire_compat):
|
||||
captured["resolved"] = _request_resolved_auth_headers.get()
|
||||
return MagicMock()
|
||||
return CallToolResult(content=[], isError=False)
|
||||
|
||||
with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server):
|
||||
with patch.object(
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
from typing import Dict, Optional
|
||||
|
||||
import pytest
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
|
|
@ -46,7 +47,7 @@ def _make_server(server_id: str, max_concurrent_requests: Optional[int]) -> MCPS
|
|||
def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyTracker):
|
||||
async def fake_create_mcp_client(server, **kwargs):
|
||||
class _ProbeClient:
|
||||
async def call_tool(self, params, host_progress_callback=None):
|
||||
async def call_tool(self, params, host_progress_callback=None, allow_input_required=False):
|
||||
tracker.enter(server.server_id)
|
||||
try:
|
||||
await asyncio.sleep(HOLD_SECONDS)
|
||||
|
|
@ -145,11 +146,11 @@ async def test_openapi_backed_server_also_respects_the_cap():
|
|||
server = _make_server("srv-openapi", max_concurrent_requests=2)
|
||||
server.spec_path = "/fake/openapi.json"
|
||||
|
||||
async def fake_openapi_handler(mcp_server, name, arguments):
|
||||
async def fake_openapi_handler(mcp_server, name, arguments, wire_compat):
|
||||
tracker.enter(mcp_server.server_id)
|
||||
try:
|
||||
await asyncio.sleep(HOLD_SECONDS)
|
||||
return "ok"
|
||||
return CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
|
||||
finally:
|
||||
tracker.exit(mcp_server.server_id)
|
||||
|
||||
|
|
|
|||
|
|
@ -12,8 +12,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp import ReadResourceResult, Resource
|
||||
from mcp.server.models import InitializationOptions
|
||||
from mcp.types import (
|
||||
INVALID_REQUEST,
|
||||
METHOD_NOT_FOUND,
|
||||
BlobResourceContents,
|
||||
CallToolResult,
|
||||
Prompt,
|
||||
|
|
@ -21,7 +23,7 @@ from mcp.types import (
|
|||
TextContent,
|
||||
TextResourceContents,
|
||||
)
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION, MODERN_PROTOCOL_VERSIONS
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
|
|
@ -6565,8 +6567,11 @@ class TestGatewayCreateInitializationOptions:
|
|||
async def connect_sse(scope, receive, send):
|
||||
yield (None, None)
|
||||
|
||||
async def record_request(read_stream, write_stream, options):
|
||||
captured["server_name"] = server.create_initialization_options().server_name
|
||||
async def record_request(
|
||||
serving_server: object, read_stream: object, write_stream: object,
|
||||
*, lifespan_state: object, init_options: InitializationOptions,
|
||||
) -> None:
|
||||
captured["server_name"] = init_options.server_name
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
|
|
@ -6612,8 +6617,8 @@ class TestGatewayCreateInitializationOptions:
|
|||
True,
|
||||
),
|
||||
patch.object(
|
||||
mcp_server.server,
|
||||
"run",
|
||||
mcp_server,
|
||||
"serve_loop",
|
||||
side_effect=record_request,
|
||||
),
|
||||
):
|
||||
|
|
@ -8021,7 +8026,7 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details():
|
|||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool",
|
||||
new=AsyncMock(return_value=[]),
|
||||
new=AsyncMock(return_value=CallToolResult(content=[], is_error=False)),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
|
||||
|
|
@ -9255,6 +9260,152 @@ async def test_call_mcp_tool_skips_failure_hook_for_upstream_auth_error():
|
|||
proxy_logging_mock.post_call_failure_hook.assert_not_awaited()
|
||||
|
||||
|
||||
def _interim_input_required_result():
|
||||
from mcp.types import InputRequiredResult
|
||||
|
||||
return InputRequiredResult.model_validate(
|
||||
{
|
||||
"resultType": "input_required",
|
||||
"inputRequests": {
|
||||
"req-1": {
|
||||
"method": "elicitation/create",
|
||||
"params": {"message": "Pick one", "requestedSchema": {"type": "object", "properties": {}}},
|
||||
}
|
||||
},
|
||||
"requestState": "state-1",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _managed_tool_returning(server, upstream_result, proxy_logging_mock):
|
||||
from litellm.proxy._experimental.mcp_server.server import global_mcp_server_manager
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[server.server_id],
|
||||
),
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server),
|
||||
patch.object(global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server),
|
||||
patch.object(global_mcp_server_manager, "server_owning_tool_name_prefix", return_value=server),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[server],
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations._list_tools_before_first_call",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.operations._handle_managed_mcp_tool",
|
||||
new_callable=AsyncMock,
|
||||
return_value=upstream_result,
|
||||
) as managed_call,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock),
|
||||
):
|
||||
yield managed_call
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_mcp_tool_legacy_interim_result_is_rejected_into_failure_accounting():
|
||||
"""An upstream input_required interim on a legacy connection cannot be carried on the wire, so it
|
||||
must come back as isError and go through the same failure accounting as any other errored call."""
|
||||
from mcp.types import CallToolResult
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import (
|
||||
INPUT_REQUIRED_UNSUPPORTED_MESSAGE,
|
||||
WireCompat,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.server import call_mcp_tool
|
||||
from litellm.proxy._types import MCPTransport, UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="server-interim",
|
||||
name="test_server",
|
||||
alias="test_server",
|
||||
server_name="test_server",
|
||||
url="https://test-server.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
mcp_info={"server_name": "test_server"},
|
||||
)
|
||||
proxy_logging_mock = _mock_mcp_proxy_logging()
|
||||
logging_obj = _mock_mcp_logging_obj()
|
||||
|
||||
with _managed_tool_returning(server, _interim_input_required_result(), proxy_logging_mock) as managed_call:
|
||||
result = await call_mcp_tool(
|
||||
name="test_server-any_tool",
|
||||
arguments={"x": 1},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"),
|
||||
litellm_logging_obj=logging_obj,
|
||||
wire_compat=WireCompat.LEGACY,
|
||||
)
|
||||
|
||||
assert managed_call.await_args.kwargs["wire_compat"] is WireCompat.LEGACY
|
||||
assert isinstance(result, CallToolResult) and result.is_error is True
|
||||
assert result.content[0].text == INPUT_REQUIRED_UNSUPPORTED_MESSAGE
|
||||
logging_obj.async_success_handler.assert_not_awaited()
|
||||
logging_obj.async_failure_handler.assert_awaited_once()
|
||||
assert str(logging_obj.async_failure_handler.await_args.args[0]) == INPUT_REQUIRED_UNSUPPORTED_MESSAGE
|
||||
proxy_logging_mock.post_call_failure_hook.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_mcp_tool_modern_interim_result_passes_through_without_completed_accounting():
|
||||
"""On a modern connection the interim result is returned with its fields intact and is neither
|
||||
logged as a completed success nor run through the post-call guardrail and success hooks."""
|
||||
from mcp.types import InputRequiredResult
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import WireCompat
|
||||
from litellm.proxy._experimental.mcp_server.server import call_mcp_tool
|
||||
from litellm.proxy._types import MCPTransport, UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="server-interim",
|
||||
name="test_server",
|
||||
alias="test_server",
|
||||
server_name="test_server",
|
||||
url="https://test-server.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
mcp_info={"server_name": "test_server"},
|
||||
)
|
||||
proxy_logging_mock = _mock_mcp_proxy_logging()
|
||||
logging_obj = _mock_mcp_logging_obj()
|
||||
interim = _interim_input_required_result()
|
||||
|
||||
with _managed_tool_returning(server, interim, proxy_logging_mock) as managed_call:
|
||||
result = await call_mcp_tool(
|
||||
name="test_server-any_tool",
|
||||
arguments={"x": 1},
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"),
|
||||
litellm_logging_obj=logging_obj,
|
||||
wire_compat=WireCompat.MODERN,
|
||||
)
|
||||
|
||||
assert managed_call.await_args.kwargs["wire_compat"] is WireCompat.MODERN
|
||||
assert isinstance(result, InputRequiredResult)
|
||||
assert result.request_state == "state-1"
|
||||
assert result.input_requests is not None and set(result.input_requests) == {"req-1"}
|
||||
logging_obj.async_success_handler.assert_not_awaited()
|
||||
logging_obj.async_failure_handler.assert_not_awaited()
|
||||
logging_obj.async_post_mcp_tool_call_hook.assert_not_awaited()
|
||||
proxy_logging_mock.post_mcp_call_hook.assert_not_awaited()
|
||||
proxy_logging_mock.post_call_failure_hook.assert_not_awaited()
|
||||
assert sorted(c.kwargs["event_type"] for c in logging_obj.has_run_logging.call_args_list) == [
|
||||
"async_success",
|
||||
"sync_success",
|
||||
], "the @client wrapper would otherwise log the interim result as a completed success on return"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aggregate_listing_reports_per_server_outcomes():
|
||||
"""A failed server must contribute a classified outcome, not just silently shrink the list:
|
||||
|
|
@ -10371,7 +10522,10 @@ async def test_tool_listing_preserves_permission_denial_when_failure_logging_fai
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("prefix,suffix", (("", ""), ("/gateway", "/")))
|
||||
async def test_legacy_sse_mount_emits_message_endpoint(prefix: str, suffix: str) -> None:
|
||||
@pytest.mark.parametrize("opening_protocol", (None, *MODERN_PROTOCOL_VERSIONS))
|
||||
async def test_legacy_sse_mount_emits_message_endpoint(
|
||||
prefix: str, suffix: str, opening_protocol: str | None,
|
||||
) -> None:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.routing import Mount
|
||||
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
|
||||
|
|
@ -10434,6 +10588,23 @@ async def test_legacy_sse_mount_emits_message_endpoint(prefix: str, suffix: str)
|
|||
await asyncio.wait_for(app(post_scope, requests.get, messages.put), 2)
|
||||
return (await messages.get())["status"]
|
||||
|
||||
if opening_protocol is not None:
|
||||
discover: Final = json.dumps({
|
||||
"jsonrpc": "2.0",
|
||||
"id": 0,
|
||||
"method": "server/discover",
|
||||
"params": {"_meta": {
|
||||
"io.modelcontextprotocol/protocolVersion": opening_protocol,
|
||||
"io.modelcontextprotocol/clientInfo": {"name": "modern-client", "version": "1"},
|
||||
"io.modelcontextprotocol/clientCapabilities": {},
|
||||
}},
|
||||
}).encode()
|
||||
assert await post(discover) == 202
|
||||
discovered_frame: Final = (await asyncio.wait_for(outgoing.get(), 2))["body"].decode()
|
||||
discovered: Final = json.loads(discovered_frame.split("data: ", 1)[1].splitlines()[0])
|
||||
assert discovered["id"] == 0
|
||||
assert discovered["error"]["code"] == METHOD_NOT_FOUND
|
||||
|
||||
initialization: Final = json.dumps(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from mcp.types import Tool as MCPTool
|
|||
from pydantic import AnyUrl, TypeAdapter
|
||||
|
||||
from litellm.constants import MCP_METADATA_TIMEOUT
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
_deserialize_json_dict,
|
||||
|
|
@ -6730,7 +6731,7 @@ class TestMCPServerManager:
|
|||
# Create mock client that tracks call_tool usage
|
||||
mock_client = AsyncMock()
|
||||
|
||||
async def mock_call_tool(params, host_progress_callback=None):
|
||||
async def mock_call_tool(params, host_progress_callback=None, allow_input_required=False):
|
||||
# Return a mock CallToolResult
|
||||
result = MagicMock(spec=CallToolResult)
|
||||
result.content = [{"type": "text", "text": "Tool executed successfully"}]
|
||||
|
|
@ -10061,7 +10062,7 @@ class _RetryFakeClient:
|
|||
self._MCPClient = MCPClient
|
||||
self.attempts = 0
|
||||
|
||||
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False):
|
||||
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False):
|
||||
self.attempts += 1
|
||||
if self._raises is not None:
|
||||
if raise_on_error:
|
||||
|
|
@ -10273,7 +10274,7 @@ class TestOBOConcurrencyLimit:
|
|||
inflight = {"current": 0, "peak": 0}
|
||||
|
||||
class _ConcurrencyRecordingClient:
|
||||
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False):
|
||||
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, allow_input_required=False):
|
||||
inflight["current"] += 1
|
||||
inflight["peak"] = max(inflight["peak"], inflight["current"])
|
||||
try:
|
||||
|
|
@ -12250,6 +12251,38 @@ class TestOpenApiHandlerRelaysUpstreamAuth:
|
|||
assert result.is_error is True
|
||||
assert "upstream returned HTTP 503" in result.content[0].text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("body", "compat", "expected_structured"),
|
||||
[
|
||||
('{"total": 1.10, "items": [ ]}', "legacy", {"total": 1.1, "items": []}),
|
||||
('{"total": 1.10, "items": [ ]}', "modern", {"total": 1.1, "items": []}),
|
||||
("[1, 2]", "legacy", None),
|
||||
("[1, 2]", "modern", [1, 2]),
|
||||
("plain text", "legacy", None),
|
||||
("plain text", "modern", None),
|
||||
],
|
||||
)
|
||||
async def test_json_bodies_keep_verbatim_text_and_gain_structured_content(self, body, compat, expected_structured):
|
||||
"""The OpenAPI arm used to stringify the response; now the text block is the upstream body
|
||||
byte for byte, exactly once, and JSON bodies carry structuredContent when the caller's revision admits it."""
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat, parse_http_body
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
manager = MCPServerManager()
|
||||
|
||||
async def handler(**_kwargs):
|
||||
return parse_http_body(body)
|
||||
|
||||
tool = MagicMock()
|
||||
tool.handler = handler
|
||||
with patch.object(global_mcp_tool_registry, "get_tool", return_value=tool):
|
||||
result = await manager._call_openapi_tool_handler(self._server(), "list_reports", {}, WireCompat(compat))
|
||||
|
||||
assert result.is_error is False
|
||||
assert [block.text for block in result.content] == [body]
|
||||
assert result.structured_content == expected_structured
|
||||
|
||||
|
||||
class TestConfigServerIdPinning:
|
||||
"""config.yaml servers may pin ``server_id`` so permission grants survive connection edits."""
|
||||
|
|
@ -14254,7 +14287,7 @@ class TestProtectedCredentialPreparation:
|
|||
caller_token: Final = _request_auth_header.set(caller)
|
||||
extra_token: Final = _request_extra_headers.set(forwarded)
|
||||
try:
|
||||
assert await tool() == "authenticated"
|
||||
assert await tool() == TextResult("authenticated")
|
||||
sent: Final = destination.calls.last.request.headers
|
||||
assert sent.get("x-api-key") == static.get("X-API-Key", (forwarded or {}).get("X-API-Key"))
|
||||
if caller:
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
|||
_request_auth_header,
|
||||
_request_extra_headers,
|
||||
_request_resolved_auth_headers,
|
||||
_request_upstream_url,
|
||||
_resolve_param_list,
|
||||
_resolve_ref,
|
||||
build_input_schema,
|
||||
|
|
@ -31,6 +32,7 @@ from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
|
|||
get_base_url,
|
||||
resolve_operation_params,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import JsonResult, TextResult
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPOpenApiUpstreamError,
|
||||
|
|
@ -40,6 +42,43 @@ from litellm.proxy._experimental.mcp_server.exceptions import (
|
|||
GET_ASYNC_CLIENT_TARGET = "litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.get_async_httpx_client"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unsupported_http_method_returns_text_without_sending_request(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function("/echo", "HEAD", {}, "https://upstream.example")
|
||||
token: Final = _request_upstream_url.set("https://outer.example/request")
|
||||
try:
|
||||
assert await tool() == TextResult("Unsupported HTTP method: head")
|
||||
assert len(respx_mock.calls) == 0
|
||||
assert _request_upstream_url.get() == "https://outer.example/request"
|
||||
finally:
|
||||
_request_upstream_url.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("body,expected", [
|
||||
(' { "ok": true }\n', JsonResult({"ok": True}, ' { "ok": true }\n')),
|
||||
(' [1, 2]\n', JsonResult([1, 2], ' [1, 2]\n')),
|
||||
('false', JsonResult(False, 'false')),
|
||||
('0', JsonResult(0, '0')),
|
||||
('""', JsonResult("", '""')),
|
||||
('null', TextResult('null')),
|
||||
('{"unfinished":', TextResult('{"unfinished":')),
|
||||
('', TextResult('')),
|
||||
])
|
||||
async def test_http_response_preserves_body_and_classifies_json(
|
||||
respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch,
|
||||
body: str, expected: TextResult | JsonResult,
|
||||
) -> None:
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
tool: Final = create_tool_function("/echo", "get", {}, "https://upstream.example")
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text=body)
|
||||
assert await tool() == expected
|
||||
assert destination.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type,value,accepted", [
|
||||
(MCPAuth.api_key, "Bearer Bearer", False), (MCPAuth.api_key, "ApiKey ApiKey", False),
|
||||
|
|
@ -63,7 +102,7 @@ async def test_authorization_validates_credentials_before_http(
|
|||
caller_token: Final = _request_auth_header.set(value)
|
||||
try:
|
||||
if accepted:
|
||||
assert await tool() == "authenticated"
|
||||
assert await tool() == TextResult("authenticated")
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["authorization"] == value
|
||||
else:
|
||||
|
|
@ -103,7 +142,7 @@ async def test_static_auth_validates_headers_after_existing_precedence(
|
|||
assert exc.value.status_code == 500
|
||||
assert destination.call_count == 0
|
||||
else:
|
||||
assert await tool() == "authenticated"
|
||||
assert await tool() == TextResult("authenticated")
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["authorization"] == expected
|
||||
finally:
|
||||
|
|
@ -124,7 +163,7 @@ async def test_static_auth_uses_configured_custom_header(
|
|||
)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
if credential:
|
||||
assert await tool() == "authenticated"
|
||||
assert await tool() == TextResult("authenticated")
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers["x-custom"] == credential
|
||||
else:
|
||||
|
|
@ -144,7 +183,7 @@ async def test_static_auth_accepts_api_key_carried_by_static_header(
|
|||
)
|
||||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated")
|
||||
if credential:
|
||||
assert await tool() == "authenticated"
|
||||
assert await tool() == TextResult("authenticated")
|
||||
assert destination.calls.last.request.headers["apikey"] == credential
|
||||
assert "x-api-key" not in destination.calls.last.request.headers
|
||||
else:
|
||||
|
|
@ -167,7 +206,7 @@ async def test_static_validation_preserves_no_auth_and_resolved_oauth(
|
|||
destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="echo")
|
||||
token: Final = _request_resolved_auth_headers.set(resolved)
|
||||
try:
|
||||
assert await tool() == "echo"
|
||||
assert await tool() == TextResult("echo")
|
||||
assert destination.call_count == 1
|
||||
assert destination.calls.last.request.headers.get("authorization") == (resolved or {}).get("Authorization")
|
||||
finally:
|
||||
|
|
@ -220,7 +259,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"repository-id": "test-repo"})
|
||||
assert result == '{"id": "123"}'
|
||||
assert result == JsonResult({"id": "123"}, '{"id": "123"}')
|
||||
|
||||
# Verify URL was constructed correctly
|
||||
call_args = async_client.get.call_args
|
||||
|
|
@ -256,7 +295,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"2fa-code": "123456"})
|
||||
assert result == "verified"
|
||||
assert result == TextResult("verified")
|
||||
|
||||
# Verify query parameter was included
|
||||
call_args = async_client.post.call_args
|
||||
|
|
@ -290,7 +329,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"user.name": "john.doe"})
|
||||
assert result == "found"
|
||||
assert result == TextResult("found")
|
||||
|
||||
call_args = async_client.get.call_args
|
||||
assert call_args[1]["params"]["user.name"] == "john.doe"
|
||||
|
|
@ -323,7 +362,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"$filter": "name eq 'test'"})
|
||||
assert result == "[]"
|
||||
assert result == JsonResult([], "[]")
|
||||
|
||||
call_args = async_client.get.call_args
|
||||
assert call_args[1]["params"]["$filter"] == "name eq 'test'"
|
||||
|
|
@ -356,7 +395,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"class": "premium"})
|
||||
assert result == "items"
|
||||
assert result == TextResult("items")
|
||||
|
||||
call_args = async_client.get.call_args
|
||||
assert call_args[1]["params"]["class"] == "premium"
|
||||
|
|
@ -407,7 +446,7 @@ class TestCreateToolFunction:
|
|||
"$filter": "active",
|
||||
}
|
||||
)
|
||||
assert result == "success"
|
||||
assert result == TextResult("success")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_body_parameter(self):
|
||||
|
|
@ -440,7 +479,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"body": {"name": "test"}})
|
||||
assert result == "created"
|
||||
assert result == TextResult("created")
|
||||
|
||||
call_args = async_client.post.call_args
|
||||
assert call_args[1]["json"] == {"name": "test"}
|
||||
|
|
@ -464,7 +503,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func()
|
||||
assert result == "ok"
|
||||
assert result == TextResult("ok")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_http_methods(self):
|
||||
|
|
@ -497,7 +536,7 @@ class TestCreateToolFunction:
|
|||
mock_client.return_value = async_client
|
||||
|
||||
result = await func(**{"repository-id": "test"})
|
||||
assert result == "success"
|
||||
assert result == TextResult("success")
|
||||
|
||||
def test_no_exec_usage(self):
|
||||
"""Verify that create_tool_function does not use exec()."""
|
||||
|
|
@ -614,7 +653,7 @@ class TestPathSecurity:
|
|||
|
||||
response = await tool_function(**{"filename": "../admin"})
|
||||
|
||||
assert "Invalid path parameter" in response
|
||||
assert isinstance(response, TextResult) and "Invalid path parameter" in response.text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_should_encode_and_request_safe_path_parameters(self):
|
||||
|
|
@ -643,7 +682,7 @@ class TestPathSecurity:
|
|||
|
||||
response = await tool_function(**{"filename": "report 2024.json"})
|
||||
|
||||
assert response == "dummy-response"
|
||||
assert response == TextResult("dummy-response")
|
||||
|
||||
# Verify URL was properly encoded
|
||||
call_args = async_client.get.call_args
|
||||
|
|
@ -1181,7 +1220,7 @@ class TestRequestExtraHeaders:
|
|||
finally:
|
||||
_request_extra_headers.reset(token)
|
||||
|
||||
assert result == "ok"
|
||||
assert result == TextResult("ok")
|
||||
call_args = async_client.get.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent.get("X-TOKEN") == "secret-value"
|
||||
|
|
@ -1204,7 +1243,7 @@ class TestRequestExtraHeaders:
|
|||
|
||||
result = await func()
|
||||
|
||||
assert result == "ok"
|
||||
assert result == TextResult("ok")
|
||||
call_args = async_client.get.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent == {"X-Static": "static-value"}
|
||||
|
|
@ -1232,7 +1271,7 @@ class TestRequestExtraHeaders:
|
|||
finally:
|
||||
_request_extra_headers.reset(token)
|
||||
|
||||
assert result == "created"
|
||||
assert result == TextResult("created")
|
||||
call_args = async_client.post.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent.get("X-Static") == "static-value"
|
||||
|
|
@ -1260,7 +1299,7 @@ class TestRequestExtraHeaders:
|
|||
finally:
|
||||
_request_extra_headers.reset(token)
|
||||
|
||||
assert result == "ok"
|
||||
assert result == TextResult("ok")
|
||||
call_args = async_client.get.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent.get("X-Tenant") == "operator-tenant"
|
||||
|
|
@ -1288,7 +1327,7 @@ class TestRequestExtraHeaders:
|
|||
finally:
|
||||
_request_extra_headers.reset(token)
|
||||
|
||||
assert result == "ok"
|
||||
assert result == TextResult("ok")
|
||||
call_args = async_client.get.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent.get("X-Tenant") == "operator-tenant"
|
||||
|
|
@ -1320,7 +1359,7 @@ class TestRequestExtraHeaders:
|
|||
_request_auth_header.reset(auth_token)
|
||||
_request_extra_headers.reset(extra_token)
|
||||
|
||||
assert result == "secure-data"
|
||||
assert result == TextResult("secure-data")
|
||||
call_args = async_client.get.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent.get("Authorization") == "Bearer byok-credential"
|
||||
|
|
@ -1380,7 +1419,7 @@ class TestRequestExtraHeaders:
|
|||
_request_extra_headers.reset(extra_token)
|
||||
_request_resolved_auth_headers.reset(resolved_token)
|
||||
|
||||
assert result == "secure-data"
|
||||
assert result == TextResult("secure-data")
|
||||
headers_sent = async_client.get.call_args[1]["headers"]
|
||||
authorization_values = [v for k, v in headers_sent.items() if k.lower() == "authorization"]
|
||||
assert authorization_values == ["Bearer resolved-oauth"]
|
||||
|
|
@ -1436,7 +1475,7 @@ class TestUpstreamStatusIsClassified:
|
|||
async def test_success_still_returns_the_body(self):
|
||||
tool, client = self._tool(200, text='{"reports": []}')
|
||||
with patch(GET_ASYNC_CLIENT_TARGET, return_value=client):
|
||||
assert await tool() == '{"reports": []}'
|
||||
assert await tool() == JsonResult({"reports": []}, '{"reports": []}')
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_401_raises_the_reauth_signal_carrying_the_challenge(self):
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from datetime import datetime, timezone
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from mcp.types import CallToolResult
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
|
|
@ -46,7 +47,7 @@ async def test_openapi_local_tool_runs_pre_call_tool_check():
|
|||
fake_tool.name = "list_pets"
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=[])
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
|
|
@ -131,7 +132,7 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises():
|
|||
pre_call = AsyncMock(
|
||||
side_effect=HTTPException(status_code=403, detail="not allowed")
|
||||
)
|
||||
handle_local = AsyncMock(return_value=[])
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
|
|
@ -191,7 +192,7 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable():
|
|||
fake_tool.name = "list_pets"
|
||||
|
||||
pre_call = AsyncMock(return_value={})
|
||||
handle_local = AsyncMock(return_value=[])
|
||||
handle_local = AsyncMock(return_value=CallToolResult(content=[], is_error=False))
|
||||
resolve_auth = MagicMock()
|
||||
|
||||
# `_get_mcp_server_from_tool_name` returns None — no server context.
|
||||
|
|
@ -275,9 +276,9 @@ async def test_openapi_local_tool_injects_resolved_oauth_token():
|
|||
fake_tool.name = "get_values"
|
||||
captured: dict = {}
|
||||
|
||||
async def handle_local(_name, _arguments):
|
||||
async def handle_local(_name, _arguments, _wire_compat):
|
||||
captured["resolved"] = _request_resolved_auth_headers.get()
|
||||
return []
|
||||
return CallToolResult(content=[], is_error=False)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
|
|
@ -603,13 +604,13 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc
|
|||
captured["resolver_credential"] = kwargs["mcp_auth_header"]
|
||||
return None, kwargs["forwarded_headers"]
|
||||
|
||||
async def capture_local(_name, _arguments):
|
||||
async def capture_local(_name, _arguments, _wire_compat):
|
||||
captured["injected"] = _request_auth_header.get()
|
||||
return []
|
||||
return CallToolResult(content=[], is_error=False)
|
||||
|
||||
async def capture_openapi_handler(_server, _name, _arguments):
|
||||
async def capture_openapi_handler(_server, _name, _arguments, _wire_compat):
|
||||
captured["injected"] = _request_auth_header.get()
|
||||
return []
|
||||
return CallToolResult(content=[], is_error=False)
|
||||
|
||||
manager = mcp_operations.global_mcp_server_manager
|
||||
with (
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(capl
|
|||
@pytest.mark.asyncio
|
||||
async def test_dispatch_uses_explicit_context_when_ambient_caller_differs():
|
||||
from mcp.server.auth.middleware.auth_context import auth_context_var
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.server import set_auth_context
|
||||
|
||||
context = prepare_context(
|
||||
|
|
@ -64,12 +65,13 @@ async def test_dispatch_uses_explicit_context_when_ambient_caller_differs():
|
|||
@pytest.mark.asyncio
|
||||
async def test_legacy_adapter_cleans_context_after_cancelled_operation():
|
||||
from types import SimpleNamespace
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
|
||||
|
||||
previous_session = server.active_mcp_session_var.get()
|
||||
previous_request = active_mcp_request_ctx_var.get()
|
||||
request = SimpleNamespace(session=object())
|
||||
request = SimpleNamespace(session=object(), protocol_version="2025-06-18")
|
||||
auth = (None, None, None, None, None, None, None)
|
||||
|
||||
async def cancelled_operation():
|
||||
|
|
@ -90,6 +92,7 @@ async def test_legacy_adapter_cleans_context_after_cancelled_operation():
|
|||
@pytest.mark.asyncio
|
||||
async def test_legacy_adapter_cleans_context_when_trace_setup_fails():
|
||||
from types import SimpleNamespace
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var
|
||||
|
||||
|
|
@ -111,6 +114,7 @@ async def test_legacy_adapter_cleans_context_when_trace_setup_fails():
|
|||
@pytest.mark.asyncio
|
||||
async def test_prompt_sampling_receives_explicit_operation_caller_headers_and_ip():
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
|
@ -201,9 +205,11 @@ def _catalog_case(method):
|
|||
@pytest.mark.parametrize("state", ["success", "denied", "upstream_failure", "scope_failure"])
|
||||
async def test_native_catalog_operations_preserve_context_results_and_failure_policy(method, state):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from fastapi import HTTPException
|
||||
from mcp.server.context import ServerRequestContext
|
||||
from mcp.types import PaginatedRequestParams
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations, server
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
|
@ -260,6 +266,7 @@ async def test_native_catalog_operations_preserve_context_results_and_failure_po
|
|||
async def test_explicit_proxy_context_rejects_catalog_operations_before_upstream_access(method):
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import METHOD_NOT_FOUND
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
operation, _, manager_method, _, _ = _catalog_case(method)
|
||||
|
|
@ -276,6 +283,7 @@ async def test_explicit_proxy_context_rejects_catalog_operations_before_upstream
|
|||
@pytest.mark.parametrize("failure", ["missing_env", "pii", "guardrail", "unexpected"])
|
||||
async def test_tool_operation_preserves_failure_messages_and_request_trace(failure):
|
||||
from mcp.types import CallToolRequest, CallToolRequestParams
|
||||
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server.utils import MCPMissingUserEnvVarsError
|
||||
|
|
@ -333,6 +341,7 @@ async def test_catalog_operation_preserves_empty_result_for_malformed_upstream_i
|
|||
@pytest.mark.parametrize("catalog_unavailable", [False, True])
|
||||
async def test_tool_listing_returns_empty_result_without_dispatch_for_unavailable_catalog(catalog_unavailable):
|
||||
from mcp.types import ListToolsRequest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
allowed = AsyncMock(
|
||||
|
|
@ -352,6 +361,7 @@ async def test_tool_listing_returns_empty_result_without_dispatch_for_unavailabl
|
|||
@pytest.mark.asyncio
|
||||
async def test_explicit_proxy_context_lists_builtin_tools_and_blocks_direct_tool_dispatch():
|
||||
from mcp.types import CallToolRequest, CallToolRequestParams, ListToolsRequest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
context = prepare_context(mcp_proxy_mode=True)
|
||||
|
|
@ -486,8 +496,10 @@ class TestChallengeMissingTokenExchangeSubject:
|
|||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_challenges_missing_subject_before_cold_listing():
|
||||
"""On a cold catalog the challenge fires before any listing or tool resolution is attempted."""
|
||||
from fastapi import HTTPException
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
|
||||
server = _server("te-exec", MCPAuth.oauth2_token_exchange)
|
||||
|
|
@ -509,3 +521,24 @@ async def test_execute_mcp_tool_challenges_missing_subject_before_cold_listing()
|
|||
)
|
||||
assert exc_info.value.status_code == 401
|
||||
listing.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("compat", ["legacy", "modern"])
|
||||
async def test_local_tool_json_array_is_converted_once_for_the_caller_revision(compat: str) -> None:
|
||||
"""The local-registry arm used to convert at MODERN and let the legacy downgrade append a second
|
||||
text block; converting at the caller's revision keeps the upstream body exactly once."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import WireCompat, parse_http_body
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
|
||||
|
||||
body = '["a","b"]'
|
||||
tool = MagicMock()
|
||||
tool.handler = AsyncMock(return_value=parse_http_body(body))
|
||||
with patch.object(global_mcp_tool_registry, "get_tool", return_value=tool):
|
||||
result = await operations._handle_local_mcp_tool("reports-list_tags", {}, WireCompat(compat))
|
||||
|
||||
assert [block.text for block in result.content] == [body]
|
||||
assert result.structured_content == (["a", "b"] if compat == "modern" else None)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
|
|
@ -7,12 +6,15 @@ from datetime import datetime
|
|||
from typing import Any, Dict, Final, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import operations as mcp_operations
|
||||
|
||||
if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11
|
||||
from exceptiongroup import BaseExceptionGroup
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
from starlette.requests import Request
|
||||
|
||||
from litellm.constants import MCP_TOOL_LISTING_TIMEOUT
|
||||
|
|
@ -29,6 +31,8 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
_OK_TOOL_RESULT: Final = CallToolResult(content=[TextContent(type="text", text='{"result": "ok"}')], is_error=False)
|
||||
|
||||
|
||||
def _rendered_log_message(call):
|
||||
message = str(call.args[0])
|
||||
|
|
@ -1472,10 +1476,10 @@ class TestListToolsRestAPI:
|
|||
monkeypatch,
|
||||
):
|
||||
"""The REST tools/list path should include tools beyond the upstream first page."""
|
||||
import litellm.experimental_mcp_client.client as mcp_client_module
|
||||
from mcp.types import ListToolsResult, PaginatedRequestParams
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm.experimental_mcp_client.client as mcp_client_module
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
|
|
@ -2470,7 +2474,7 @@ class TestCallToolRestAPI:
|
|||
|
||||
async def fake_execute_mcp_tool(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return {"result": "ok"}
|
||||
return _OK_TOOL_RESULT
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
|
|
@ -2530,13 +2534,89 @@ class TestCallToolRestAPI:
|
|||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
assert result == {"result": "ok"}
|
||||
assert result == _OK_TOOL_RESULT
|
||||
assert captured["name"] == "demo-tool"
|
||||
assert captured["arguments"] == {"foo": "bar"}
|
||||
assert captured["allowed_mcp_servers"] == [stub_server]
|
||||
assert captured["oauth2_headers"] is None
|
||||
fire_logging.assert_awaited_once()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("structured", "expected_structured", "expected_texts"),
|
||||
[
|
||||
({"a": 1}, {"a": 1}, ['{"a": 1}']),
|
||||
([1, 2], None, ['{"a": 1}', "[1, 2]"]),
|
||||
],
|
||||
)
|
||||
async def test_rest_keeps_its_serialization_shape_with_legacy_structured_admission(
|
||||
self, monkeypatch, structured, expected_structured, expected_texts
|
||||
):
|
||||
"""REST has no negotiated revision, so it admits object structuredContent only and downgrades
|
||||
anything else losslessly, while the response keeps the SDK model shape (resultType included)
|
||||
rather than being run through the MCP legacy wire serializer."""
|
||||
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
||||
async def fake_get_allowed_mcp_servers(*args, **kwargs):
|
||||
return ["server-1"]
|
||||
|
||||
class StubServer:
|
||||
server_id = "server-1"
|
||||
alias = "server-1"
|
||||
server_name = "server-1"
|
||||
name = "stub"
|
||||
allowed_tools = None
|
||||
mcp_info = {"server_name": "stub"}
|
||||
available_on_public_internet = True
|
||||
auth_type = None
|
||||
|
||||
stub_server = StubServer()
|
||||
|
||||
async def fake_add_litellm_data_to_request(**kwargs):
|
||||
return kwargs.get("data", {})
|
||||
|
||||
async def fake_execute_mcp_tool(**kwargs):
|
||||
return CallToolResult(
|
||||
content=[TextContent(type="text", text='{"a": 1}')],
|
||||
structuredContent=structured,
|
||||
isError=False,
|
||||
)
|
||||
|
||||
async def fake_fire_logging(logging_obj, result, start_time, end_time, **kwargs):
|
||||
return result
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", fake_get_allowed_mcp_servers
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda server_id: stub_server if server_id == "server-1" else None,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.add_litellm_data_to_request", fake_add_litellm_data_to_request, raising=False
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {}, raising=False)
|
||||
monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", fake_execute_mcp_tool, raising=False)
|
||||
monkeypatch.setattr(rest_endpoints, "_fire_mcp_tool_call_logging", fake_fire_logging, raising=False)
|
||||
|
||||
request = _build_request(
|
||||
path="/mcp-rest/tools/call",
|
||||
method="POST",
|
||||
json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {}},
|
||||
)
|
||||
|
||||
result = await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=UserAPIKeyAuth())
|
||||
|
||||
assert isinstance(result, CallToolResult)
|
||||
dumped = result.model_dump(by_alias=True, mode="json", exclude_none=True)
|
||||
assert dumped.get("structuredContent") == expected_structured
|
||||
assert [block["text"] for block in dumped["content"]] == expected_texts
|
||||
assert dumped["resultType"] == "complete"
|
||||
assert dumped["isError"] is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("auth_type", "per_user_oauth", "expected"),
|
||||
|
|
@ -2580,7 +2660,7 @@ class TestCallToolRestAPI:
|
|||
|
||||
async def fake_execute_mcp_tool(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return {"result": "ok"}
|
||||
return _OK_TOOL_RESULT
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager, "get_allowed_mcp_servers", fake_get_allowed_mcp_servers
|
||||
|
|
@ -2607,7 +2687,7 @@ class TestCallToolRestAPI:
|
|||
|
||||
result = await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=UserAPIKeyAuth())
|
||||
|
||||
assert result == {"result": "ok"}
|
||||
assert result == _OK_TOOL_RESULT
|
||||
assert captured["oauth2_headers"] == expected
|
||||
assert captured["raw_headers"]["authorization"] == "Bearer user-subject-token"
|
||||
|
||||
|
|
@ -2637,7 +2717,7 @@ class TestCallToolRestAPI:
|
|||
return kwargs.get("data", {})
|
||||
|
||||
async def fake_execute_mcp_tool(**kwargs):
|
||||
return {"content": [{"type": "text", "text": "jane@example.com"}]}
|
||||
return CallToolResult(content=[TextContent(type="text", text="jane@example.com")], is_error=False)
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False)
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -2659,7 +2739,7 @@ class TestCallToolRestAPI:
|
|||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {}, raising=False)
|
||||
monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", fake_execute_mcp_tool, raising=False)
|
||||
masked_result = {"content": [{"type": "text", "text": "<EMAIL_ADDRESS>"}]}
|
||||
masked_result = CallToolResult(content=[TextContent(type="text", text="<EMAIL_ADDRESS>")], is_error=False)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"_fire_mcp_tool_call_logging",
|
||||
|
|
@ -2714,9 +2794,9 @@ class TestCallToolRestAPI:
|
|||
|
||||
async def fake_execute_mcp_tool(**kwargs):
|
||||
captured.update(kwargs)
|
||||
return {"result": "ok"}
|
||||
return _OK_TOOL_RESULT
|
||||
|
||||
fire_logging = AsyncMock(return_value={"result": "ok"})
|
||||
fire_logging = AsyncMock(return_value=_OK_TOOL_RESULT)
|
||||
monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,243 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from mcp.types import CallToolResult, ImageContent, InputRequiredResult, TextContent, Tool
|
||||
from mcp_types.methods import serialize_server_result
|
||||
from mcp_types.version import KNOWN_PROTOCOL_VERSIONS, MODERN_PROTOCOL_VERSIONS
|
||||
from pydantic import JsonValue, ValidationError
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.result_conversion import (
|
||||
INPUT_REQUIRED_UNSUPPORTED_MESSAGE,
|
||||
JsonResult,
|
||||
TextResult,
|
||||
WireCompat,
|
||||
complete_call_tool_result,
|
||||
error_text_result,
|
||||
handler_outcome,
|
||||
parse_http_body,
|
||||
to_call_tool_result,
|
||||
to_gateway_tool,
|
||||
wire_compat_for,
|
||||
)
|
||||
|
||||
BOTH: Final = (WireCompat.LEGACY, WireCompat.MODERN)
|
||||
|
||||
|
||||
def _interim() -> InputRequiredResult:
|
||||
return InputRequiredResult.model_validate(
|
||||
{
|
||||
"resultType": "input_required",
|
||||
"inputRequests": {
|
||||
"req-1": {
|
||||
"method": "elicitation/create",
|
||||
"params": {"message": "Pick one", "requestedSchema": {"type": "object", "properties": {}}},
|
||||
}
|
||||
},
|
||||
"requestState": "abc",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _wire(result: CallToolResult | InputRequiredResult, version: str) -> dict[str, object]:
|
||||
return serialize_server_result(
|
||||
"tools/call", version, result.model_dump(by_alias=True, mode="json", exclude_none=True)
|
||||
)
|
||||
|
||||
|
||||
class TestWireCompatFor:
|
||||
def test_only_modern_revisions_map_to_modern(self):
|
||||
for version in KNOWN_PROTOCOL_VERSIONS:
|
||||
expected: Final = WireCompat.MODERN if version in MODERN_PROTOCOL_VERSIONS else WireCompat.LEGACY
|
||||
assert wire_compat_for(version) is expected, version
|
||||
assert wire_compat_for("1999-01-01") is WireCompat.LEGACY
|
||||
|
||||
|
||||
class TestParseHttpBody:
|
||||
@pytest.mark.parametrize("body", ["", " ", "{not json", "null"])
|
||||
def test_non_structured_bodies_stay_text(self, body: str):
|
||||
assert parse_http_body(body) == TextResult(body)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"body, value",
|
||||
[
|
||||
('{"a": 1}', {"a": 1}),
|
||||
("[1, 2]", [1, 2]),
|
||||
("1.10", 1.1),
|
||||
("true", True),
|
||||
('"hi"', "hi"),
|
||||
],
|
||||
)
|
||||
def test_json_bodies_keep_original_text(self, body: str, value: object):
|
||||
assert parse_http_body(body) == JsonResult(value=value, original_text=body)
|
||||
|
||||
def test_handler_outcome_stringifies_unknown_values(self):
|
||||
assert handler_outcome(42) == TextResult("42")
|
||||
assert handler_outcome(TextResult("x")) == TextResult("x")
|
||||
|
||||
|
||||
class TestTextAndJsonArms:
|
||||
@pytest.mark.parametrize("compat", BOTH)
|
||||
def test_text_result(self, compat: WireCompat):
|
||||
result = to_call_tool_result(TextResult("plain"), compat)
|
||||
assert isinstance(result, CallToolResult)
|
||||
assert result.is_error is False
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == ["plain"]
|
||||
assert result.structured_content is None
|
||||
|
||||
@pytest.mark.parametrize("compat", BOTH)
|
||||
def test_json_object_is_structured_everywhere_and_text_is_verbatim(self, compat: WireCompat):
|
||||
body: Final = '{"n": 1.10,\n"k": "v"}'
|
||||
result = to_call_tool_result(parse_http_body(body), compat)
|
||||
assert isinstance(result, CallToolResult)
|
||||
assert result.structured_content == {"n": 1.1, "k": "v"}
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == [body]
|
||||
|
||||
@pytest.mark.parametrize("body", ["[1, 2]", "3", "true", '"s"'])
|
||||
def test_non_object_json_is_structured_only_on_modern(self, body: str):
|
||||
legacy = to_call_tool_result(parse_http_body(body), WireCompat.LEGACY)
|
||||
modern = to_call_tool_result(parse_http_body(body), WireCompat.MODERN)
|
||||
assert isinstance(legacy, CallToolResult) and isinstance(modern, CallToolResult)
|
||||
assert legacy.structured_content is None
|
||||
assert modern.structured_content == json.loads(body)
|
||||
for result in (legacy, modern):
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == [body]
|
||||
|
||||
@pytest.mark.parametrize("compat", BOTH)
|
||||
def test_json_null_keeps_text_and_claims_no_structured_field(self, compat: WireCompat):
|
||||
result = to_call_tool_result(parse_http_body("null"), compat)
|
||||
assert isinstance(result, CallToolResult)
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == ["null"]
|
||||
assert "structuredContent" not in _wire(result, "2026-07-28")
|
||||
|
||||
|
||||
class TestSdkResultArm:
|
||||
def _incoming(self, content: list[TextContent]) -> CallToolResult:
|
||||
return CallToolResult(content=content, structured_content=[1, 2], meta={"trace": "t1"}, is_error=False)
|
||||
|
||||
def test_modern_passes_through_the_same_object(self):
|
||||
incoming = self._incoming([])
|
||||
assert to_call_tool_result(incoming, WireCompat.MODERN) is incoming
|
||||
|
||||
def test_legacy_downgrade_with_empty_content_appends_json_text(self):
|
||||
incoming = self._incoming([])
|
||||
result = to_call_tool_result(incoming, WireCompat.LEGACY)
|
||||
assert isinstance(result, CallToolResult)
|
||||
assert result.structured_content is None
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == ["[1, 2]"]
|
||||
assert result.meta == {"trace": "t1"}
|
||||
assert incoming.structured_content == [1, 2] and incoming.content == []
|
||||
|
||||
def test_legacy_downgrade_keeps_unrelated_content_and_appends_json_text(self):
|
||||
incoming = self._incoming([TextContent(type="text", text="Done")])
|
||||
result = to_call_tool_result(incoming, WireCompat.LEGACY)
|
||||
assert isinstance(result, CallToolResult)
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == ["Done", "[1, 2]"]
|
||||
assert incoming.content == [TextContent(type="text", text="Done")]
|
||||
assert incoming.structured_content == [1, 2]
|
||||
|
||||
def test_legacy_keeps_object_structured_content(self):
|
||||
incoming = CallToolResult(content=[], structured_content={"a": 1}, is_error=False)
|
||||
assert to_call_tool_result(incoming, WireCompat.LEGACY) is incoming
|
||||
|
||||
@pytest.mark.parametrize("value", [False, 0, "", []])
|
||||
def test_legacy_downgrade_preserves_falsy_values_and_non_text_blocks(self, value: JsonValue) -> None:
|
||||
incoming: Final = CallToolResult(
|
||||
content=[
|
||||
ImageContent(type="image", data="AA==", mime_type="image/png"),
|
||||
TextContent(type="text", text="Done"),
|
||||
],
|
||||
structured_content=value,
|
||||
meta={"trace": "t1"},
|
||||
is_error=True,
|
||||
)
|
||||
before: Final = incoming.model_dump(by_alias=True)
|
||||
result: Final = to_call_tool_result(incoming, WireCompat.LEGACY)
|
||||
assert isinstance(result, CallToolResult)
|
||||
assert result.content == [*incoming.content, TextContent(type="text", text=json.dumps(value))]
|
||||
assert result.structured_content is None
|
||||
assert result.meta == incoming.meta
|
||||
assert result.is_error is True
|
||||
assert incoming.model_dump(by_alias=True) == before
|
||||
|
||||
def test_is_error_survives_downgrade(self):
|
||||
incoming = CallToolResult(content=[], structured_content=7, is_error=True)
|
||||
result = to_call_tool_result(incoming, WireCompat.LEGACY)
|
||||
assert isinstance(result, CallToolResult) and result.is_error is True
|
||||
|
||||
|
||||
class TestInterimAndExceptionArms:
|
||||
def test_modern_interim_passes_through(self):
|
||||
interim = _interim()
|
||||
assert to_call_tool_result(interim, WireCompat.MODERN) is interim
|
||||
|
||||
def test_legacy_interim_becomes_error_result(self):
|
||||
result = to_call_tool_result(_interim(), WireCompat.LEGACY)
|
||||
assert isinstance(result, CallToolResult)
|
||||
assert result.is_error is True
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == [INPUT_REQUIRED_UNSUPPORTED_MESSAGE]
|
||||
|
||||
def test_complete_call_tool_result_never_returns_interim(self):
|
||||
result = complete_call_tool_result(_interim(), WireCompat.MODERN)
|
||||
assert isinstance(result, CallToolResult) and result.is_error is True
|
||||
|
||||
@pytest.mark.parametrize("compat", BOTH)
|
||||
def test_exception_arm_matches_error_text_result(self, compat: WireCompat):
|
||||
exc = ValueError("boom")
|
||||
result = to_call_tool_result(exc, compat)
|
||||
assert result == error_text_result(exc)
|
||||
assert isinstance(result, CallToolResult) and result.is_error is True
|
||||
assert [c.text for c in result.content if isinstance(c, TextContent)] == ["ValueError: boom"]
|
||||
|
||||
|
||||
class TestSdkWireSerialization:
|
||||
@pytest.mark.parametrize("version", KNOWN_PROTOCOL_VERSIONS)
|
||||
def test_converted_results_serialize_on_their_negotiated_revision(self, version: str):
|
||||
compat = wire_compat_for(version)
|
||||
for body in ('{"a": 1}', "[1, 2]", "3", "null", "text"):
|
||||
result = to_call_tool_result(parse_http_body(body), compat)
|
||||
frame = _wire(result, version)
|
||||
assert frame["content"] == [{"type": "text", "text": body}]
|
||||
structured = json.loads(body) if body != "text" else None
|
||||
expects_structured = structured is not None and (
|
||||
compat is WireCompat.MODERN or isinstance(structured, dict)
|
||||
)
|
||||
assert ("structuredContent" in frame) is expects_structured, (version, body)
|
||||
if expects_structured:
|
||||
assert frame["structuredContent"] == structured
|
||||
assert ("resultType" in frame) is (compat is WireCompat.MODERN), (version, body)
|
||||
|
||||
@pytest.mark.parametrize("version", KNOWN_PROTOCOL_VERSIONS)
|
||||
def test_downgraded_sdk_result_serializes_where_the_raw_one_would_not(self, version: str):
|
||||
incoming = CallToolResult(content=[TextContent(type="text", text="Done")], structured_content=[1, 2])
|
||||
converted = to_call_tool_result(incoming, wire_compat_for(version))
|
||||
frame = _wire(converted, version)
|
||||
if version in MODERN_PROTOCOL_VERSIONS:
|
||||
assert frame["structuredContent"] == [1, 2]
|
||||
return
|
||||
with pytest.raises(ValidationError):
|
||||
_wire(incoming, version)
|
||||
assert "structuredContent" not in frame
|
||||
assert frame["content"] == [{"type": "text", "text": "Done"}, {"type": "text", "text": "[1, 2]"}]
|
||||
|
||||
def test_modern_interim_serializes_with_its_fields_intact(self):
|
||||
frame = _wire(_interim(), "2026-07-28")
|
||||
assert frame["resultType"] == "input_required"
|
||||
assert frame["requestState"] == "abc"
|
||||
assert frame["inputRequests"]["req-1"]["params"]["message"] == "Pick one"
|
||||
|
||||
|
||||
class TestToGatewayTool:
|
||||
def test_rename_is_a_deep_copy_that_keeps_every_other_field(self):
|
||||
tool = Tool(
|
||||
name="orig",
|
||||
description="d",
|
||||
inputSchema={"type": "object", "properties": {"q": {"type": "string"}}},
|
||||
_meta={"owner": "x"},
|
||||
)
|
||||
renamed = to_gateway_tool(tool, "srv-orig")
|
||||
assert renamed.name == "srv-orig"
|
||||
assert tool.name == "orig"
|
||||
assert renamed.input_schema == tool.input_schema and renamed.input_schema is not tool.input_schema
|
||||
assert renamed.meta == {"owner": "x"}
|
||||
assert renamed.description == "d"
|
||||
Loading…
Add table
Reference in a new issue