mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
fix lazymcp review regressions
This commit is contained in:
parent
7ba28bb6dd
commit
9df9d65048
6 changed files with 183 additions and 53 deletions
|
|
@ -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),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue