Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 57 additions & 41 deletions litellm/proxy/_experimental/mcp_server/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -270,6 +270,9 @@ def _jsonrpc_text_has_top_level_method(text: str) -> bool:
global_mcp_tool_registry,
)
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
is_tool_name_prefixed,
normalize_server_name,
split_server_prefix_from_name,
)

Expand Down Expand Up @@ -2483,47 +2486,60 @@ async def execute_mcp_tool( # noqa: PLR0915
None,
)

# Resolve the actual MCP server up-front so the permission check uses
# the canonical server.name even when the tool name is prefixed with a
# short ID (LITELLM_USE_SHORT_MCP_TOOL_PREFIX) that doesn't match the
# server's display name directly.
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
if mcp_server is None and requested_server is not None:
# REST callers may pass the raw tool name (no prefix) plus a
# ``requested_server_id``. The mapping might only contain the
# prefixed form, so retry the lookup with every known prefix of
# the requested server before treating the tool as unresolved —
# otherwise the tool_server_mismatch guard below is silently
# bypassed.
for known_prefix in iter_known_server_prefixes(requested_server):
candidate = global_mcp_server_manager._get_mcp_server_from_tool_name(
add_server_prefix_to_name(name, known_prefix)
)
if candidate is not None:
mcp_server = candidate
break
if mcp_server is not None:
server_name = mcp_server.name

# REST /mcp-rest/tools/call passes server_id — tool must belong to that server
if requested_server is not None:
if (
mcp_server is not None
and mcp_server.server_id != requested_server.server_id
):
raise HTTPException(
status_code=403,
detail={
"error": "tool_server_mismatch",
"message": (
f"Tool '{name}' belongs to MCP server '{mcp_server.name}' "
f"but request specified server_id for '{requested_server.name}'."
),
},
)
if mcp_server is None:
mcp_server = requested_server
server_name = requested_server.name
name_is_prefixed = False
if requested_server is not None and MCP_TOOL_PREFIX_SEPARATOR in name:
all_registry_prefixes: Set[str] = set()
for registry_server in global_mcp_server_manager.get_registry().values():
for known_prefix in iter_known_server_prefixes(registry_server):
all_registry_prefixes.add(normalize_server_name(known_prefix))
name_is_prefixed = is_tool_name_prefixed(
name, known_server_prefixes=all_registry_prefixes
)

if requested_server is not None and not name_is_prefixed:
# REST callers may pass server_id with the upstream tool name (no
# LiteLLM prefix). The first segment is not a registered server
# prefix, so the whole string is the upstream tool name and may
# legitimately contain the separator (e.g. "text-to-speech").
# server_id is authoritative for routing and auth.
mcp_server = requested_server
server_name = requested_server.name
original_tool_name = name
else:
# Resolve from tool name (MCP JSON-RPC or prefixed REST tool names).
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
if mcp_server is None and requested_server is not None:
for known_prefix in iter_known_server_prefixes(requested_server):
candidate = (
global_mcp_server_manager._get_mcp_server_from_tool_name(
add_server_prefix_to_name(name, known_prefix)
)
)
if candidate is not None:
mcp_server = candidate
break
if mcp_server is not None:
server_name = mcp_server.name

if requested_server is not None:
if (
mcp_server is not None
and mcp_server.server_id != requested_server.server_id
):
raise HTTPException(
status_code=403,
detail={
"error": "tool_server_mismatch",
"message": (
f"Tool '{name}' belongs to MCP server "
f"'{mcp_server.name}' but request specified "
f"server_id for '{requested_server.name}'."
),
},
)
if mcp_server is None:
mcp_server = requested_server
server_name = requested_server.name

# Only enforce server-level permissions when we can resolve a server
if server_name:
Expand Down
38 changes: 28 additions & 10 deletions tests/local_testing/test_azure_anthropic_sync_post.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,43 +2,61 @@
``_get_httpx_client`` + ``HTTPHandler.post`` (same pattern as Azure Anthropic sync path:
``_get_httpx_client(params={"timeout": ...})`` then ``post(..., timeout=...)``).

Uses https://httpbin.org/delay/10 with ``timeout=5`` — the handler must raise :class:`~litellm.exceptions.Timeout`
before the 10s delay completes. Skips if httpbin is unreachable.
A local server stalls longer than the per-request ``timeout`` but well under the client
default, so the handler must raise :class:`~litellm.exceptions.Timeout` from the per-request
override rather than completing under the (much larger) client default.

Lives under ``local_testing`` (not ``make test-unit``).
"""

import json
import os
import sys
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

import httpx
import pytest

sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))

from litellm.exceptions import Timeout as LitellmTimeout
from litellm.llms.custom_httpx.http_handler import _get_httpx_client

_HTTPBIN_DELAY_S = 10
_PER_REQUEST_TIMEOUT_S = 5.0
_SERVER_DELAY_S = 5
_PER_REQUEST_TIMEOUT_S = 1.0
_CLIENT_DEFAULT_TIMEOUT_S = 60.0


class _SlowHandler(BaseHTTPRequestHandler):
def do_POST(self):
time.sleep(_SERVER_DELAY_S)
try:
self.send_response(200)
self.end_headers()
self.wfile.write(b"{}")
except OSError:
pass

def log_message(self, *args):
pass


def test_post_delay_exceeds_per_request_timeout_raises():
try:
httpx.get("https://httpbin.org/get", timeout=5.0)
except Exception as e:
pytest.skip(f"httpbin.org unreachable: {e}")
server = ThreadingHTTPServer(("127.0.0.1", 0), _SlowHandler)
threading.Thread(target=server.serve_forever, daemon=True).start()
host, port = server.server_address

handler = _get_httpx_client(params={"timeout": _CLIENT_DEFAULT_TIMEOUT_S})
try:
with pytest.raises(LitellmTimeout):
handler.post(
f"https://httpbin.org/delay/{_HTTPBIN_DELAY_S}",
f"http://{host}:{port}/delay",
headers={"content-type": "application/json"},
data=json.dumps({"model": "claude", "messages": []}),
timeout=_PER_REQUEST_TIMEOUT_S,
)
finally:
handler.close()
server.shutdown()
server.server_close()
Loading
Loading