mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
feat(mcp/v2): route call_tool through the egress transport
Factor the connect-and-call step of _call_regular_mcp_tool into an overridable _open_and_call_tool seam, so v2 overrides only the transport (resolve() + UpstreamConnection.call_tool) while inheriting the guardrail/hook, gather, timeout, and result-extraction tail unchanged. The v1 seam builds headers, creates the MCPClient, calls the tool, and caches the upstream initialize instructions; v2 defers to it for what it does not own yet (unmapped modes, OpenAPI, the mcp_auth_header override/inbound-token path, and guardrail-injected headers such as the JWT signer) and otherwise resolves the credential and calls the tool, mapping any failure via _egress_item_failure. Header building moves behind a shared _build_egress_headers helper that _v2_connection calls: static configured headers (with per-user env-var interpolation, raising the 412 precondition on the call path) plus, on the call path only, the caller headers the server forwards. The forwarder always strips Authorization since the upstream credential is resolve()'s, never a forwarded inbound header. UpstreamConnection.call_tool gains progress-callback support for parity. Validated: 8 unit tests (the call path exercised against an in-process FastMCP echo server); proxy boots, tools-list unchanged, and a live tools/call on a none-mode server returns the upstream result through the v2 path.
This commit is contained in:
parent
fcdda93e19
commit
577ee9680a
4 changed files with 237 additions and 43 deletions
|
|
@ -3426,6 +3426,77 @@ class MCPServerManager:
|
|||
GuardrailRaisedException: If guardrails block the call
|
||||
HTTPException: If an HTTP error occurs
|
||||
"""
|
||||
tasks.append(
|
||||
asyncio.create_task(
|
||||
self._open_and_call_tool(
|
||||
mcp_server,
|
||||
original_tool_name,
|
||||
arguments,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
hook_extra_headers=hook_extra_headers,
|
||||
host_progress_callback=host_progress_callback,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
_timeout = (
|
||||
mcp_server.timeout if mcp_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
)
|
||||
try:
|
||||
mcp_responses = await asyncio.wait_for(
|
||||
asyncio.gather(*tasks), timeout=_timeout
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
raise HTTPException(
|
||||
status_code=504,
|
||||
detail={
|
||||
"error": "timeout",
|
||||
"message": f"MCP tool call timed out after {_timeout}s",
|
||||
},
|
||||
)
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
|
||||
# If proxy_logging_obj is None, the tool call result is at index 0
|
||||
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
|
||||
result_index = 1 if proxy_logging_obj else 0
|
||||
result = mcp_responses[result_index]
|
||||
|
||||
return cast(CallToolResult, result)
|
||||
|
||||
async def _open_and_call_tool(
|
||||
self,
|
||||
mcp_server: MCPServer,
|
||||
original_tool_name: str,
|
||||
arguments: Dict[str, Any],
|
||||
*,
|
||||
mcp_auth_header: Optional[str],
|
||||
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
|
||||
oauth2_headers: Optional[Dict[str, str]],
|
||||
raw_headers: Optional[Dict[str, str]],
|
||||
hook_extra_headers: Optional[Dict[str, str]],
|
||||
host_progress_callback: Optional["ProgressFnT"],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> CallToolResult:
|
||||
"""Open the upstream connection and invoke one tool, returning its result.
|
||||
|
||||
The transport seam factored out of _call_regular_mcp_tool so v2 can override only the
|
||||
connect-and-call step (resolve() + UpstreamConnection) while inheriting the
|
||||
guardrail/hook, gather, timeout, and result-extraction tail. This v1 implementation
|
||||
builds the outbound headers, creates the MCPClient, calls the tool, and caches the
|
||||
upstream's initialize instructions.
|
||||
"""
|
||||
# Get server-specific auth header if available (case-insensitive)
|
||||
# FIX: Added case-insensitive matching to handle auth header keys that may not match
|
||||
# the exact case of server alias/name (e.g., '1litellmagcgateway' vs '1LiteLLMAGCGateway')
|
||||
|
|
@ -3538,44 +3609,9 @@ class MCPServerManager:
|
|||
arguments=arguments,
|
||||
)
|
||||
|
||||
async def _call_tool_via_client(client, params):
|
||||
return await client.call_tool(
|
||||
params, host_progress_callback=host_progress_callback
|
||||
)
|
||||
|
||||
tasks.append(
|
||||
asyncio.create_task(_call_tool_via_client(client, call_tool_params))
|
||||
result = await client.call_tool(
|
||||
call_tool_params, host_progress_callback=host_progress_callback
|
||||
)
|
||||
|
||||
_timeout = (
|
||||
mcp_server.timeout if mcp_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
)
|
||||
try:
|
||||
mcp_responses = await asyncio.wait_for(
|
||||
asyncio.gather(*tasks), timeout=_timeout
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
raise HTTPException(
|
||||
status_code=504,
|
||||
detail={
|
||||
"error": "timeout",
|
||||
"message": f"MCP tool call timed out after {_timeout}s",
|
||||
},
|
||||
)
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
raise e
|
||||
|
||||
# If proxy_logging_obj is None, the tool call result is at index 0
|
||||
# If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task)
|
||||
result_index = 1 if proxy_logging_obj else 0
|
||||
result = mcp_responses[result_index]
|
||||
self._remember_upstream_initialize_instructions(mcp_server, client)
|
||||
|
||||
return cast(CallToolResult, result)
|
||||
|
|
@ -4538,6 +4574,8 @@ def _make_global_mcp_server_manager() -> MCPServerManager:
|
|||
_global_mcp_server_manager: Optional[MCPServerManager] = None
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.shared.session import ProgressFnT
|
||||
|
||||
global_mcp_server_manager: MCPServerManager
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -27,7 +27,14 @@ from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerM
|
|||
from litellm.proxy._types import MCPTransport
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp.types import GetPromptResult, Prompt, ReadResourceResult, Resource
|
||||
from mcp.shared.session import ProgressFnT
|
||||
from mcp.types import (
|
||||
CallToolResult,
|
||||
GetPromptResult,
|
||||
Prompt,
|
||||
ReadResourceResult,
|
||||
Resource,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
|
||||
|
|
@ -82,6 +89,8 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
subject_token: Optional[str] = None,
|
||||
forward_caller_headers: bool = False,
|
||||
raise_on_missing_env: bool = False,
|
||||
) -> Result[UpstreamConnection, CredError]:
|
||||
"""The shared egress seam: resolve auth via resolve() and build the UpstreamConnection.
|
||||
|
||||
|
|
@ -106,10 +115,14 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
)
|
||||
if isinstance(auth, Error):
|
||||
return Error(auth.error)
|
||||
resolved_static = await self._resolve_static_headers_with_env_vars(
|
||||
server, user_api_key_auth, raise_on_missing=False
|
||||
egress = await self._build_egress_headers(
|
||||
server,
|
||||
user_api_key_auth,
|
||||
raw_headers,
|
||||
forward_caller_headers=forward_caller_headers,
|
||||
raise_on_missing_env=raise_on_missing_env,
|
||||
)
|
||||
headers = {**(extra_headers or {}), **(resolved_static or {})} or None
|
||||
headers = {**(extra_headers or {}), **egress} or None
|
||||
is_stdio = server.transport == MCPTransport.stdio
|
||||
return Ok(
|
||||
UpstreamConnection(
|
||||
|
|
@ -123,6 +136,53 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
)
|
||||
)
|
||||
|
||||
async def _build_egress_headers(
|
||||
self,
|
||||
server: MCPServer,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
raw_headers: Optional[Dict[str, str]],
|
||||
*,
|
||||
forward_caller_headers: bool,
|
||||
raise_on_missing_env: bool,
|
||||
) -> Dict[str, str]:
|
||||
"""The non-credential egress headers: configured static headers (with per-user env-var
|
||||
interpolation) plus, on the call path, the caller headers the server forwards upstream.
|
||||
|
||||
The credential (Authorization) is resolve()'s, never this dict. raise_on_missing_env
|
||||
propagates MCPMissingUserEnvVarsError (412 + setup URL) on the call path; list ops degrade.
|
||||
forward_caller_headers gates caller-header forwarding (call only; list/prompts/resources do
|
||||
not forward, matching v1).
|
||||
"""
|
||||
static = await self._resolve_static_headers_with_env_vars(
|
||||
server, user_api_key_auth, raise_on_missing=raise_on_missing_env
|
||||
)
|
||||
forwarded = (
|
||||
self._forwarded_request_headers(server, raw_headers)
|
||||
if forward_caller_headers
|
||||
else {}
|
||||
)
|
||||
return {**forwarded, **(static or {})}
|
||||
|
||||
@staticmethod
|
||||
def _forwarded_request_headers(
|
||||
server: MCPServer,
|
||||
raw_headers: Optional[Dict[str, str]],
|
||||
) -> Dict[str, str]:
|
||||
"""Caller request headers the server is configured to forward (server.extra_headers pulled
|
||||
from raw_headers). Authorization is always stripped: the upstream credential is resolve()'s,
|
||||
so the forwarder never ships inbound auth upstream (the credential-isolation invariant; the
|
||||
passthrough/override paths that would forward an inbound token defer to v1).
|
||||
"""
|
||||
if not server.extra_headers or not raw_headers:
|
||||
return {}
|
||||
normalized = {k.lower(): v for k, v in raw_headers.items()}
|
||||
return {
|
||||
header: normalized[header.lower()]
|
||||
for header in server.extra_headers
|
||||
if header.lower() != "authorization"
|
||||
and normalized.get(header.lower()) is not None
|
||||
}
|
||||
|
||||
async def _get_tools_from_server(
|
||||
self,
|
||||
server: MCPServer,
|
||||
|
|
@ -266,6 +326,69 @@ class MCPServerManagerV2(MCPServerManager):
|
|||
self._egress_item_failure(server, result.error)
|
||||
return result.ok
|
||||
|
||||
async def _open_and_call_tool(
|
||||
self,
|
||||
mcp_server: MCPServer,
|
||||
original_tool_name: str,
|
||||
arguments: Dict[str, object],
|
||||
*,
|
||||
mcp_auth_header: Optional[str],
|
||||
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
|
||||
oauth2_headers: Optional[Dict[str, str]],
|
||||
raw_headers: Optional[Dict[str, str]],
|
||||
hook_extra_headers: Optional[Dict[str, str]],
|
||||
host_progress_callback: Optional[ProgressFnT],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth],
|
||||
) -> CallToolResult:
|
||||
from litellm.proxy.gateway.mcp.result import Error
|
||||
|
||||
# The effective inbound credential is the per-server header if present, else the deprecated
|
||||
# mcp_auth_header (matches v1); an override here means the request defers to v1.
|
||||
server_auth_header: Optional[Union[str, Dict[str, str]]] = mcp_auth_header
|
||||
if mcp_server_auth_headers:
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
lookup_mcp_server_auth_in_headers,
|
||||
)
|
||||
|
||||
found = lookup_mcp_server_auth_in_headers(
|
||||
mcp_server_auth_headers,
|
||||
alias=mcp_server.alias,
|
||||
server_name=mcp_server.server_name,
|
||||
)
|
||||
if found is not None:
|
||||
server_auth_header = found
|
||||
|
||||
# Defer to v1 for guardrail-injected headers (JWT signer etc.), which the v2 path does not
|
||||
# apply, in addition to the usual _should_defer cases (unmapped mode, OpenAPI, override).
|
||||
if hook_extra_headers or self._should_defer(mcp_server, server_auth_header):
|
||||
return await super()._open_and_call_tool(
|
||||
mcp_server,
|
||||
original_tool_name,
|
||||
arguments,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
hook_extra_headers=hook_extra_headers,
|
||||
host_progress_callback=host_progress_callback,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
)
|
||||
conn = await self._v2_connection(
|
||||
mcp_server,
|
||||
user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
forward_caller_headers=True,
|
||||
raise_on_missing_env=True,
|
||||
)
|
||||
if isinstance(conn, Error):
|
||||
self._egress_item_failure(mcp_server, conn.error)
|
||||
result = await conn.ok.call_tool(
|
||||
original_tool_name, arguments, host_progress_callback
|
||||
)
|
||||
if isinstance(result, Error):
|
||||
self._egress_item_failure(mcp_server, result.error)
|
||||
return result.ok
|
||||
|
||||
def _egress_list_failure(
|
||||
self, server: MCPServer, error: "CredError | ConnError"
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ if TYPE_CHECKING:
|
|||
)
|
||||
from mcp import ReadResourceResult, Resource
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.shared.session import ProgressFnT
|
||||
from mcp.types import CallToolResult, GetPromptResult, Prompt
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
|
|
@ -193,10 +194,15 @@ class UpstreamConnection:
|
|||
return await self._run(op)
|
||||
|
||||
async def call_tool(
|
||||
self, name: str, arguments: Dict[str, object]
|
||||
self,
|
||||
name: str,
|
||||
arguments: Dict[str, object],
|
||||
progress_callback: Optional[ProgressFnT] = None,
|
||||
) -> Result[CallToolResult, ConnError]:
|
||||
async def op(session: ClientSession) -> CallToolResult:
|
||||
return await session.call_tool(name, arguments)
|
||||
return await session.call_tool(
|
||||
name, arguments, progress_callback=progress_callback
|
||||
)
|
||||
|
||||
return await self._run(op)
|
||||
|
||||
|
|
|
|||
|
|
@ -159,3 +159,30 @@ async def test_v2_override_gets_prompt_via_upstream_connection(echo_server_url):
|
|||
)
|
||||
result = await manager.get_prompt_from_server(server, "greeting", {"name": "Tin"})
|
||||
assert result.messages
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v2_override_calls_tool_via_upstream_connection(echo_server_url):
|
||||
# Tool-call path: the factored _open_and_call_tool seam routes through resolve() +
|
||||
# UpstreamConnection.call_tool (v2) for the none mode, returning the upstream CallToolResult.
|
||||
manager = MCPServerManagerV2()
|
||||
server = MCPServer(
|
||||
server_id="echo1",
|
||||
name="echo1",
|
||||
transport=MCPTransport.http,
|
||||
url=echo_server_url,
|
||||
auth_type=MCPAuth.none,
|
||||
)
|
||||
result = await manager._open_and_call_tool(
|
||||
server,
|
||||
"echo",
|
||||
{"text": "hi"},
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
hook_extra_headers=None,
|
||||
host_progress_callback=None,
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
assert any("echo: hi" in getattr(c, "text", "") for c in result.content)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue