fix lazymcp review regressions

This commit is contained in:
jibanez-staticduo 2026-05-19 08:14:42 +02:00
parent 7ba28bb6dd
commit 9df9d65048
No known key found for this signature in database
6 changed files with 183 additions and 53 deletions

View file

@ -805,6 +805,9 @@ if MCP_AVAILABLE:
match_list = [
s.lower() for s in iter_known_server_prefixes(server) if s
]
server_name = getattr(server, "name", None)
if server_name:
match_list.append(str(server_name).lower())
if server_or_group.lower() in match_list:
filtered_server[server.server_id] = server
@ -828,6 +831,9 @@ if MCP_AVAILABLE:
f"Could not resolve '{server_or_group}' as access group: {e}"
)
if mcp_servers is not None:
return list(filtered_server.values())
if filtered_server:
return list(filtered_server.values())
@ -2005,6 +2011,9 @@ if MCP_AVAILABLE:
try:
from litellm.proxy.proxy_server import user_api_key_cache
# DualCache does not expose prefix invalidation, so this intentionally
# mirrors existing targeted invalidation paths and only touches the
# process-local dict. Redis entries expire quickly via the cache TTL.
in_mem = getattr(user_api_key_cache, "in_memory_cache", None)
cache_dict = getattr(in_mem, "cache_dict", {}) if in_mem else {}
for key in [k for k in cache_dict if str(k).startswith("lazymcp:")]:
@ -2059,6 +2068,28 @@ if MCP_AVAILABLE:
)
return apply_tool_overrides(tools, server)
def _build_lazymcp_catalog_description(servers: List[Dict[str, Any]]) -> str:
description_lines = [
"Describe MCP servers and tools available through the LiteLLM LazyMCP gateway.",
"",
"Available MCP servers:",
]
if servers:
description_lines.extend(
f"- {item['name']}: {item['description']}" for item in servers
)
else:
description_lines.append("- No MCP servers are available for this request.")
description_lines.extend(
[
"",
'Call mcp_describe with {"server":"<name>"} to list tools for one server with input schemas.',
'Call mcp_describe with {"server":"<name>","tool":"<tool_name>"} to get details for one tool with its input schema.',
'Call mcp_call with {"server":"<name>","tool":"<tool_name>","arguments":{...}} to execute a tool.',
]
)
return "\n".join(description_lines)
async def _get_lazymcp_catalog(
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_auth_header: Optional[str],
@ -2096,6 +2127,7 @@ if MCP_AVAILABLE:
return {
**cached,
"servers": filtered_servers,
"description": _build_lazymcp_catalog_description(filtered_servers),
"server_count": len(filtered_servers),
"tool_count": sum(
server.get("tool_count", 0) for server in filtered_servers
@ -2135,28 +2167,9 @@ if MCP_AVAILABLE:
}
)
description_lines = [
"Describe MCP servers and tools available through the LiteLLM LazyMCP gateway.",
"",
"Available MCP servers:",
]
if servers:
description_lines.extend(
f"- {item['name']}: {item['description']}" for item in servers
)
else:
description_lines.append("- No MCP servers are available for this request.")
description_lines.extend(
[
"",
'Call mcp_describe with {"server":"<name>"} to list tools for one server with input schemas.',
'Call mcp_describe with {"server":"<name>","tool":"<tool_name>"} to get details for one tool with its input schema.',
'Call mcp_call with {"server":"<name>","tool":"<tool_name>","arguments":{...}} to execute a tool.',
]
)
catalog = {
"servers": servers,
"description": "\n".join(description_lines),
"description": _build_lazymcp_catalog_description(servers),
"server_count": len(servers),
"tool_count": sum(item["tool_count"] for item in servers),
}

View file

@ -15869,6 +15869,23 @@ async def root_lazymcp_route(request: Request):
raise HTTPException(status_code=500, detail=f"Internal server error: {str(e)}")
async def _lazymcp_forward_as_path(path_segment: str, request: Request):
"""Rewrite path to /lazymcp/{path_segment} and stream the LazyMCP response.
LazyMCP uses a separate session manager from the standard MCP endpoint, so
this stays as a small wrapper instead of sharing _mcp_forward_as_path.
"""
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_lazymcp,
)
scope = dict(request.scope)
scope["path"] = f"/lazymcp/{path_segment}"
return await _stream_mcp_asgi_response(
handle_streamable_http_lazymcp, scope, request.receive
)
@app.api_route(
"/lazymcp/{mcp_server_name}/",
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],
@ -15878,7 +15895,7 @@ async def root_lazymcp_route(request: Request):
methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"],
)
async def dynamic_lazymcp_route(mcp_server_name: str, request: Request):
"""Handle dynamic LazyMCP server routes like /lazymcp/github_mcp."""
"""Handle /lazymcp/{name} for MCP servers, toolsets, access groups, and CSV lists."""
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
@ -15893,14 +15910,28 @@ async def dynamic_lazymcp_route(mcp_server_name: str, request: Request):
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(
mcp_server_name, client_ip=client_ip
)
scope = dict(request.scope)
scope["path"] = f"/lazymcp/{mcp_server_name}"
if mcp_server is None and prisma_client is not None:
if mcp_server is not None:
return await _lazymcp_forward_as_path(mcp_server_name, request)
if "," in mcp_server_name:
resolved_tokens = await _resolve_mcp_csv_tokens(mcp_server_name, client_ip)
if not resolved_tokens:
raise HTTPException(
status_code=404,
detail=(
f"No MCP server, toolset, or access group in "
f"'{mcp_server_name}' resolved to a known target"
),
)
return await _lazymcp_forward_as_path(",".join(resolved_tokens), request)
if prisma_client is not None:
toolset = await global_mcp_server_manager.get_toolset_by_name_cached(
prisma_client, mcp_server_name
)
if toolset is not None:
scope = dict(request.scope)
scope["path"] = "/lazymcp"
token = _mcp_active_toolset_id.set(toolset.toolset_id)
try:
@ -15910,10 +15941,12 @@ async def dynamic_lazymcp_route(mcp_server_name: str, request: Request):
finally:
_mcp_active_toolset_id.reset(token)
# Defer all remaining names (server, access-group, or invalid target) to
# the LazyMCP handler, which applies the existing group/permission resolver.
return await _stream_mcp_asgi_response(
handle_streamable_http_lazymcp, scope, request.receive
if await _is_mcp_access_group_cached(mcp_server_name):
return await _lazymcp_forward_as_path(mcp_server_name, request)
raise HTTPException(
status_code=404,
detail=f"MCP server, toolset, or access group '{mcp_server_name}' not found",
)
except HTTPException as e:
raise e

View file

@ -248,23 +248,36 @@ class LiteLLM_Proxy_MCP_Handler:
active_toolset_id: Optional[str] = None
if effective_filter and len(effective_filter) == 1:
requested_scope = effective_filter[0]
if not global_mcp_server_manager.get_mcp_server_by_name(requested_scope):
try:
from litellm.proxy.proxy_server import prisma_client
if global_mcp_server_manager.get_mcp_server_by_name(requested_scope):
return effective_filter, active_toolset_id
try:
from litellm.proxy.proxy_server import _is_mcp_access_group_cached
if prisma_client is not None:
toolset = (
await global_mcp_server_manager.get_toolset_by_name_cached(
prisma_client, requested_scope
)
if await _is_mcp_access_group_cached(requested_scope):
return effective_filter, active_toolset_id
except Exception as _e:
verbose_logger.debug(
f"Could not resolve LazyMCP scope '{requested_scope}' as access group: {_e}"
)
try:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is not None:
toolset = (
await global_mcp_server_manager.get_toolset_by_name_cached(
prisma_client, requested_scope
)
if toolset is not None:
active_toolset_id = toolset.toolset_id
effective_filter = None
except Exception as _e:
verbose_logger.debug(
f"Could not resolve LazyMCP scope '{requested_scope}' as toolset: {_e}"
)
if toolset is not None:
active_toolset_id = toolset.toolset_id
effective_filter = None
else:
effective_filter = []
except Exception as _e:
verbose_logger.debug(
f"Could not resolve LazyMCP scope '{requested_scope}' as toolset: {_e}"
)
effective_filter = []
return effective_filter, active_toolset_id
@staticmethod
@ -463,6 +476,8 @@ class LiteLLM_Proxy_MCP_Handler:
client_ip=client_ip,
)
# LazyMCP keeps the sentinel so unresolved Responses/Chat client IPs
# remain fail-closed in the catalog/IP filter instead of broadening.
standard_client_ip = (
None if client_ip == INVALID_MCP_CLIENT_IP_SENTINEL else client_ip
)

View file

@ -479,6 +479,9 @@ async def test_lazymcp_cached_catalog_rechecks_current_visibility(monkeypatch):
assert catalog["server_count"] == 1
assert catalog["tool_count"] == 1
assert [server["name"] for server in catalog["servers"]] == ["visible"]
assert "visible" in catalog["description"]
assert "revoked" not in catalog["description"]
assert "must not leak" not in catalog["description"]
@pytest.mark.asyncio
@ -521,6 +524,8 @@ async def test_lazymcp_cached_catalog_hides_all_revoked_servers(monkeypatch):
assert catalog["server_count"] == 0
assert catalog["tool_count"] == 0
assert catalog["servers"] == []
assert "revoked" not in catalog["description"]
assert "No MCP servers are available" in catalog["description"]
def test_invalidating_toolset_cache_tolerates_lazymcp_invalidation_error():
@ -788,12 +793,6 @@ def test_lazymcp_dynamic_route_falls_back_for_non_toolset(monkeypatch):
async def fake_get_toolset(_prisma_client, _toolset_name):
return None
async def fake_stream_response(_handle_fn, scope, _receive):
from starlette.responses import Response
assert scope["path"] == "/lazymcp/github"
return Response("ok", media_type="text/event-stream")
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object())
monkeypatch.setattr(
"litellm.proxy.auth.ip_address_utils.IPAddressUtils.get_mcp_client_ip",
@ -808,12 +807,13 @@ def test_lazymcp_dynamic_route_falls_back_for_non_toolset(monkeypatch):
fake_get_toolset,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response
"litellm.proxy.proxy_server._is_mcp_access_group_cached",
AsyncMock(return_value=False),
)
response = TestClient(app).get("/lazymcp/github", follow_redirects=False)
assert response.status_code == 200
assert response.status_code == 404
def test_lazymcp_toolset_route_returns_404_for_missing_toolset(monkeypatch):

View file

@ -37,6 +37,7 @@ _GET_ACCESS_GROUP_SERVERS = (
"MCPRequestHandler._get_mcp_servers_from_access_groups"
)
_FORWARD = "litellm.proxy.proxy_server._mcp_forward_as_path"
_LAZYMCP_FORWARD = "litellm.proxy.proxy_server._lazymcp_forward_as_path"
_RESOLVE_CSV = "litellm.proxy.proxy_server._resolve_mcp_csv_tokens"
@ -486,3 +487,50 @@ async def test_dynamic_mcp_route_empty_access_group_returns_404():
await dynamic_mcp_route("empty_group", request)
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_dynamic_lazymcp_route_unknown_name_returns_404():
from litellm.proxy.proxy_server import dynamic_lazymcp_route
request = _make_request("/lazymcp/does_not_exist")
fake_mgr = MagicMock()
fake_mgr.get_mcp_server_by_name = MagicMock(return_value=None)
fake_mgr.get_toolset_by_name_cached = AsyncMock(return_value=None)
with (
patch(_MCP_MANAGER, fake_mgr),
patch(_PRISMA, new=MagicMock()),
patch(_IS_ACCESS_GROUP, new=AsyncMock(return_value=False)),
):
with pytest.raises(HTTPException) as exc_info:
await dynamic_lazymcp_route("does_not_exist", request)
assert exc_info.value.status_code == 404
assert "does_not_exist" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_dynamic_lazymcp_route_resolves_access_group_without_broadening():
from starlette.responses import Response
from litellm.proxy.proxy_server import dynamic_lazymcp_route
request = _make_request("/lazymcp/dev_group")
fake_mgr = MagicMock()
fake_mgr.get_mcp_server_by_name = MagicMock(return_value=None)
fake_mgr.get_toolset_by_name_cached = AsyncMock(return_value=None)
fake_forward = AsyncMock(return_value=Response(content=b"{}", status_code=200))
with (
patch(_MCP_MANAGER, fake_mgr),
patch(_PRISMA, new=MagicMock()),
patch(_IS_ACCESS_GROUP, new=AsyncMock(return_value=True)),
patch(_LAZYMCP_FORWARD, new=fake_forward),
):
response = await dynamic_lazymcp_route("dev_group", request)
assert response.status_code == 200
fake_forward.assert_awaited_once_with("dev_group", request)

View file

@ -284,7 +284,7 @@ def test_get_requested_mcp_servers_handles_lazymcp_variants():
@pytest.mark.asyncio
async def test_resolve_lazymcp_scope_handles_server_toolset_and_errors(monkeypatch):
server_manager = types.SimpleNamespace(
get_mcp_server_by_name=MagicMock(side_effect=[object(), None, None]),
get_mcp_server_by_name=MagicMock(side_effect=[object(), None, None, None]),
get_toolset_by_name_cached=AsyncMock(
side_effect=[
types.SimpleNamespace(toolset_id="toolset-1"),
@ -292,7 +292,10 @@ async def test_resolve_lazymcp_scope_handles_server_toolset_and_errors(monkeypat
]
),
)
proxy_module = types.SimpleNamespace(prisma_client=object())
proxy_module = types.SimpleNamespace(
prisma_client=object(),
_is_mcp_access_group_cached=AsyncMock(return_value=False),
)
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
assert await LiteLLM_Proxy_MCP_Handler._resolve_lazymcp_scope(
@ -303,7 +306,25 @@ async def test_resolve_lazymcp_scope_handles_server_toolset_and_errors(monkeypat
) == (None, "toolset-1")
assert await LiteLLM_Proxy_MCP_Handler._resolve_lazymcp_scope(
["broken"], server_manager
) == (["broken"], None)
) == ([], None)
@pytest.mark.asyncio
async def test_resolve_lazymcp_scope_keeps_access_group(monkeypatch):
server_manager = types.SimpleNamespace(
get_mcp_server_by_name=MagicMock(return_value=None),
get_toolset_by_name_cached=AsyncMock(return_value=None),
)
proxy_module = types.SimpleNamespace(
prisma_client=object(),
_is_mcp_access_group_cached=AsyncMock(return_value=True),
)
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
assert await LiteLLM_Proxy_MCP_Handler._resolve_lazymcp_scope(
["dev_group"], server_manager
) == (["dev_group"], None)
server_manager.get_toolset_by_name_cached.assert_not_awaited()
@pytest.mark.asyncio