From 03d67173c32d458991611577cf5fcae495c75726 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 13 May 2026 15:01:56 +0200 Subject: [PATCH] feat(mcp): add LazyMCP gateway endpoint Add a compact LazyMCP gateway over the existing MCP infrastructure, preserving MCP permissions, IP filtering, OAuth headers, toolset grants, and Responses API execution behavior while exposing only mcp_describe, mcp_call, and mcp_status to agents. --- .../mcp_server/mcp_server_manager.py | 11 + .../proxy/_experimental/mcp_server/server.py | 785 +++++++++++++- litellm/proxy/litellm_pre_call_utils.py | 8 +- litellm/proxy/proxy_server.py | 133 +++ litellm/responses/main.py | 5 + .../responses/mcp/chat_completions_handler.py | 9 + .../mcp/litellm_proxy_mcp_handler.py | 308 +++++- .../responses/mcp/mcp_streaming_iterator.py | 7 + litellm/responses/utils.py | 17 + litellm/types/proxy/litellm_pre_call_utils.py | 3 + .../mcp_server/test_mcp_server.py | 984 ++++++++++++++++++ .../mcp/test_litellm_proxy_mcp_handler.py | 712 +++++++++++++ .../src/components/mcp_tools/mcp_connect.tsx | 43 +- .../components/mcp_tools/mcp_servers.test.tsx | 31 + .../src/components/mcp_tools/mcp_servers.tsx | 4 + 15 files changed, 2981 insertions(+), 79 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d1b49039e8e..071de5eb781 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1083,6 +1083,17 @@ class MCPServerManager: ] for k in keys_to_remove: cache_dict.pop(k, None) + try: + from litellm.proxy._experimental.mcp_server.server import ( + invalidate_lazymcp_cache, + ) + + invalidate_lazymcp_cache() + except Exception as lazy_exc: + verbose_logger.debug( + "invalidate_toolset_cache: failed to evict LazyMCP entries: %s", + lazy_exc, + ) except Exception as e: verbose_logger.warning( f"invalidate_toolset_cache: failed to evict in-memory entries: {e}" diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 0a74a92f9ce..08046ef6f87 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -6,6 +6,9 @@ LiteLLM MCP Server Routes import asyncio import contextlib +import hashlib +import json +import re import time import types import traceback @@ -257,13 +260,25 @@ if MCP_AVAILABLE: stateless=True, ) + lazymcp_server: Server = Server( + name=f"{LITELLM_MCP_SERVER_NAME}-lazymcp", + version=LITELLM_MCP_SERVER_VERSION, + ) + lazy_session_manager = StreamableHTTPSessionManager( + app=lazymcp_server, + event_store=None, + json_response=False, + stateless=True, + ) + # Context managers for proper lifecycle management _session_manager_cm = None _sse_session_manager_cm = None + _lazy_session_manager_cm = None async def initialize_session_managers(): """Initialize the session managers. Can be called from main app lifespan.""" - global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm, _lazy_session_manager_cm # Use async lock to prevent concurrent initialization async with _INITIALIZATION_LOCK: @@ -275,10 +290,12 @@ if MCP_AVAILABLE: # Start the session managers with context managers _session_manager_cm = session_manager.run() _sse_session_manager_cm = sse_session_manager.run() + _lazy_session_manager_cm = lazy_session_manager.run() # Enter the context managers await _session_manager_cm.__aenter__() await _sse_session_manager_cm.__aenter__() + await _lazy_session_manager_cm.__aenter__() _SESSION_MANAGERS_INITIALIZED = True verbose_logger.info( @@ -287,7 +304,7 @@ if MCP_AVAILABLE: async def shutdown_session_managers(): """Shutdown the session managers.""" - global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm, _lazy_session_manager_cm if _SESSION_MANAGERS_INITIALIZED: verbose_logger.info("Shutting down MCP session managers...") @@ -297,11 +314,14 @@ if MCP_AVAILABLE: await _session_manager_cm.__aexit__(None, None, None) if _sse_session_manager_cm: await _sse_session_manager_cm.__aexit__(None, None, None) + if _lazy_session_manager_cm: + await _lazy_session_manager_cm.__aexit__(None, None, None) except Exception as e: verbose_logger.exception(f"Error during session manager shutdown: {e}") _session_manager_cm = None _sse_session_manager_cm = None + _lazy_session_manager_cm = None _SESSION_MANAGERS_INITIALIZED = False @contextlib.asynccontextmanager @@ -504,6 +524,70 @@ if MCP_AVAILABLE: return response + @lazymcp_server.list_tools() + async def list_lazymcp_tools() -> List[MCPTool]: + try: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + client_ip, + ) = get_auth_context() + catalog = await _get_lazymcp_catalog( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + return _get_lazymcp_gateway_tools(catalog.get("description")) + except Exception as e: + verbose_logger.exception("Error in LazyMCP list_tools endpoint: %s", e) + return _get_lazymcp_gateway_tools() + + @lazymcp_server.call_tool() + async def lazymcp_tool_call( + name: str, arguments: Dict[str, Any] | None + ) -> CallToolResult: + arguments = arguments or {} + try: + if name == "mcp_describe": + return _make_lazymcp_text_result(await _lazymcp_describe(arguments)) + if name == "mcp_status": + return _make_lazymcp_text_result(await _lazymcp_status()) + if name == "mcp_call": + return await _lazymcp_call(arguments) + return CallToolResult( + content=[ + TextContent( + text=json.dumps({"error": "Unknown LazyMCP tool."}), + type="text", + ) + ], + isError=True, + ) + except Exception as e: + verbose_logger.exception("LazyMCP tool call failed: %s", e) + return CallToolResult( + content=[ + TextContent( + text=json.dumps( + { + "error": "Upstream MCP tool call failed.", + "details": str(e), + } + ), + type="text", + ) + ], + isError=True, + ) + @server.list_prompts() async def list_prompts() -> List[Prompt]: """ @@ -1217,6 +1301,7 @@ if MCP_AVAILABLE: log_list_tools_to_spendlogs: bool = False, list_tools_log_source: Optional[str] = None, litellm_trace_id: Optional[str] = None, + client_ip: Optional[str] = None, ) -> List[MCPTool]: """ Helper method to fetch tools from MCP servers based on server filtering criteria. @@ -1308,6 +1393,7 @@ if MCP_AVAILABLE: allowed_mcp_servers = await _get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, + client_ip=client_ip, ) # Pre-fetch OAuth credentials only when at least one server uses OAuth2, @@ -1780,6 +1866,560 @@ if MCP_AVAILABLE: return managed_tools + LAZYMCP_TOOL_NAMES = ("mcp_describe", "mcp_call", "mcp_status") + LAZYMCP_CACHE_TTL_SECONDS = 300 + LAZYMCP_UNAVAILABLE_SERVER_ERROR = { + "error": "MCP server is not available for this request." + } + LAZYMCP_UNAVAILABLE_TOOL_ERROR = { + "error": "Tool is not available for this request." + } + + def _hash_lazymcp_value(value: Any) -> Optional[str]: + if value is None: + return None + encoded = json.dumps(value, sort_keys=True, default=str) + return hashlib.sha256(encoded.encode("utf-8")).hexdigest() + + def _safe_lazymcp_text(value: Optional[str], fallback: str) -> str: + text = re.sub(r"\s+", " ", value or "").strip() or fallback + text = re.sub(r"https?://\S+", "[url]", text) + if len(text) > 160: + text = text[:157].rstrip() + "..." + return text + + def _get_lazymcp_server_label(mcp_server: MCPServer) -> str: + return str( + mcp_server.alias + or mcp_server.server_name + or mcp_server.name + or mcp_server.server_id + ) + + def _get_lazymcp_server_description(mcp_server: MCPServer) -> str: + mcp_info = mcp_server.mcp_info or {} + description = getattr(mcp_server, "description", None) or mcp_info.get( + "description" + ) + return _safe_lazymcp_text(description, "No description configured.") + + def _summarize_lazymcp_schema(schema: Any) -> Dict[str, Any]: + if not isinstance(schema, dict): + return {} + properties = schema.get("properties") + return { + "type": schema.get("type", "object"), + "required": schema.get("required", []), + "properties": ( + sorted(properties.keys()) if isinstance(properties, dict) else [] + ), + } + + def _lazymcp_tool_to_summary( + tool: MCPTool, include_schema: bool = False + ) -> Dict[str, Any]: + summary: Dict[str, Any] = { + "name": tool.name, + "description": _safe_lazymcp_text( + getattr(tool, "description", None), "No description configured." + ), + } + schema = getattr(tool, "inputSchema", None) + if include_schema: + summary["input_schema"] = schema or {} + else: + summary["input_schema_summary"] = _summarize_lazymcp_schema(schema) + return summary + + def _lazymcp_cache_scope( + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_auth_header: Optional[str], + mcp_servers: Optional[List[str]], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + client_ip: Optional[str], + ) -> str: + object_permission = getattr(user_api_key_auth, "object_permission", None) + normalized_raw_headers = { + str(key).lower(): value + for key, value in (raw_headers or {}).items() + if isinstance(key, str) + } + object_permission_payload = None + if object_permission is not None: + object_permission_payload = ( + object_permission.model_dump(mode="json") + if hasattr(object_permission, "model_dump") + else str(object_permission) + ) + + scope_payload = { + "api_key_hash": _hash_lazymcp_value( + getattr(user_api_key_auth, "api_key", None) + ), + "user_id": getattr(user_api_key_auth, "user_id", None), + "team_id": getattr(user_api_key_auth, "team_id", None), + "mcp_servers": mcp_servers or [], + "active_toolset": _mcp_active_toolset_id.get(), + "client_ip": client_ip, + "mcp_auth_header_hash": _hash_lazymcp_value(mcp_auth_header), + "mcp_server_auth_headers_hash": _hash_lazymcp_value( + mcp_server_auth_headers + ), + "oauth2_headers_hash": _hash_lazymcp_value(oauth2_headers), + "header_mcp_servers": normalized_raw_headers.get("x-mcp-servers"), + "header_mcp_access_groups": normalized_raw_headers.get( + "x-mcp-access-groups" + ), + "raw_header_names": sorted(normalized_raw_headers.keys()), + "raw_header_values_hash": _hash_lazymcp_value(normalized_raw_headers), + "object_permission": object_permission_payload, + } + encoded = json.dumps(scope_payload, sort_keys=True, default=str) + return hashlib.sha256(encoded.encode("utf-8")).hexdigest() + + async def _lazymcp_cache_get(key: str) -> Optional[Any]: + try: + from litellm.proxy.proxy_server import user_api_key_cache + + return await user_api_key_cache.async_get_cache(key=key) + except Exception as e: + verbose_logger.debug("LazyMCP cache get failed for %s: %s", key, e) + return None + + async def _lazymcp_cache_set(key: str, value: Any) -> None: + try: + from litellm.proxy.proxy_server import user_api_key_cache + + await user_api_key_cache.async_set_cache( + key=key, + value=value, + ttl=LAZYMCP_CACHE_TTL_SECONDS, + ) + except Exception as e: + verbose_logger.debug("LazyMCP cache set failed for %s: %s", key, e) + + def invalidate_lazymcp_cache() -> None: + """Evict LazyMCP entries from the in-memory DualCache layer only.""" + try: + from litellm.proxy.proxy_server import user_api_key_cache + + 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:")]: + cache_dict.pop(key, None) + except Exception as e: + verbose_logger.warning("invalidate_lazymcp_cache failed: %s", e) + + async def _get_lazymcp_allowed_servers( + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_servers: Optional[List[str]], + client_ip: Optional[str], + ) -> List[MCPServer]: + user_api_key_auth = await _merge_toolset_permissions(user_api_key_auth) + return await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + + async def _get_lazymcp_server_tools( + server: MCPServer, + user_api_key_auth: Optional[UserAPIKeyAuth], + 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]], + ) -> List[MCPTool]: + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + if extra_headers is None and server.auth_type == MCPAuth.oauth2: + extra_headers = await _get_user_oauth_extra_headers_from_db( + server=server, + user_api_key_auth=user_api_key_auth, + ) + tools = await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, + raw_headers=raw_headers, + ) + tools = filter_tools_by_allowed_tools(tools, server) + tools = await filter_tools_by_key_team_permissions( + tools=tools, + server_id=server.server_id, + user_api_key_auth=user_api_key_auth, + ) + return apply_tool_overrides(tools, server) + + async def _get_lazymcp_catalog( + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_auth_header: Optional[str], + mcp_servers: Optional[List[str]], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + client_ip: Optional[str], + ) -> Dict[str, Any]: + scope_hash = _lazymcp_cache_scope( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + cache_key = f"lazymcp:catalog:{scope_hash}" + cached = await _lazymcp_cache_get(cache_key) + if isinstance(cached, dict): + return cached + + allowed_servers = await _get_lazymcp_allowed_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + servers: List[Dict[str, Any]] = [] + for server_item in allowed_servers: + try: + tools = await _get_lazymcp_server_tools( + server=server_item, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + except Exception as e: + verbose_logger.exception( + "LazyMCP failed to list tools for server %s: %s", + _get_lazymcp_server_label(server_item), + e, + ) + tools = [] + servers.append( + { + "server_id": server_item.server_id, + "name": _get_lazymcp_server_label(server_item), + "description": _get_lazymcp_server_description(server_item), + "tool_count": len(tools), + "tools": [_lazymcp_tool_to_summary(tool) for tool in tools], + } + ) + + 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":""} to list tools for one server with input schemas.', + 'Call mcp_describe with {"server":"","tool":""} to get details for one tool with its input schema.', + 'Call mcp_call with {"server":"","tool":"","arguments":{...}} to execute a tool.', + ] + ) + catalog = { + "servers": servers, + "description": "\n".join(description_lines), + "server_count": len(servers), + "tool_count": sum(item["tool_count"] for item in servers), + } + await _lazymcp_cache_set(cache_key, catalog) + return catalog + + def _find_lazymcp_server( + catalog: Dict[str, Any], server_name: str + ) -> Optional[Dict[str, Any]]: + requested = server_name.lower() + for item in catalog.get("servers", []): + if str(item.get("name", "")).lower() == requested: + return item + return None + + async def _lazymcp_describe(arguments: Dict[str, Any]) -> Dict[str, Any]: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + client_ip, + ) = get_auth_context() + catalog = await _get_lazymcp_catalog( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + server_name = arguments.get("server") + tool_name = arguments.get("tool") + if not server_name: + return { + "servers": [ + { + "name": item["name"], + "description": item["description"], + "tool_count": item["tool_count"], + } + for item in catalog.get("servers", []) + ] + } + + server_item = _find_lazymcp_server(catalog, str(server_name)) + if server_item is None: + return LAZYMCP_UNAVAILABLE_SERVER_ERROR + + allowed_servers = await _get_lazymcp_allowed_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + selected_server = next( + ( + server + for server in allowed_servers + if _get_lazymcp_server_label(server).lower() == str(server_name).lower() + ), + None, + ) + if selected_server: + tools = await _get_lazymcp_server_tools( + selected_server, + user_api_key_auth, + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) + if tool_name: + for tool in tools: + if tool.name == tool_name: + return { + "server": _get_lazymcp_server_label(selected_server), + "tool": _lazymcp_tool_to_summary(tool, include_schema=True), + } + return LAZYMCP_UNAVAILABLE_TOOL_ERROR + return { + "server": _get_lazymcp_server_label(selected_server), + "description": server_item["description"], + "tools": [ + _lazymcp_tool_to_summary(tool, include_schema=True) + for tool in tools + ], + } + + # The catalog can be served from a short-lived cache, so re-check the + # current permission/IP-filter result before returning server/tool + # details. If access was revoked after the catalog was cached, do not + # leak stale cached tool metadata. + return LAZYMCP_UNAVAILABLE_SERVER_ERROR + + async def _lazymcp_status() -> Dict[str, Any]: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + client_ip, + ) = get_auth_context() + catalog = await _get_lazymcp_catalog( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + return { + "enabled": True, + "mode": "lazymcp", + "visible_server_count": catalog.get("server_count", 0), + "visible_tool_count": catalog.get("tool_count", 0), + "route_restricted": bool(mcp_servers), + "requested_server_count": len(mcp_servers or []), + "toolset_scoped": _mcp_active_toolset_id.get() is not None, + } + + async def _lazymcp_call(arguments: Dict[str, Any]) -> CallToolResult: + server_name = arguments.get("server") + tool_name = arguments.get("tool") + tool_arguments = arguments.get("arguments") + if ( + not isinstance(server_name, str) + or not isinstance(tool_name, str) + or not isinstance(tool_arguments, dict) + ): + return CallToolResult( + content=[ + TextContent( + text=json.dumps( + { + "error": "Invalid LazyMCP arguments.", + "details": "mcp_call requires server, tool, and arguments.", + } + ), + type="text", + ) + ], + isError=True, + ) + + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + client_ip, + ) = get_auth_context() + allowed_servers = await _get_lazymcp_allowed_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + selected_server = next( + ( + server + for server in allowed_servers + if _get_lazymcp_server_label(server).lower() == server_name.lower() + ), + None, + ) + if selected_server is None: + return CallToolResult( + content=[ + TextContent( + text=json.dumps(LAZYMCP_UNAVAILABLE_SERVER_ERROR), type="text" + ) + ], + isError=True, + ) + visible_tools = await _get_lazymcp_server_tools( + selected_server, + user_api_key_auth, + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) + selected_tool = next( + (tool for tool in visible_tools if tool.name == tool_name), None + ) + if selected_tool is None: + return CallToolResult( + content=[ + TextContent( + text=json.dumps(LAZYMCP_UNAVAILABLE_TOOL_ERROR), type="text" + ) + ], + isError=True, + ) + return await call_mcp_tool( + name=selected_tool.name, + arguments=tool_arguments, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=[_get_lazymcp_server_label(selected_server)], + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + metadata={ + "source": "lazymcp", + "lazy_mcp": True, + "lazy_mcp_gateway_tool": "mcp_call", + "lazy_mcp_server": server_name, + "lazy_mcp_tool": tool_name, + "delegated_tool_name": tool_name, + "server_name": server_name, + }, + ) + + def _make_lazymcp_text_result(payload: Dict[str, Any]) -> CallToolResult: + return CallToolResult( + content=[TextContent(text=json.dumps(payload, default=str), type="text")], + isError=bool(payload.get("error")), + ) + + def _get_lazymcp_gateway_tools(description: Optional[str] = None) -> List[MCPTool]: + describe_description = description or ( + "Describe MCP servers and tools available through the LiteLLM LazyMCP gateway." + ) + return [ + MCPTool( + name="mcp_describe", + description=describe_description, + inputSchema={ + "type": "object", + "properties": { + "server": { + "type": "string", + "description": "Optional MCP server name or alias to inspect.", + }, + "tool": { + "type": "string", + "description": "Optional tool name within the selected MCP server.", + }, + }, + }, + ), + MCPTool( + name="mcp_call", + description="Execute a tool from an MCP server available through the LiteLLM LazyMCP gateway. Use mcp_describe first to inspect available servers, tools, and schemas.", + inputSchema={ + "type": "object", + "properties": { + "server": { + "type": "string", + "description": "MCP server name or alias.", + }, + "tool": { + "type": "string", + "description": "Tool name to execute on the selected MCP server.", + }, + "arguments": { + "type": "object", + "description": "Arguments to pass to the selected MCP tool.", + }, + }, + "required": ["server", "tool", "arguments"], + "additionalProperties": False, + }, + ), + MCPTool( + name="mcp_status", + description="Report safe diagnostics for the current LiteLLM LazyMCP view.", + inputSchema={ + "type": "object", + "properties": {}, + "additionalProperties": False, + }, + ), + ] + async def _list_mcp_prompts( user_api_key_auth: Optional[UserAPIKeyAuth] = None, mcp_auth_header: Optional[str] = None, @@ -2722,6 +3362,86 @@ if MCP_AVAILABLE: ] return False + async def _prepare_mcp_request_context( + scope: Scope, + path: str, + ) -> Tuple[ + Optional[UserAPIKeyAuth], + Optional[str], + Optional[List[str]], + Optional[Dict[str, Dict[str, str]]], + Optional[Dict[str, str]], + Optional[Dict[str, str]], + Optional[str], + ]: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) = await extract_mcp_auth_context(scope, path) + + client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) + + for server_name in mcp_servers or []: + server = global_mcp_server_manager.get_mcp_server_by_name( + server_name, client_ip=client_ip + ) + if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers: + if server.needs_user_oauth_token: + stored_oauth_headers = await _get_user_oauth_extra_headers_from_db( + server=server, + user_api_key_auth=user_api_key_auth, + ) + if stored_oauth_headers: + continue + + request = StarletteRequest(scope) + base_url = get_request_base_url(request) + authorization_uri = ( + f"Bearer authorization_uri=" + f"{base_url}/.well-known/oauth-authorization-server/{server_name}" + ) + + raise HTTPException( + status_code=401, + detail="Unauthorized", + headers={"www-authenticate": authorization_uri}, + ) + + scope["headers"] = [ + (k, v) + for k, v in scope.get("headers", []) + if k.lower() != b"x-mcp-toolset-id" + ] + + active_toolset_id = _mcp_active_toolset_id.get() + if active_toolset_id and user_api_key_auth is not None: + user_api_key_auth = await _apply_toolset_scope( + user_api_key_auth, active_toolset_id + ) + + set_auth_context( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + return ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + client_ip, + ) + async def _apply_toolset_scope( user_api_key_auth: UserAPIKeyAuth, toolset_id: str, @@ -2943,10 +3663,8 @@ if MCP_AVAILABLE: mcp_server_auth_headers, oauth2_headers, raw_headers, - ) = await extract_mcp_auth_context(scope, path) - - # Extract client IP for MCP access control - _client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope)) + _client_ip, + ) = await _prepare_mcp_request_context(scope, path) verbose_logger.debug( f"MCP request mcp_servers (header/path): {mcp_servers}" @@ -3023,17 +3741,6 @@ if MCP_AVAILABLE: if _debug_headers: send = MCPDebug.wrap_send_with_debug_headers(send, _debug_headers) - # Set the auth context variable for easy access in MCP functions - set_auth_context( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - client_ip=_client_ip, - ) - # Ensure session managers are initialized if not _SESSION_MANAGERS_INITIALIZED: await initialize_session_managers() @@ -3077,6 +3784,48 @@ if MCP_AVAILABLE: # If we can't send a proper response, re-raise the original error raise e + async def handle_streamable_http_lazymcp( + scope: Scope, receive: Receive, send: Send + ) -> None: + """Handle LazyMCP requests through StreamableHTTP.""" + try: + path = scope.get("path", "") + if path.startswith("/lazymcp/"): + scope["path"] = "/mcp/" + path[len("/lazymcp/") :] + elif path.startswith("/lazymcp"): + scope["path"] = "/mcp" + path[len("/lazymcp") :] + path = scope.get("path", "") + await _prepare_mcp_request_context(scope, path) + + if not _SESSION_MANAGERS_INITIALIZED: + await initialize_session_managers() + await asyncio.sleep(0.1) + + handled = await _handle_stale_mcp_session( + scope, receive, send, lazy_session_manager + ) + if handled: + return + await lazy_session_manager.handle_request(scope, receive, send) + except HTTPException: + raise + except Exception as e: + verbose_logger.exception(f"Error handling LazyMCP request: {e}") + try: + from starlette.responses import JSONResponse + from starlette.status import HTTP_500_INTERNAL_SERVER_ERROR + + error_response = JSONResponse( + status_code=HTTP_500_INTERNAL_SERVER_ERROR, + content={"error": "LazyMCP request failed", "details": str(e)}, + ) + await error_response(scope, receive, send) + except Exception as response_error: + verbose_logger.exception( + f"Failed to send LazyMCP error response: {response_error}" + ) + raise e + async def handle_sse_mcp(scope: Scope, receive: Receive, send: Send) -> None: """Handle MCP requests through SSE.""" try: @@ -3160,6 +3909,8 @@ if MCP_AVAILABLE: # Mount the MCP handlers app.mount("/", handle_streamable_http_mcp) app.mount("/mcp", handle_streamable_http_mcp) + app.mount("/lazymcp", handle_streamable_http_lazymcp) + app.mount("/lazymcp/{mcp_server_name}", handle_streamable_http_lazymcp) app.mount("/{mcp_server_name}/mcp", handle_streamable_http_mcp) app.mount("/sse", handle_sse_mcp) app.add_middleware(AuthContextMiddleware) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 8cb9a11ffee..d115e54cbe7 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -29,6 +29,7 @@ from litellm.proxy.common_utils.callback_utils import ( get_metadata_variable_name_from_kwargs, ) from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy.auth.ip_address_utils import IPAddressUtils # Cache special headers as a frozenset for O(1) lookup performance _SPECIAL_HEADERS_CACHE = frozenset( @@ -1415,7 +1416,12 @@ async def add_litellm_data_to_request( # noqa: PLR0915 if "user" not in data: data["user"] = user - data["secret_fields"] = SecretFields(raw_headers=_raw_headers) + data["secret_fields"] = SecretFields( + raw_headers=_raw_headers, + mcp_client_ip=IPAddressUtils.get_mcp_client_ip( + request, general_settings=general_settings + ), + ) ## Dynamic api version (Azure OpenAI endpoints) ## try: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5d89d3fa9c5..b99fae8bd86 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15639,6 +15639,57 @@ async def _stream_mcp_asgi_response( # Toolset-namespaced MCP routes - handle /toolset/{toolset_name}/mcp # Must be declared BEFORE /{mcp_server_name}/mcp to avoid being swallowed by the catchall. +@app.api_route( + "/toolset/{toolset_name}/lazymcp/", + methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], +) +@app.api_route( + "/toolset/{toolset_name}/lazymcp", + methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], +) +async def toolset_lazymcp_route(toolset_name: str, request: Request): + """Namespace a toolset as its own LazyMCP endpoint.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.proxy._experimental.mcp_server.server import ( + _mcp_active_toolset_id, + handle_streamable_http_lazymcp, + ) + + if prisma_client is None: + raise HTTPException(status_code=503, detail="Database not available") + + toolset = await global_mcp_server_manager.get_toolset_by_name_cached( + prisma_client, toolset_name + ) + if toolset is None: + raise HTTPException( + status_code=404, + detail=f"Toolset '{toolset_name}' not found", + ) + + scope = dict(request.scope) + scope["path"] = "/lazymcp" + + token = _mcp_active_toolset_id.set(toolset.toolset_id) + try: + return await _stream_mcp_asgi_response( + handle_streamable_http_lazymcp, scope, request.receive + ) + finally: + _mcp_active_toolset_id.reset(token) + + except HTTPException as e: + raise e + except Exception as e: + verbose_proxy_logger.error( + f"Error handling toolset LazyMCP route for {toolset_name}: {str(e)}" + ) + raise HTTPException(status_code=500, detail=f"Internal server error: {str(e)}") + + @app.api_route( "/toolset/{toolset_name}/mcp", methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], @@ -15791,6 +15842,88 @@ async def _is_mcp_access_group_cached(name: str) -> bool: return result +@app.api_route( + "/lazymcp/", + methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], +) +@app.api_route( + "/lazymcp", + methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], +) +async def root_lazymcp_route(request: Request): + """Handle root LazyMCP route like /lazymcp.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_lazymcp, + ) + + scope = dict(request.scope) + scope["path"] = "/lazymcp" + return await _stream_mcp_asgi_response( + handle_streamable_http_lazymcp, scope, request.receive + ) + except HTTPException as e: + raise e + except Exception as e: + verbose_proxy_logger.error(f"Error handling root LazyMCP route: {str(e)}") + raise HTTPException(status_code=500, detail=f"Internal server error: {str(e)}") + + +@app.api_route( + "/lazymcp/{mcp_server_name}/", + methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], +) +@app.api_route( + "/lazymcp/{mcp_server_name}", + 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.""" + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.proxy._experimental.mcp_server.server import ( + _mcp_active_toolset_id, + handle_streamable_http_lazymcp, + ) + from litellm.proxy.auth.ip_address_utils import IPAddressUtils + + client_ip = IPAddressUtils.get_mcp_client_ip(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: + toolset = await global_mcp_server_manager.get_toolset_by_name_cached( + prisma_client, mcp_server_name + ) + if toolset is not None: + scope["path"] = "/lazymcp" + token = _mcp_active_toolset_id.set(toolset.toolset_id) + try: + return await _stream_mcp_asgi_response( + handle_streamable_http_lazymcp, scope, request.receive + ) + 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 + ) + except HTTPException as e: + raise e + except Exception as e: + verbose_proxy_logger.error( + f"Error handling dynamic LazyMCP route for {mcp_server_name}: {str(e)}" + ) + raise HTTPException(status_code=500, detail=f"Internal server error: {str(e)}") + + # Dynamic MCP server routes - handle /{mcp_server_name}/mcp @app.api_route( "/{mcp_server_name}/mcp", diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 4ee9235af7d..9eb39979395 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -192,6 +192,7 @@ async def aresponses_api_with_mcp( mcp_auth_header: Optional[str] = None mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None secret_fields = kwargs.get("secret_fields") + client_ip = ResponsesAPIRequestUtils.get_verified_mcp_client_ip(secret_fields) if secret_fields and isinstance(secret_fields, dict): ( mcp_auth_header, @@ -212,6 +213,7 @@ async def aresponses_api_with_mcp( litellm_trace_id=kwargs.get("litellm_trace_id"), mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, + client_ip=client_ip, ) openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( original_mcp_tools @@ -260,6 +262,7 @@ async def aresponses_api_with_mcp( user_api_key_auth=user_api_key_auth, base_item_id=base_item_id, pre_processed_mcp_tools=original_mcp_tools, + client_ip=client_ip, ) return LiteLLM_Proxy_MCP_Handler._create_mcp_streaming_response( @@ -335,6 +338,7 @@ async def aresponses_api_with_mcp( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers_from_request, + client_ip=client_ip, litellm_call_id=kwargs.get("litellm_call_id"), litellm_trace_id=kwargs.get("litellm_trace_id"), ) @@ -399,6 +403,7 @@ async def aresponses_api_with_mcp( mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, + client_ip=client_ip, ) final_response = ( LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response( diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 24b5db28571..233019e2456 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -117,6 +117,9 @@ async def acompletion_with_mcp( # noqa: PLR0915 user_api_key_auth = kwargs.get("user_api_key_auth") or ( (kwargs.get("metadata", {}) or {}).get("user_api_key_auth") ) + client_ip = ResponsesAPIRequestUtils.get_verified_mcp_client_ip( + kwargs.get("secret_fields") + ) # Extract MCP auth headers before fetching tools (needed for dynamic auth) ( @@ -139,6 +142,7 @@ async def acompletion_with_mcp( # noqa: PLR0915 litellm_trace_id=kwargs.get("litellm_trace_id"), mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, + client_ip=client_ip, ) openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( @@ -216,6 +220,7 @@ async def acompletion_with_mcp( # noqa: PLR0915 mcp_server_auth_headers, oauth2_headers, raw_headers, + client_ip, litellm_call_id, litellm_trace_id, openai_tools, @@ -229,6 +234,7 @@ async def acompletion_with_mcp( # noqa: PLR0915 self.mcp_server_auth_headers = mcp_server_auth_headers self.oauth2_headers = oauth2_headers self.raw_headers = raw_headers + self.client_ip = client_ip self.litellm_call_id = litellm_call_id self.litellm_trace_id = litellm_trace_id self.openai_tools = openai_tools @@ -454,6 +460,7 @@ async def acompletion_with_mcp( # noqa: PLR0915 mcp_server_auth_headers=self.mcp_server_auth_headers, oauth2_headers=self.oauth2_headers, raw_headers=self.raw_headers, + client_ip=self.client_ip, litellm_call_id=self.litellm_call_id, litellm_trace_id=self.litellm_trace_id, ) @@ -514,6 +521,7 @@ async def acompletion_with_mcp( # noqa: PLR0915 mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + client_ip=client_ip, litellm_call_id=kwargs.get("litellm_call_id"), litellm_trace_id=kwargs.get("litellm_trace_id"), openai_tools=openai_tools, @@ -635,6 +643,7 @@ async def acompletion_with_mcp( # noqa: PLR0915 mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + client_ip=client_ip, litellm_call_id=kwargs.get("litellm_call_id"), litellm_trace_id=kwargs.get("litellm_trace_id"), ) diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 94cff6922b5..7f3ad81434a 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -1,3 +1,4 @@ +import json import re import traceback from datetime import datetime @@ -43,6 +44,8 @@ ToolParam = Any LITELLM_PROXY_MCP_SERVER_URL = "litellm_proxy" LITELLM_PROXY_MCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/" +LITELLM_PROXY_LAZYMCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/lazymcp/" +LITELLM_PROXY_LAZYMCP_TOOL_SERVER_MAP_PREFIX = "lazymcp:" # Matches any URL whose path ends with /mcp/ — covers both root-path # (http://host:port/mcp/name) and sub-path (http://host/base/mcp/name) proxy deployments. @@ -50,6 +53,7 @@ LITELLM_PROXY_MCP_SERVER_URL_PREFIX = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/" # in a "server not found" error from the internal gateway, not a silent failure or data leak, # so this broad pattern is intentional and preferred over anchoring to localhost only. _PROXY_MCP_PATH_RE = re.compile(r"^https?://.+/mcp/([^/]+)$") +_PROXY_LAZYMCP_PATH_RE = re.compile(r"^https?://.+/lazymcp(?:/([^/]+))?$") class LiteLLM_Proxy_MCP_Handler: @@ -59,6 +63,34 @@ class LiteLLM_Proxy_MCP_Handler: This handles when a user passes mcp server_url="litellm_proxy" in their tools. """ + @staticmethod + def _encode_lazymcp_tool_server_map_value( + mcp_servers: Optional[List[str]], toolset_id: Optional[str] + ) -> str: + payload = {"mcp_servers": mcp_servers or [], "toolset_id": toolset_id} + return f"{LITELLM_PROXY_LAZYMCP_TOOL_SERVER_MAP_PREFIX}{json.dumps(payload, sort_keys=True)}" + + @staticmethod + def _decode_lazymcp_tool_server_map_value( + value: Optional[str], + ) -> Optional[Dict[str, Any]]: + if not isinstance(value, str) or not value.startswith( + LITELLM_PROXY_LAZYMCP_TOOL_SERVER_MAP_PREFIX + ): + return None + try: + decoded = json.loads( + value[len(LITELLM_PROXY_LAZYMCP_TOOL_SERVER_MAP_PREFIX) :] + ) + except Exception: + return {"mcp_servers": [], "toolset_id": None} + if not isinstance(decoded, dict): + return {"mcp_servers": [], "toolset_id": None} + mcp_servers = decoded.get("mcp_servers") + if not isinstance(mcp_servers, list): + decoded["mcp_servers"] = [] + return decoded + @staticmethod def _should_use_litellm_mcp_gateway(tools: Optional[Iterable[ToolParam]]) -> bool: """ @@ -77,6 +109,10 @@ class LiteLLM_Proxy_MCP_Handler: server_url ): return True + if isinstance(server_url, str) and _PROXY_LAZYMCP_PATH_RE.match( + server_url + ): + return True return False @staticmethod @@ -111,7 +147,17 @@ class LiteLLM_Proxy_MCP_Handler: } mcp_tools_with_litellm_proxy.append(rewritten) else: - other_tools.append(tool) + lazy_match = _PROXY_LAZYMCP_PATH_RE.match(server_url) + if lazy_match: + rewritten_url = ( + f"{LITELLM_PROXY_MCP_SERVER_URL}/lazymcp" + ) + if lazy_match.group(1): + rewritten_url = f"{LITELLM_PROXY_LAZYMCP_SERVER_URL_PREFIX}{lazy_match.group(1)}" + rewritten = {**tool, "server_url": rewritten_url} + mcp_tools_with_litellm_proxy.append(rewritten) + else: + other_tools.append(tool) else: other_tools.append(tool) else: @@ -170,38 +216,13 @@ class LiteLLM_Proxy_MCP_Handler: return user_api_key_auth @staticmethod - async def _get_mcp_tools_from_manager( - user_api_key_auth: Any, + def _get_requested_mcp_servers( mcp_tools_with_litellm_proxy: Optional[Iterable[ToolParam]], - litellm_trace_id: Optional[str] = None, - mcp_auth_header: Optional[str] = None, - mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, - ) -> tuple[List[MCPTool], List[str]]: - """ - Get available tools from the MCP server manager. - - Args: - user_api_key_auth: User authentication info for access control - mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy" - mcp_auth_header: Optional deprecated auth header for MCP servers - mcp_server_auth_headers: Optional server-specific auth headers (e.g. from x-mcp-{alias}-*) - - Returns: - List of MCP tools - List names of allowed MCP servers - """ - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - from litellm.proxy._experimental.mcp_server.server import ( - _get_allowed_mcp_servers_from_mcp_server_names, - _get_tools_from_mcp_servers, - ) - + ) -> tuple[List[str], bool]: mcp_servers: List[str] = [] + use_lazymcp = False if mcp_tools_with_litellm_proxy: for _tool in mcp_tools_with_litellm_proxy: - # if user specifies servers as server_url: litellm_proxy/mcp/zapier,github then return zapier,github server_url = ( _tool.get("server_url", "") if isinstance(_tool, dict) else "" ) @@ -209,10 +230,101 @@ class LiteLLM_Proxy_MCP_Handler: LITELLM_PROXY_MCP_SERVER_URL_PREFIX ): mcp_servers.append(server_url.split("/")[-1]) + elif isinstance(server_url, str) and server_url.startswith( + LITELLM_PROXY_LAZYMCP_SERVER_URL_PREFIX + ): + use_lazymcp = True + mcp_servers.append(server_url.split("/")[-1]) + elif server_url == f"{LITELLM_PROXY_MCP_SERVER_URL}/lazymcp": + use_lazymcp = True + return mcp_servers, use_lazymcp + + @staticmethod + async def _resolve_lazymcp_scope( + effective_filter: Optional[List[str]], + global_mcp_server_manager: Any, + ) -> tuple[Optional[List[str]], Optional[str]]: + 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 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}" + ) + return effective_filter, active_toolset_id + + @staticmethod + async def _get_lazymcp_gateway_tools( + user_api_key_auth: Any, + effective_filter: Optional[List[str]], + active_toolset_id: Optional[str], + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + client_ip: Optional[str], + ) -> tuple[List[MCPTool], List[str]]: + from litellm.proxy._experimental.mcp_server.server import ( + _mcp_active_toolset_id, + _apply_toolset_scope, + _get_lazymcp_gateway_tools, + _get_lazymcp_catalog, + ) + + token = ( + _mcp_active_toolset_id.set(active_toolset_id) + if active_toolset_id is not None + else None + ) + try: + if active_toolset_id is not None and user_api_key_auth is not None: + user_api_key_auth = await _apply_toolset_scope( + user_api_key_auth, active_toolset_id + ) + catalog = await _get_lazymcp_catalog( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=effective_filter, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=None, + raw_headers=None, + client_ip=client_ip, + ) + finally: + if token is not None: + _mcp_active_toolset_id.reset(token) + return _get_lazymcp_gateway_tools(catalog.get("description")), [ + LiteLLM_Proxy_MCP_Handler._encode_lazymcp_tool_server_map_value( + effective_filter, active_toolset_id + ) + ] + + @staticmethod + async def _get_standard_mcp_tools( + user_api_key_auth: Any, + mcp_servers: List[str], + global_mcp_server_manager: Any, + mcp_auth_header: Optional[str], + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]], + litellm_trace_id: Optional[str], + client_ip: Optional[str] = None, + ) -> tuple[List[MCPTool], List[str]]: + from litellm.proxy._experimental.mcp_server.server import ( + _get_allowed_mcp_servers_from_mcp_server_names, + _get_tools_from_mcp_servers, + ) - # Resolve toolset names: collect all toolset IDs first, then apply their - # combined permissions in a single pass so multiple toolsets are unioned - # rather than the last one overwriting the others. resolved_mcp_servers: List[str] = [] resolved_toolset_ids: List[str] = [] for name in mcp_servers: @@ -227,7 +339,6 @@ class LiteLLM_Proxy_MCP_Handler: ) ) if toolset is not None: - # Access control: only allow if the key explicitly grants this toolset. if user_api_key_auth is not None: from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_view, @@ -241,8 +352,6 @@ class LiteLLM_Proxy_MCP_Handler: if op else None ) - # None means no grants configured → deny (consistent with - # fetch_mcp_toolsets which returns [] for unconfigured keys) if ( granted is None or toolset.toolset_id not in granted @@ -252,14 +361,11 @@ class LiteLLM_Proxy_MCP_Handler: ) continue resolved_toolset_ids.append(toolset.toolset_id) - # Don't add to resolved_mcp_servers — toolset scope - # restricts via object_permission, not server name filter. continue except Exception as _e: verbose_logger.debug(f"Could not resolve '{name}' as toolset: {_e}") resolved_mcp_servers.append(name) - # Apply all resolved toolsets at once (union), avoiding permission overwrite. if resolved_toolset_ids and user_api_key_auth is not None: user_api_key_auth = ( await LiteLLM_Proxy_MCP_Handler._apply_toolset_permissions( @@ -269,10 +375,6 @@ class LiteLLM_Proxy_MCP_Handler: ) ) - # When toolsets were resolved we updated object_permission.mcp_servers to the - # full union (toolset server IDs + direct server names). Passing a name-based - # filter here would exclude those toolset server IDs (which are UUIDs, not - # names), so use None and let the auth object's mcp_servers do the filtering. effective_server_filter = ( None if resolved_toolset_ids else (resolved_mcp_servers or None) ) @@ -285,6 +387,7 @@ class LiteLLM_Proxy_MCP_Handler: log_list_tools_to_spendlogs=True, list_tools_log_source="responses", litellm_trace_id=litellm_trace_id, + client_ip=client_ip, ) allowed_mcp_server_ids = ( @@ -293,7 +396,6 @@ class LiteLLM_Proxy_MCP_Handler: allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( # type: ignore[attr-defined] allowed_mcp_server_ids ) - allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers=effective_server_filter, allowed_mcp_servers=allowed_mcp_servers, @@ -313,6 +415,66 @@ class LiteLLM_Proxy_MCP_Handler: return tools, server_names + @staticmethod + async def _get_mcp_tools_from_manager( + user_api_key_auth: Any, + mcp_tools_with_litellm_proxy: Optional[Iterable[ToolParam]], + litellm_trace_id: Optional[str] = None, + mcp_auth_header: Optional[str] = None, + mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, + client_ip: Optional[str] = None, + ) -> tuple[List[MCPTool], List[str]]: + """ + Get available tools from the MCP server manager. + + Args: + user_api_key_auth: User authentication info for access control + mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy" + mcp_auth_header: Optional deprecated auth header for MCP servers + mcp_server_auth_headers: Optional server-specific auth headers (e.g. from x-mcp-{alias}-*) + + Returns: + List of MCP tools + List names of allowed MCP servers + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + mcp_servers, use_lazymcp = LiteLLM_Proxy_MCP_Handler._get_requested_mcp_servers( + mcp_tools_with_litellm_proxy + ) + + if use_lazymcp: + effective_filter = mcp_servers or None + active_toolset_id: Optional[str] = None + effective_filter, active_toolset_id = ( + await LiteLLM_Proxy_MCP_Handler._resolve_lazymcp_scope( + effective_filter, global_mcp_server_manager + ) + ) + return await LiteLLM_Proxy_MCP_Handler._get_lazymcp_gateway_tools( + user_api_key_auth=user_api_key_auth, + effective_filter=effective_filter, + active_toolset_id=active_toolset_id, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + client_ip=client_ip, + ) + + standard_client_ip = ( + None if client_ip == "__invalid_mcp_client_ip__" else client_ip + ) + return await LiteLLM_Proxy_MCP_Handler._get_standard_mcp_tools( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + global_mcp_server_manager=global_mcp_server_manager, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + litellm_trace_id=litellm_trace_id, + client_ip=standard_client_ip, + ) + @staticmethod def _deduplicate_mcp_tools( mcp_tools: List[MCPTool], allowed_mcp_servers: List[str] @@ -427,6 +589,7 @@ class LiteLLM_Proxy_MCP_Handler: litellm_trace_id: Optional[str] = None, mcp_auth_header: Optional[str] = None, mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, + client_ip: Optional[str] = None, ) -> tuple[List[Any], dict[str, str]]: """ Process MCP tools through filtering and deduplication pipeline without OpenAI transformation. @@ -454,6 +617,7 @@ class LiteLLM_Proxy_MCP_Handler: litellm_trace_id=litellm_trace_id, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, + client_ip=client_ip, ) # Step 2: Filter tools based on allowed_tools parameter @@ -652,6 +816,7 @@ class LiteLLM_Proxy_MCP_Handler: mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, + client_ip: Optional[str] = None, litellm_call_id: Optional[str] = None, litellm_trace_id: Optional[str] = None, ) -> List[Dict[str, Any]]: @@ -689,6 +854,63 @@ class LiteLLM_Proxy_MCP_Handler: # Import here to avoid circular import from litellm.proxy.proxy_server import proxy_logging_obj + lazymcp_scope = ( + LiteLLM_Proxy_MCP_Handler._decode_lazymcp_tool_server_map_value( + tool_server_map.get(tool_name) + ) + ) + if ( + tool_name in {"mcp_describe", "mcp_call", "mcp_status"} + and lazymcp_scope is not None + ): + from litellm.proxy._experimental.mcp_server.server import ( + _mcp_active_toolset_id, + _apply_toolset_scope, + lazymcp_tool_call, + set_auth_context, + ) + + lazy_mcp_servers = lazymcp_scope.get("mcp_servers") or None + if not isinstance(lazy_mcp_servers, list): + lazy_mcp_servers = None + lazy_toolset_id = lazymcp_scope.get("toolset_id") + scoped_user_api_key_auth = user_api_key_auth + if ( + isinstance(lazy_toolset_id, str) + and user_api_key_auth is not None + ): + scoped_user_api_key_auth = await _apply_toolset_scope( + user_api_key_auth, lazy_toolset_id + ) + token = ( + _mcp_active_toolset_id.set(lazy_toolset_id) + if isinstance(lazy_toolset_id, str) + else None + ) + try: + set_auth_context( + user_api_key_auth=scoped_user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=lazy_mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + result = await lazymcp_tool_call(tool_name, parsed_arguments) + finally: + if token is not None: + _mcp_active_toolset_id.reset(token) + result_text = LiteLLM_Proxy_MCP_Handler._parse_mcp_result(result) + tool_results.append( + { + "tool_call_id": tool_call_id, + "result": result_text, + "name": tool_name, + } + ) + continue + server_name = tool_server_map[tool_name] # Remove the server name prefix if the tool name includes it. diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index 42c46dff47c..a01cedde4e4 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -30,6 +30,7 @@ async def create_mcp_list_tools_events( user_api_key_auth: Any, base_item_id: str, pre_processed_mcp_tools: List[Any], + client_ip: Optional[str] = None, ) -> List[ResponsesAPIStreamingResponse]: """Create MCP discovery events using pre-processed tools from the parent""" @@ -325,6 +326,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): secret_fields = self.original_request_params.get("secret_fields") if secret_fields and isinstance(secret_fields, dict): raw_headers_from_request = secret_fields.get("raw_headers") + from litellm.responses.utils import ResponsesAPIRequestUtils + + self.client_ip = ResponsesAPIRequestUtils.get_verified_mcp_client_ip( + secret_fields + ) # Extract MCP-specific headers self.mcp_auth_header: Optional[str] = None @@ -666,6 +672,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): mcp_server_auth_headers=self.mcp_server_auth_headers, oauth2_headers=self.oauth2_headers, raw_headers=self.raw_headers, + client_ip=getattr(self, "client_ip", None), litellm_call_id=self.litellm_call_id, litellm_trace_id=self.litellm_trace_id, ) diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 74e4d7a533a..f2ccdc4cc7f 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -885,6 +885,23 @@ class ResponsesAPIRequestUtils: raw_headers_from_request, ) + @staticmethod + def get_verified_mcp_client_ip( + secret_fields: Optional[Dict[str, Any]], + ) -> str: + """Return the verified MCP client IP or a fail-closed sentinel. + + LazyMCP access control uses this value for IP filtering. When no verified + IP is available, return a non-None sentinel so internal-only servers stay + hidden instead of bypassing filtering. + """ + + if secret_fields and isinstance(secret_fields, dict): + client_ip = secret_fields.get("mcp_client_ip") + if isinstance(client_ip, str) and client_ip.strip(): + return client_ip.strip() + return "__invalid_mcp_client_ip__" + class ResponseAPILoggingUtils: @staticmethod diff --git a/litellm/types/proxy/litellm_pre_call_utils.py b/litellm/types/proxy/litellm_pre_call_utils.py index e0d3f3dac66..8f25d593cb3 100644 --- a/litellm/types/proxy/litellm_pre_call_utils.py +++ b/litellm/types/proxy/litellm_pre_call_utils.py @@ -1,3 +1,5 @@ +from typing import Optional + from typing_extensions import TypedDict @@ -22,3 +24,4 @@ class SecretFields(TypedDict): """ raw_headers: dict + mcp_client_ip: Optional[str] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index e1eddfc9c7a..8412ff671ff 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1,9 +1,12 @@ import asyncio +import json +import types from datetime import datetime, timedelta from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException +from fastapi.testclient import TestClient from mcp import ReadResourceResult, Resource from mcp.types import ( BlobResourceContents, @@ -13,14 +16,58 @@ from mcp.types import ( ) from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, LiteLLM_MCPServerTable, MCPTransport, UserAPIKeyAuth, ) +from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer +def _parse_mcp_text_result(result): + return json.loads(result.content[0].text) + + +def test_lazymcp_root_trailing_slash_does_not_redirect(monkeypatch): + from litellm.proxy.proxy_server import app + + async def fake_stream_response(handle_fn, scope, receive): + from starlette.responses import Response + + assert scope["path"] == "/lazymcp" + return Response("ok", media_type="text/event-stream") + + monkeypatch.setattr( + "litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response + ) + + response = TestClient(app).get("/lazymcp/", follow_redirects=False) + + assert response.status_code == 200 + assert "location" not in response.headers + + +def test_lazymcp_root_without_trailing_slash_still_routes(monkeypatch): + from litellm.proxy.proxy_server import app + + async def fake_stream_response(handle_fn, scope, receive): + from starlette.responses import Response + + assert scope["path"] == "/lazymcp" + return Response("ok", media_type="text/event-stream") + + monkeypatch.setattr( + "litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response + ) + + response = TestClient(app).get("/lazymcp", follow_redirects=False) + + assert response.status_code == 200 + assert "location" not in response.headers + + @pytest.fixture(autouse=True) def cleanup_mcp_global_state(): """Clean up MCP global state before and after each test. @@ -109,6 +156,943 @@ async def test_mcp_server_tool_call_body_contains_request_data(): assert body["arguments"] == tool_arguments +@pytest.mark.asyncio +async def test_lazymcp_lists_only_gateway_tools_and_describes_visible_servers(): + try: + from mcp.types import Tool as MCPTool + + from litellm.proxy._experimental.mcp_server.server import ( + _get_lazymcp_gateway_tools, + _lazymcp_describe, + set_auth_context, + ) + except ImportError: + pytest.skip("MCP server not available") + + visible_server = MCPServer( + server_id="visible-server", + name="memory", + alias="memory", + server_name="memory", + transport=MCPTransport.http, + mcp_info={"description": "Long-term memory search and storage."}, + ) + hidden_server = MCPServer( + server_id="hidden-server", + name="hidden", + alias="hidden", + server_name="hidden", + transport=MCPTransport.http, + mcp_info={"description": "Should not leak."}, + ) + tool = MCPTool( + name="memory-search", + description="Search memory", + inputSchema={"type": "object", "properties": {"query": {"type": "string"}}}, + ) + + set_auth_context(UserAPIKeyAuth(api_key="sk-test", user_id="user")) + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[visible_server]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", + AsyncMock(return_value=[tool]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._lazymcp_cache_get", + AsyncMock(return_value=None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._lazymcp_cache_set", + AsyncMock(), + ), + ): + gateway_tools = _get_lazymcp_gateway_tools("catalog") + catalog = await _lazymcp_describe({}) + server_detail = await _lazymcp_describe({"server": "memory"}) + tool_detail = await _lazymcp_describe( + {"server": "memory", "tool": "memory-search"} + ) + legacy_tool_detail = await _lazymcp_describe( + { + "server": "memory", + "tool": "memory-search", + "include_schemas": True, + } + ) + hidden_detail = await _lazymcp_describe({"server": hidden_server.name}) + + assert [tool.name for tool in gateway_tools] == [ + "mcp_describe", + "mcp_call", + "mcp_status", + ] + assert catalog["servers"] == [ + { + "name": "memory", + "description": "Long-term memory search and storage.", + "tool_count": 1, + } + ] + assert server_detail["tools"][0]["name"] == "memory-search" + assert server_detail["tools"][0]["input_schema"] == tool.inputSchema + assert "input_schema_summary" not in server_detail["tools"][0] + assert tool_detail["tool"]["input_schema"] == tool.inputSchema + assert legacy_tool_detail == tool_detail + assert "include_schemas" not in gateway_tools[0].inputSchema.get("properties", {}) + assert hidden_detail == {"error": "MCP server is not available for this request."} + + +@pytest.mark.asyncio +async def test_lazymcp_describe_denies_stale_cached_server_after_access_revoked(): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _lazymcp_describe, + set_auth_context, + ) + except ImportError: + pytest.skip("MCP server not available") + + cached_catalog = { + "description": "cached catalog", + "server_count": 1, + "tool_count": 1, + "servers": [ + { + "name": "revoked", + "description": "stale cached server", + "tool_count": 1, + "tools": [ + { + "name": "revoked-tool", + "description": "must not leak", + "input_schema": {"type": "object"}, + } + ], + } + ], + } + + set_auth_context(UserAPIKeyAuth(api_key="sk-test", user_id="user")) + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._lazymcp_cache_get", + AsyncMock(return_value=cached_catalog), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", + AsyncMock(side_effect=AssertionError("stale server should not load tools")), + ), + ): + server_detail = await _lazymcp_describe({"server": "revoked"}) + tool_detail = await _lazymcp_describe( + {"server": "revoked", "tool": "revoked-tool"} + ) + + assert server_detail == {"error": "MCP server is not available for this request."} + assert tool_detail == {"error": "MCP server is not available for this request."} + + +@pytest.mark.asyncio +async def test_lazymcp_call_rechecks_permissions_and_delegates_to_mcp_call(): + try: + from mcp.types import Tool as MCPTool + + from litellm.proxy._experimental.mcp_server.server import ( + _lazymcp_call, + set_auth_context, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="server-id", + name="github", + alias="github", + server_name="github", + transport=MCPTransport.http, + ) + tool = MCPTool( + name="github-create_issue", + description="Create issue", + inputSchema={"type": "object"}, + ) + delegated_result = MagicMock() + delegated_result.content = [] + + set_auth_context(UserAPIKeyAuth(api_key="sk-test", user_id="user")) + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + AsyncMock(return_value=[server]), + ) as allowed_mock, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", + AsyncMock(return_value=[tool]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + AsyncMock(return_value=delegated_result), + ) as call_mock, + ): + result = await _lazymcp_call( + { + "server": "github", + "tool": "github-create_issue", + "arguments": {"title": "Bug"}, + } + ) + + assert result is delegated_result + assert allowed_mock.await_count >= 1 + call_mock.assert_awaited_once() + assert call_mock.await_args.kwargs["name"] == "github-create_issue" + + +@pytest.mark.asyncio +async def test_lazymcp_call_uses_unavailable_server_error_when_server_missing(): + try: + from litellm.proxy._experimental.mcp_server.server import _lazymcp_call + except ImportError: + pytest.skip("MCP server not available") + + with patch( + "litellm.proxy._experimental.mcp_server.server._get_lazymcp_allowed_servers", + AsyncMock(return_value=[]), + ): + result = await _lazymcp_call( + {"server": "missing", "tool": "tool", "arguments": {}} + ) + + payload = json.loads(result.content[0].text) + assert payload["error"] == "MCP server is not available for this request." + + +def test_lazymcp_cache_get_set_and_invalidate_paths(monkeypatch): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server_module + except ImportError: + pytest.skip("MCP server not available") + + cache_dict = {"lazymcp:catalog:abc": {"servers": []}, "other": "value"} + fake_cache = MagicMock(in_memory_cache=MagicMock(cache_dict=cache_dict)) + fake_cache.async_get_cache = AsyncMock(return_value={"cached": True}) + fake_cache.async_set_cache = AsyncMock() + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", + fake_cache, + raising=False, + ) + + assert asyncio.run(mcp_server_module._lazymcp_cache_get("lazymcp:catalog:abc")) == { + "cached": True + } + asyncio.run(mcp_server_module._lazymcp_cache_set("lazymcp:catalog:abc", {"x": 1})) + fake_cache.async_set_cache.assert_awaited_once_with( + key="lazymcp:catalog:abc", + value={"x": 1}, + ttl=mcp_server_module.LAZYMCP_CACHE_TTL_SECONDS, + ) + + mcp_server_module.invalidate_lazymcp_cache() + + assert "lazymcp:catalog:abc" not in cache_dict + assert cache_dict["other"] == "value" + + +def test_lazymcp_cache_helpers_tolerate_cache_errors(monkeypatch): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server_module + except ImportError: + pytest.skip("MCP server not available") + + fake_cache = MagicMock(in_memory_cache=MagicMock(cache_dict={})) + fake_cache.async_get_cache = AsyncMock(side_effect=RuntimeError("get failed")) + fake_cache.async_set_cache = AsyncMock(side_effect=RuntimeError("set failed")) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", + fake_cache, + raising=False, + ) + + assert asyncio.run(mcp_server_module._lazymcp_cache_get("lazymcp:broken")) is None + asyncio.run(mcp_server_module._lazymcp_cache_set("lazymcp:broken", {})) + fake_cache.async_set_cache.assert_awaited_once() + + +def test_invalidating_toolset_cache_tolerates_lazymcp_invalidation_error(): + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + with ( + patch( + "litellm.proxy.proxy_server.user_api_key_cache", + MagicMock(in_memory_cache=MagicMock(cache_dict={})), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.invalidate_lazymcp_cache", + MagicMock(side_effect=RuntimeError("lazy failed")), + ), + ): + manager.invalidate_toolset_cache(toolset_id="abc") + + +@pytest.mark.asyncio +async def test_lazymcp_session_manager_shutdown_resets_lazy_manager(monkeypatch): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server_module + except ImportError: + pytest.skip("MCP server not available") + + class _AsyncExit: + def __init__(self): + self.closed = False + + async def __aexit__(self, *_args): + self.closed = True + + lazy_manager = _AsyncExit() + monkeypatch.setattr(mcp_server_module, "_SESSION_MANAGERS_INITIALIZED", True) + monkeypatch.setattr(mcp_server_module, "_session_manager_cm", None) + monkeypatch.setattr(mcp_server_module, "_sse_session_manager_cm", None) + monkeypatch.setattr(mcp_server_module, "_lazy_session_manager_cm", lazy_manager) + + await mcp_server_module.shutdown_session_managers() + + assert lazy_manager.closed is True + assert mcp_server_module._lazy_session_manager_cm is None + assert mcp_server_module._SESSION_MANAGERS_INITIALIZED is False + + +@pytest.mark.asyncio +async def test_lazymcp_lifespan_initializes_and_shuts_down(monkeypatch): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server_module + except ImportError: + pytest.skip("MCP server not available") + + events = [] + + async def fake_initialize(): + events.append("initialize") + + async def fake_shutdown(): + events.append("shutdown") + + monkeypatch.setattr( + mcp_server_module, "initialize_session_managers", fake_initialize + ) + monkeypatch.setattr(mcp_server_module, "shutdown_session_managers", fake_shutdown) + + async with mcp_server_module.lifespan(object()): + events.append("inside") + + assert events == ["initialize", "inside", "shutdown"] + + +@pytest.mark.asyncio +async def test_lazymcp_list_tools_endpoint_returns_gateway_tools_and_fallback( + monkeypatch, +): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server_module + except ImportError: + pytest.skip("MCP server not available") + + monkeypatch.setattr( + mcp_server_module, + "get_auth_context", + MagicMock(return_value=(None, None, None, None, None, None, "127.0.0.1")), + ) + monkeypatch.setattr( + mcp_server_module, + "_get_lazymcp_catalog", + AsyncMock(return_value={"description": "custom catalog"}), + ) + + tools = await mcp_server_module.list_lazymcp_tools() + + assert tools[0].name == "mcp_describe" + assert tools[0].description == "custom catalog" + + monkeypatch.setattr( + mcp_server_module, + "_get_lazymcp_catalog", + AsyncMock(side_effect=RuntimeError("catalog failed")), + ) + + fallback_tools = await mcp_server_module.list_lazymcp_tools() + + assert fallback_tools[0].name == "mcp_describe" + + +@pytest.mark.asyncio +async def test_lazymcp_tool_call_dispatches_and_handles_errors(monkeypatch): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server_module + except ImportError: + pytest.skip("MCP server not available") + + monkeypatch.setattr( + mcp_server_module, + "_lazymcp_describe", + AsyncMock(return_value={"ok": "describe"}), + ) + monkeypatch.setattr( + mcp_server_module, + "_lazymcp_status", + AsyncMock(return_value={"ok": "status"}), + ) + delegated_result = MagicMock() + monkeypatch.setattr( + mcp_server_module, + "_lazymcp_call", + AsyncMock(return_value=delegated_result), + ) + + describe_result = await mcp_server_module.lazymcp_tool_call("mcp_describe", {}) + status_result = await mcp_server_module.lazymcp_tool_call("mcp_status", {}) + call_result = await mcp_server_module.lazymcp_tool_call("mcp_call", {}) + unknown_result = await mcp_server_module.lazymcp_tool_call("unknown", {}) + + assert json.loads(describe_result.content[0].text) == {"ok": "describe"} + assert json.loads(status_result.content[0].text) == {"ok": "status"} + assert call_result is delegated_result + assert json.loads(unknown_result.content[0].text) == { + "error": "Unknown LazyMCP tool." + } + + monkeypatch.setattr( + mcp_server_module, + "_lazymcp_status", + AsyncMock(side_effect=RuntimeError("status failed")), + ) + + error_result = await mcp_server_module.lazymcp_tool_call("mcp_status", {}) + + assert error_result.isError is True + assert "status failed" in json.loads(error_result.content[0].text)["details"] + + +@pytest.mark.asyncio +async def test_lazymcp_toolset_route_sets_scope_and_streams(monkeypatch): + from litellm.proxy.proxy_server import app + + async def fake_get_toolset(_prisma_client, toolset_name): + return types.SimpleNamespace(toolset_id=f"toolset-{toolset_name}") + + async def fake_stream_response(_handle_fn, scope, _receive): + from litellm.proxy._experimental.mcp_server.server import _mcp_active_toolset_id + from starlette.responses import Response + + assert scope["path"] == "/lazymcp" + assert _mcp_active_toolset_id.get() == "toolset-dev" + return Response("ok", media_type="text/event-stream") + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object()) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_toolset_by_name_cached", + fake_get_toolset, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response + ) + + response = TestClient(app).get("/toolset/dev/lazymcp", follow_redirects=False) + + assert response.status_code == 200 + + +def test_lazymcp_root_route_streams_without_redirect(monkeypatch): + from litellm.proxy.proxy_server import app + + async def fake_stream_response(_handle_fn, scope, _receive): + from starlette.responses import Response + + assert scope["path"] == "/lazymcp" + return Response("ok", media_type="text/event-stream") + + monkeypatch.setattr( + "litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response + ) + + response = TestClient(app).get("/lazymcp", follow_redirects=False) + + assert response.status_code == 200 + + +def test_lazymcp_root_route_returns_500_on_unexpected_error(monkeypatch): + from litellm.proxy.proxy_server import app + + async def fake_stream_response(_handle_fn, _scope, _receive): + raise RuntimeError("boom") + + monkeypatch.setattr( + "litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response + ) + + response = TestClient(app).get("/lazymcp", follow_redirects=False) + + assert response.status_code == 500 + assert "boom" in response.json()["detail"] + + +@pytest.mark.asyncio +async def test_lazymcp_dynamic_route_handles_toolset_and_fallback(monkeypatch): + from litellm.proxy.proxy_server import app + + async def fake_get_toolset(_prisma_client, toolset_name): + return types.SimpleNamespace(toolset_id=f"toolset-{toolset_name}") + + async def fake_stream_response(_handle_fn, scope, _receive): + from litellm.proxy._experimental.mcp_server.server import _mcp_active_toolset_id + from starlette.responses import Response + + assert scope["path"] == "/lazymcp" + assert _mcp_active_toolset_id.get() == "toolset-dev" + 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", + MagicMock(return_value="127.0.0.1"), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_by_name", + MagicMock(return_value=None), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_toolset_by_name_cached", + fake_get_toolset, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response + ) + + response = TestClient(app).get("/lazymcp/dev", follow_redirects=False) + + assert response.status_code == 200 + + +def test_lazymcp_dynamic_route_falls_back_for_non_toolset(monkeypatch): + from litellm.proxy.proxy_server import app + + 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", + MagicMock(return_value="127.0.0.1"), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_mcp_server_by_name", + MagicMock(return_value=None), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_toolset_by_name_cached", + fake_get_toolset, + ) + monkeypatch.setattr( + "litellm.proxy.proxy_server._stream_mcp_asgi_response", fake_stream_response + ) + + response = TestClient(app).get("/lazymcp/github", follow_redirects=False) + + assert response.status_code == 200 + + +def test_lazymcp_toolset_route_returns_404_for_missing_toolset(monkeypatch): + from litellm.proxy.proxy_server import app + + async def fake_get_toolset(_prisma_client, _toolset_name): + return None + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", object()) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.get_toolset_by_name_cached", + fake_get_toolset, + ) + + response = TestClient(app).get("/toolset/missing/lazymcp", follow_redirects=False) + + assert response.status_code == 404 + + +def test_lazymcp_toolset_route_returns_503_without_database(monkeypatch): + from litellm.proxy.proxy_server import app + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + + response = TestClient(app).get("/toolset/dev/lazymcp", follow_redirects=False) + + assert response.status_code == 503 + + +@pytest.mark.asyncio +async def test_invalidating_toolset_cache_clears_lazymcp_cache(monkeypatch): + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + cache_dict = {"toolset_perms:abc": 1, "lazymcp:catalog:xyz": 2, "other": 3} + invalidate_mock = MagicMock() + with ( + patch( + "litellm.proxy.proxy_server.user_api_key_cache", + MagicMock(in_memory_cache=MagicMock(cache_dict=cache_dict)), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.invalidate_lazymcp_cache", + invalidate_mock, + ), + ): + manager.invalidate_toolset_cache(toolset_id="abc") + + assert "toolset_perms:abc" not in cache_dict + assert cache_dict["other"] == 3 + invalidate_mock.assert_called_once() + + +def test_lazymcp_cache_scope_hashes_auth_context_and_route_scope(): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _lazymcp_cache_scope, + ) + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="sk-secret", user_id="user") + base_scope = _lazymcp_cache_scope( + user_api_key_auth=user_auth, + mcp_auth_header="Bearer upstream-one", + mcp_servers=["github"], + mcp_server_auth_headers={"github": {"authorization": "Bearer a"}}, + oauth2_headers={"Authorization": "Bearer user-a"}, + raw_headers={"x-mcp-servers": "github"}, + client_ip="127.0.0.1", + ) + changed_auth_scope = _lazymcp_cache_scope( + user_api_key_auth=user_auth, + mcp_auth_header="Bearer upstream-two", + mcp_servers=["github"], + mcp_server_auth_headers={"github": {"authorization": "Bearer a"}}, + oauth2_headers={"Authorization": "Bearer user-a"}, + raw_headers={"x-mcp-servers": "github"}, + client_ip="127.0.0.1", + ) + changed_route_scope = _lazymcp_cache_scope( + user_api_key_auth=user_auth, + mcp_auth_header="Bearer upstream-one", + mcp_servers=["slack"], + mcp_server_auth_headers={"github": {"authorization": "Bearer a"}}, + oauth2_headers={"Authorization": "Bearer user-a"}, + raw_headers={"x-mcp-servers": "github"}, + client_ip="127.0.0.1", + ) + + assert base_scope != changed_auth_scope + assert base_scope != changed_route_scope + assert "sk-secret" not in base_scope + assert "upstream-one" not in base_scope + + +def test_lazymcp_summary_helpers_sanitize_and_summarize(): + try: + from mcp.types import Tool as MCPTool + + from litellm.proxy._experimental.mcp_server.server import ( + _get_lazymcp_server_description, + _lazymcp_tool_to_summary, + _safe_lazymcp_text, + _summarize_lazymcp_schema, + ) + except ImportError: + pytest.skip("MCP server not available") + + long_text = "See https://example.com/" + (" very long" * 30) + assert "[url]" in _safe_lazymcp_text(long_text, "fallback") + assert _safe_lazymcp_text(None, "fallback") == "fallback" + assert _summarize_lazymcp_schema(None) == {} + + schema = { + "type": "object", + "required": ["query"], + "properties": {"query": {"type": "string"}}, + } + assert _summarize_lazymcp_schema(schema) == { + "type": "object", + "required": ["query"], + "properties": ["query"], + } + + tool = MCPTool(name="search", description=long_text, inputSchema=schema) + summary = _lazymcp_tool_to_summary(tool) + detailed_summary = _lazymcp_tool_to_summary(tool, include_schema=True) + assert summary["input_schema_summary"]["properties"] == ["query"] + assert detailed_summary["input_schema"] == schema + + server = MCPServer( + server_id="server-id", + name="fallback-name", + alias=None, + server_name=None, + transport=MCPTransport.http, + mcp_info={"description": long_text}, + ) + assert "[url]" in _get_lazymcp_server_description(server) + + +def test_verified_mcp_client_ip_returns_value_or_fail_closed(): + assert ( + ResponsesAPIRequestUtils.get_verified_mcp_client_ip( + {"mcp_client_ip": " 10.0.0.8 "} + ) + == "10.0.0.8" + ) + assert ( + ResponsesAPIRequestUtils.get_verified_mcp_client_ip({}) + == "__invalid_mcp_client_ip__" + ) + + +@pytest.mark.asyncio +async def test_lazymcp_listing_uses_stored_oauth_when_request_header_missing(): + try: + from mcp.types import Tool as MCPTool + + from litellm.proxy._experimental.mcp_server.server import ( + _get_lazymcp_server_tools, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="oauth-server", + name="oauth", + alias="oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + tool = MCPTool(name="oauth-search", description="Search", inputSchema={}) + user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + AsyncMock(return_value={"Authorization": "Bearer stored-user-token"}), + ) as oauth_mock, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", + AsyncMock(return_value=[tool]), + ) as list_tools_mock, + ): + tools = await _get_lazymcp_server_tools( + server=server, + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers=None, + ) + + assert tools == [tool] + oauth_mock.assert_awaited_once() + assert list_tools_mock.await_args.kwargs["extra_headers"] == { + "Authorization": "Bearer stored-user-token" + } + + +@pytest.mark.asyncio +async def test_lazymcp_toolset_scope_applies_once_for_non_admin_key(): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_lazymcp_allowed_servers, + _mcp_active_toolset_id, + _prepare_mcp_request_context, + ) + except ImportError: + pytest.skip("MCP server not available") + + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="perm-toolset", + mcp_toolsets=["toolset-123"], + ) + user_auth = UserAPIKeyAuth( + api_key="sk-test", + user_id="user", + object_permission=object_permission, + ) + scoped_auth_seen = {} + server = MCPServer( + server_id="server-id", + name="github", + alias="github", + server_name="github", + transport=MCPTransport.http, + ) + + async def fake_apply_toolset_scope(auth, toolset_id): + scoped_auth_seen["toolset_id"] = toolset_id + updated_op = auth.object_permission.model_copy( + update={ + "mcp_servers": ["server-id"], + "mcp_tool_permissions": {"server-id": ["search"]}, + "mcp_toolsets": [], + } + ) + return auth.model_copy(update={"object_permission": updated_op}) + + async def fake_get_allowed_mcp_servers(user_api_key_auth): + scoped_auth_seen["mcp_toolsets_after_scope"] = ( + user_api_key_auth.object_permission.mcp_toolsets + ) + return ["server-id"] + + token = _mcp_active_toolset_id.set("toolset-123") + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._apply_toolset_scope", + AsyncMock(side_effect=fake_apply_toolset_scope), + ) as apply_scope_mock, + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.process_mcp_request", + AsyncMock( + return_value=( + user_auth, + None, + None, + None, + None, + {}, + ) + ), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.IPAddressUtils.get_mcp_client_ip", + MagicMock(return_value="127.0.0.1"), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers", + AsyncMock(side_effect=fake_get_allowed_mcp_servers), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id", + MagicMock(return_value=server), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.filter_server_ids_by_ip_with_info", + MagicMock(return_value=(["server-id"], 0)), + ), + ): + scoped_auth, _, _, _, _, _, client_ip = await _prepare_mcp_request_context( + {"type": "http", "path": "/mcp", "headers": []}, "/mcp" + ) + allowed_servers = await _get_lazymcp_allowed_servers( + user_api_key_auth=scoped_auth, + mcp_servers=None, + client_ip=client_ip, + ) + finally: + _mcp_active_toolset_id.reset(token) + + assert [server.server_id for server in allowed_servers] == ["server-id"] + apply_scope_mock.assert_awaited_once() + assert scoped_auth_seen == { + "toolset_id": "toolset-123", + "mcp_toolsets_after_scope": [], + } + + +@pytest.mark.asyncio +async def test_lazymcp_access_group_route_uses_existing_group_resolution(): + try: + from mcp.types import Tool as MCPTool + + from litellm.proxy._experimental.mcp_server.server import ( + _get_lazymcp_catalog, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="server-group-id", + name="github", + alias="github", + server_name="github", + transport=MCPTransport.http, + mcp_info={"description": "Group visible server."}, + ) + tool = MCPTool(name="github-search", description="Search", inputSchema={}) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers", + AsyncMock(return_value=["server-group-id"]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id", + MagicMock(return_value=server), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.filter_server_ids_by_ip_with_info", + MagicMock(return_value=(["server-group-id"], 0)), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler._get_mcp_servers_from_access_groups", + AsyncMock(return_value=["server-group-id"]), + ) as group_mock, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", + AsyncMock(return_value=[tool]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._lazymcp_cache_get", + AsyncMock(return_value=None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._lazymcp_cache_set", + AsyncMock(), + ), + ): + catalog = await _get_lazymcp_catalog( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-test", user_id="user"), + mcp_auth_header=None, + mcp_servers=["devtools"], + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers=None, + client_ip="127.0.0.1", + ) + + group_mock.assert_awaited_once_with(["devtools"]) + assert catalog["servers"][0]["name"] == "github" + + def test_prepare_mcp_server_headers_case_insensitive_extra_headers(): try: from litellm.proxy._experimental.mcp_server.server import ( diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index 94712813783..0155802415e 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1,5 +1,6 @@ import sys import types +import inspect from unittest.mock import AsyncMock, MagicMock import pytest @@ -10,6 +11,7 @@ from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) from typing import Any, cast +from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse from litellm.types.utils import ModelResponse from litellm.types.responses.main import OutputFunctionToolCall @@ -213,6 +215,207 @@ def test_create_follow_up_input_handles_response_function_tool_call(): ] +def test_parse_mcp_tools_recognizes_lazymcp_urls(): + tools, other_tools = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools( + [ + {"type": "mcp", "server_url": "https://host.example/lazymcp"}, + {"type": "mcp", "server_url": "https://host.example/lazymcp/github"}, + {"type": "mcp", "server_url": "https://host.example/mcp/github"}, + ] + ) + + assert other_tools == [] + assert [tool["server_url"] for tool in tools] == [ + "litellm_proxy/lazymcp", + "litellm_proxy/lazymcp/github", + "litellm_proxy/mcp/github", + ] + + +def test_should_use_litellm_mcp_gateway_callable_as_static_method(): + assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway( + [{"type": "mcp", "server_url": "litellm_proxy/lazymcp/github"}] + ) + + +def test_decode_lazymcp_tool_server_map_value_handles_invalid_payloads(): + assert LiteLLM_Proxy_MCP_Handler._decode_lazymcp_tool_server_map_value(None) is None + assert ( + LiteLLM_Proxy_MCP_Handler._decode_lazymcp_tool_server_map_value("not-lazymcp") + is None + ) + assert LiteLLM_Proxy_MCP_Handler._decode_lazymcp_tool_server_map_value( + "lazymcp:not-json" + ) == {"mcp_servers": [], "toolset_id": None} + assert LiteLLM_Proxy_MCP_Handler._decode_lazymcp_tool_server_map_value( + "lazymcp:[]" + ) == {"mcp_servers": [], "toolset_id": None} + assert LiteLLM_Proxy_MCP_Handler._decode_lazymcp_tool_server_map_value( + 'lazymcp:{"mcp_servers":"github"}' + ) == {"mcp_servers": []} + + +def test_should_use_litellm_mcp_gateway_matches_proxy_urls(): + assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway( + [{"type": "mcp", "server_url": "https://proxy.example/mcp/github"}] + ) + assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway( + [{"type": "mcp", "server_url": "https://proxy.example/lazymcp/github"}] + ) + assert not LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway( + [{"type": "function", "server_url": "https://proxy.example/lazymcp/github"}] + ) + + +def test_get_requested_mcp_servers_handles_lazymcp_variants(): + servers, use_lazymcp = LiteLLM_Proxy_MCP_Handler._get_requested_mcp_servers( + [ + {"type": "mcp", "server_url": "litellm_proxy/mcp/github"}, + {"type": "mcp", "server_url": "litellm_proxy/lazymcp/slack"}, + {"type": "mcp", "server_url": "litellm_proxy/lazymcp"}, + ] + ) + + assert servers == ["github", "slack"] + assert use_lazymcp is True + + +@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_toolset_by_name_cached=AsyncMock( + side_effect=[ + types.SimpleNamespace(toolset_id="toolset-1"), + RuntimeError("db"), + ] + ), + ) + proxy_module = types.SimpleNamespace(prisma_client=object()) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + + assert await LiteLLM_Proxy_MCP_Handler._resolve_lazymcp_scope( + ["github"], server_manager + ) == (["github"], None) + assert await LiteLLM_Proxy_MCP_Handler._resolve_lazymcp_scope( + ["toolset"], server_manager + ) == (None, "toolset-1") + assert await LiteLLM_Proxy_MCP_Handler._resolve_lazymcp_scope( + ["broken"], server_manager + ) == (["broken"], None) + + +@pytest.mark.asyncio +async def test_lazymcp_catalog_uses_verified_client_ip(monkeypatch): + captured = {} + + async def fake_get_lazymcp_catalog(**kwargs): + captured.update(kwargs) + return {"description": "ok"} + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._get_lazymcp_catalog", + fake_get_lazymcp_catalog, + ) + + await LiteLLM_Proxy_MCP_Handler._get_lazymcp_gateway_tools( + user_api_key_auth=None, + effective_filter=["github"], + active_toolset_id=None, + mcp_auth_header=None, + mcp_server_auth_headers=None, + client_ip="10.0.0.8", + ) + + assert captured["client_ip"] == "10.0.0.8" + + +@pytest.mark.asyncio +async def test_lazymcp_catalog_uses_fail_closed_client_ip(monkeypatch): + from litellm.responses.utils import ResponsesAPIRequestUtils + + captured = {} + + async def fake_get_lazymcp_catalog(**kwargs): + captured.update(kwargs) + return {"description": "ok"} + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._get_lazymcp_catalog", + fake_get_lazymcp_catalog, + ) + + await LiteLLM_Proxy_MCP_Handler._get_lazymcp_gateway_tools( + user_api_key_auth=None, + effective_filter=None, + active_toolset_id=None, + mcp_auth_header=None, + mcp_server_auth_headers=None, + client_ip=ResponsesAPIRequestUtils.get_verified_mcp_client_ip(None), + ) + + assert captured["client_ip"] == "__invalid_mcp_client_ip__" + + +@pytest.mark.asyncio +async def test_lazymcp_catalog_rejects_unauthorized_toolset(monkeypatch): + get_catalog_mock = AsyncMock(return_value={"description": "blocked"}) + apply_scope_mock = AsyncMock( + side_effect=HTTPException(status_code=403, detail="forbidden") + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._get_lazymcp_catalog", + get_catalog_mock, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._apply_toolset_scope", + apply_scope_mock, + ) + + with pytest.raises(HTTPException) as exc_info: + await LiteLLM_Proxy_MCP_Handler._get_lazymcp_gateway_tools( + user_api_key_auth=types.SimpleNamespace(api_key="sk-test"), + effective_filter=None, + active_toolset_id="toolset-blocked", + mcp_auth_header=None, + mcp_server_auth_headers=None, + client_ip="10.0.0.8", + ) + + assert exc_info.value.status_code == 403 + apply_scope_mock.assert_awaited_once() + get_catalog_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_lazymcp_catalog_allowed_toolset_uses_scoped_auth(monkeypatch): + user_auth = types.SimpleNamespace(api_key="sk-test") + scoped_auth = types.SimpleNamespace(api_key="sk-scoped") + get_catalog_mock = AsyncMock(return_value={"description": "ok"}) + apply_scope_mock = AsyncMock(return_value=scoped_auth) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._get_lazymcp_catalog", + get_catalog_mock, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._apply_toolset_scope", + apply_scope_mock, + ) + + await LiteLLM_Proxy_MCP_Handler._get_lazymcp_gateway_tools( + user_api_key_auth=user_auth, + effective_filter=None, + active_toolset_id="toolset-allowed", + mcp_auth_header=None, + mcp_server_auth_headers=None, + client_ip="10.0.0.8", + ) + + apply_scope_mock.assert_awaited_once_with(user_auth, "toolset-allowed") + assert get_catalog_mock.await_args is not None + assert get_catalog_mock.await_args.kwargs["user_api_key_auth"] is scoped_auth + + @pytest.mark.asyncio async def test_execute_tool_calls_strips_server_prefix(monkeypatch): call_tool_mock = _setup_mcp_call_environment(monkeypatch) @@ -279,6 +482,326 @@ async def test_execute_tool_calls_keeps_tool_name_when_equal_to_server(monkeypat assert call_tool_mock.await_args.kwargs["name"] == tool_name +@pytest.mark.asyncio +async def test_execute_tool_calls_does_not_hijack_standard_mcp_name_collision( + monkeypatch, +): + call_tool_mock = _setup_mcp_call_environment(monkeypatch) + tool_name = "mcp_call" + tool_calls = [ + { + "id": "call-standard-mcp", + "function": {"name": tool_name, "arguments": "{}"}, + } + ] + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={tool_name: "standard-server"}, + tool_calls=tool_calls, + user_api_key_auth=None, + ) + + assert call_tool_mock.await_count == 1 + assert call_tool_mock.await_args is not None + assert call_tool_mock.await_args.kwargs["server_name"] == "standard-server" + assert call_tool_mock.await_args.kwargs["name"] == tool_name + + +@pytest.mark.asyncio +async def test_execute_tool_calls_passes_lazymcp_route_scope(monkeypatch): + proxy_module = types.SimpleNamespace(proxy_logging_obj=object()) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + + captured = {} + + def fake_set_auth_context(**kwargs): + captured.update(kwargs) + + async def fake_lazymcp_tool_call(_name, _arguments): + return _DummyMCPResult() + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + fake_set_auth_context, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.lazymcp_tool_call", + fake_lazymcp_tool_call, + ) + tool_server_map_value = ( + LiteLLM_Proxy_MCP_Handler._encode_lazymcp_tool_server_map_value( + ["github"], None + ) + ) + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={"mcp_call": tool_server_map_value}, + tool_calls=[ + { + "id": "call-lazy", + "function": { + "name": "mcp_call", + "arguments": '{"server":"github","tool":"search","arguments":{}}', + }, + } + ], + user_api_key_auth=None, + ) + + assert captured["mcp_servers"] == ["github"] + + +@pytest.mark.asyncio +async def test_execute_tool_calls_passes_lazymcp_client_ip_and_scoped_permissions( + monkeypatch, +): + proxy_module = types.SimpleNamespace(proxy_logging_obj=object()) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + + captured = {} + + def fake_set_auth_context(**kwargs): + captured.update(kwargs) + + async def fake_apply_toolset_scope(user_api_key_auth, toolset_id): + captured["toolset_scope"] = { + "user_api_key_auth": user_api_key_auth, + "toolset_id": toolset_id, + } + return user_api_key_auth + + async def fake_lazymcp_tool_call(_name, _arguments): + return _DummyMCPResult() + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + fake_set_auth_context, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._apply_toolset_scope", + fake_apply_toolset_scope, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.lazymcp_tool_call", + fake_lazymcp_tool_call, + ) + + user_auth = types.SimpleNamespace(api_key="sk-test") + tool_server_map_value = ( + LiteLLM_Proxy_MCP_Handler._encode_lazymcp_tool_server_map_value( + ["github"], "toolset-123" + ) + ) + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={"mcp_call": tool_server_map_value}, + tool_calls=[ + { + "id": "call-lazy", + "function": { + "name": "mcp_call", + "arguments": '{"server":"github","tool":"search","arguments":{}}', + }, + } + ], + user_api_key_auth=user_auth, + client_ip="10.0.0.8", + ) + + assert captured["client_ip"] == "10.0.0.8" + assert captured["toolset_scope"] == { + "user_api_key_auth": user_auth, + "toolset_id": "toolset-123", + } + + +@pytest.mark.asyncio +async def test_execute_tool_calls_rejects_unauthorized_lazymcp_toolset(monkeypatch): + proxy_module = types.SimpleNamespace(proxy_logging_obj=object()) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + + async def fake_lazymcp_tool_call(_name, _arguments): + return _DummyMCPResult() + + apply_scope_mock = AsyncMock( + side_effect=HTTPException(status_code=403, detail="forbidden") + ) + lazymcp_tool_call_mock = AsyncMock(side_effect=fake_lazymcp_tool_call) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._apply_toolset_scope", + apply_scope_mock, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.lazymcp_tool_call", + lazymcp_tool_call_mock, + ) + tool_server_map_value = ( + LiteLLM_Proxy_MCP_Handler._encode_lazymcp_tool_server_map_value( + None, "toolset-blocked" + ) + ) + + results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={"mcp_call": tool_server_map_value}, + tool_calls=[ + { + "id": "call-lazy-blocked", + "function": {"name": "mcp_call", "arguments": "{}"}, + } + ], + user_api_key_auth=types.SimpleNamespace(api_key="sk-test"), + ) + + assert results[0]["tool_call_id"] == "call-lazy-blocked" + assert "forbidden" in results[0]["result"] + apply_scope_mock.assert_awaited_once() + lazymcp_tool_call_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_execute_tool_calls_allowed_lazymcp_toolset_uses_scoped_auth( + monkeypatch, +): + proxy_module = types.SimpleNamespace(proxy_logging_obj=object()) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + + user_auth = types.SimpleNamespace(api_key="sk-test") + scoped_auth = types.SimpleNamespace(api_key="sk-scoped") + captured = {} + + def fake_set_auth_context(**kwargs): + captured.update(kwargs) + + async def fake_lazymcp_tool_call(_name, _arguments): + return _DummyMCPResult() + + apply_scope_mock = AsyncMock(return_value=scoped_auth) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server._apply_toolset_scope", + apply_scope_mock, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + fake_set_auth_context, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.lazymcp_tool_call", + fake_lazymcp_tool_call, + ) + tool_server_map_value = ( + LiteLLM_Proxy_MCP_Handler._encode_lazymcp_tool_server_map_value( + None, "toolset-allowed" + ) + ) + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={"mcp_call": tool_server_map_value}, + tool_calls=[ + { + "id": "call-lazy-allowed", + "function": {"name": "mcp_call", "arguments": "{}"}, + } + ], + user_api_key_auth=user_auth, + ) + + apply_scope_mock.assert_awaited_once_with(user_auth, "toolset-allowed") + assert captured["user_api_key_auth"] is scoped_auth + + +@pytest.mark.asyncio +async def test_execute_tool_calls_ignores_spoofed_lazymcp_forwarded_header( + monkeypatch, +): + proxy_module = types.SimpleNamespace(proxy_logging_obj=object()) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + + captured = {} + + def fake_set_auth_context(**kwargs): + captured.update(kwargs) + + async def fake_lazymcp_tool_call(_name, _arguments): + return _DummyMCPResult() + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + fake_set_auth_context, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.lazymcp_tool_call", + fake_lazymcp_tool_call, + ) + tool_server_map_value = ( + LiteLLM_Proxy_MCP_Handler._encode_lazymcp_tool_server_map_value( + ["internal"], None + ) + ) + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={"mcp_call": tool_server_map_value}, + tool_calls=[ + { + "id": "call-lazy", + "function": { + "name": "mcp_call", + "arguments": '{"server":"internal","tool":"search","arguments":{}}', + }, + } + ], + user_api_key_auth=None, + raw_headers={"x-forwarded-for": "10.0.0.1"}, + client_ip="203.0.113.9", + ) + + assert captured["client_ip"] == "203.0.113.9" + + +@pytest.mark.asyncio +async def test_execute_tool_calls_passes_lazymcp_toolset_scope(monkeypatch): + proxy_module = types.SimpleNamespace(proxy_logging_obj=object()) + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) + + captured = {} + + def fake_set_auth_context(**kwargs): + captured.update(kwargs) + + async def fake_lazymcp_tool_call(_name, _arguments): + from litellm.proxy._experimental.mcp_server.server import _mcp_active_toolset_id + + captured["active_toolset"] = _mcp_active_toolset_id.get() + return _DummyMCPResult() + + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + fake_set_auth_context, + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.server.lazymcp_tool_call", + fake_lazymcp_tool_call, + ) + tool_server_map_value = ( + LiteLLM_Proxy_MCP_Handler._encode_lazymcp_tool_server_map_value( + None, "toolset-123" + ) + ) + + await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map={"mcp_status": tool_server_map_value}, + tool_calls=[ + { + "id": "call-lazy-status", + "function": {"name": "mcp_status", "arguments": "{}"}, + } + ], + user_api_key_auth=None, + ) + + assert captured["mcp_servers"] is None + assert captured["active_toolset"] == "toolset-123" + + @pytest.mark.asyncio async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkeypatch): """ @@ -401,3 +924,192 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch assert mock_get_tools.await_args is not None assert mock_get_tools.await_args.kwargs["log_list_tools_to_spendlogs"] is True assert mock_get_tools.await_args.kwargs["list_tools_log_source"] == "responses" + + +@pytest.mark.asyncio +async def test_standard_mcp_preserves_missing_client_ip_behavior(monkeypatch): + captured = {} + + async def fake_standard_tools(**kwargs): + captured.update(kwargs) + return [], [] + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_get_standard_mcp_tools", + fake_standard_tools, + ) + fake_manager = types.SimpleNamespace(get_mcp_server_by_name=MagicMock()) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + + await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + user_api_key_auth=None, + mcp_tools_with_litellm_proxy=[ + {"type": "mcp", "server_url": "litellm_proxy/mcp/standard"} + ], + client_ip="__invalid_mcp_client_ip__", + ) + + assert captured["client_ip"] is None + + +@pytest.mark.asyncio +async def test_standard_mcp_keeps_verified_client_ip(monkeypatch): + captured = {} + + async def fake_standard_tools(**kwargs): + captured.update(kwargs) + return [], [] + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_get_standard_mcp_tools", + fake_standard_tools, + ) + fake_manager = types.SimpleNamespace(get_mcp_server_by_name=MagicMock()) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + + await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + user_api_key_auth=None, + mcp_tools_with_litellm_proxy=[ + {"type": "mcp", "server_url": "litellm_proxy/mcp/standard"} + ], + client_ip="10.0.0.7", + ) + + assert captured["client_ip"] == "10.0.0.7" + + +@pytest.mark.asyncio +async def test_responses_non_streaming_auto_execution_passes_verified_client_ip( + monkeypatch, +): + from litellm.responses import main as responses_main + + tools = [{"type": "mcp", "server_url": "litellm_proxy/lazymcp/internal"}] + captured_execute_kwargs = {} + process_calls = [] + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_parse_mcp_tools", + staticmethod(lambda _tools: (tools, [])), + ) + + async def fake_process(**kwargs): + process_calls.append(kwargs) + return ([], {"mcp_call": 'lazymcp:{"mcp_servers":["internal"]}'}) + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_process_mcp_tools_without_openai_transform", + fake_process, + ) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_transform_mcp_tools_to_openai", + staticmethod(lambda *_args, **_kwargs: []), + ) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_should_auto_execute_tools", + staticmethod(lambda **_kwargs: True), + ) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_extract_tool_calls_from_response", + staticmethod( + lambda **_kwargs: [ + { + "id": "call-1", + "function": { + "name": "mcp_call", + "arguments": '{"server":"internal","tool":"search","arguments":{}}', + }, + } + ] + ), + ) + + async def fake_execute(**kwargs): + captured_execute_kwargs.update(kwargs) + return [{"tool_call_id": "call-1", "result": "executed"}] + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_execute_tool_calls", + fake_execute, + ) + monkeypatch.setattr( + responses_main.ResponsesAPIRequestUtils, + "extract_mcp_headers_from_request", + staticmethod(lambda **_kwargs: (None, None, None, None)), + ) + monkeypatch.setattr( + responses_main, + "aresponses", + AsyncMock( + side_effect=[ + ResponsesAPIResponse( + id="resp-1", + model="test-model", + created_at=123, + output=[], + usage=ResponseAPIUsage( + input_tokens=1, output_tokens=1, total_tokens=2 + ), + ), + ResponsesAPIResponse( + id="resp-2", + model="test-model", + created_at=124, + output=[], + usage=ResponseAPIUsage( + input_tokens=1, output_tokens=1, total_tokens=2 + ), + ), + ] + ), + ) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_make_follow_up_call", + AsyncMock( + return_value=ResponsesAPIResponse( + id="resp-2", + model="test-model", + created_at=124, + output=[], + usage=ResponseAPIUsage(input_tokens=1, output_tokens=1, total_tokens=2), + ) + ), + ) + + await responses_main.aresponses_api_with_mcp( + model="test-model", + input="hello", + tools=tools, + secret_fields={"mcp_client_ip": "10.0.0.7"}, + ) + + assert captured_execute_kwargs["client_ip"] == "10.0.0.7" + assert [call["client_ip"] for call in process_calls] == [ + "10.0.0.7", + "10.0.0.7", + ] + + +def test_chat_streaming_iterator_execution_threads_client_ip(): + from litellm.responses.mcp import chat_completions_handler + + source = inspect.getsource(chat_completions_handler.acompletion_with_mcp) + + assert "client_ip=client_ip" in source + assert "self.client_ip = client_ip" in source + assert "client_ip=self.client_ip" in source diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx index 1c82859e062..152930360c2 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx @@ -112,10 +112,14 @@ const FeatureCard: React.FC = ({ interface MCPConnectProps { currentServerAccessGroups?: string[]; + mode?: "mcp" | "lazymcp"; } -const MCPConnect: React.FC = ({ currentServerAccessGroups = [] }) => { +const MCPConnect: React.FC = ({ currentServerAccessGroups = [], mode = "mcp" }) => { const proxyBaseUrl = getProxyBaseUrl(); + const endpointPath = mode === "lazymcp" ? "/lazymcp" : "/mcp"; + const endpointName = mode === "lazymcp" ? "LazyMCP" : "MCP"; + const serverLabel = mode === "lazymcp" ? "litellm-lazymcp" : "litellm"; const [copiedStates, setCopiedStates] = useState>({}); const [serverHeaders, setServerHeaders] = useState>({ openai: [], @@ -234,9 +238,9 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] } title="MCP Server Information" - description="Connection details for your LiteLLM MCP server" + description={`Connection details for your LiteLLM ${endpointName} server`} > - + = ({ currentServerAccessGroups = [] "tools": [ { "type": "mcp", - "server_label": "litellm", - "server_url": "litellm_proxy", + "server_label": "${serverLabel}", + "server_url": "litellm_proxy${mode === "lazymcp" ? "/lazymcp" : ""}", "require_approval": "never", "headers": { "x-litellm-api-key": "Bearer YOUR_LITELLM_VIRTUAL_KEY", @@ -317,9 +321,9 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] } title="MCP Server Information" - description="Connection details for your LiteLLM MCP server" + description={`Connection details for your LiteLLM ${endpointName} server`} > - + = ({ currentServerAccessGroups = [] "tools": [ { "type": "mcp", - "server_label": "litellm", - "server_url": "${proxyBaseUrl}/mcp", + "server_label": "${serverLabel}", + "server_url": "${proxyBaseUrl}${endpointPath}", "require_approval": "never", "headers": { "x-litellm-api-key": "Bearer YOUR_LITELLM_API_KEY", @@ -368,7 +372,7 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] - Use tools directly from Cursor IDE with LiteLLM MCP. Enable your AI assistant to perform real-world tasks + Use tools directly from Cursor IDE with LiteLLM {endpointName}. Enable your AI assistant to perform real-world tasks without leaving your coding environment. @@ -405,8 +409,8 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] = ({ currentServerAccessGroups = [] - Connect to LiteLLM MCP using HTTP transport. Compatible with any MCP client that supports HTTP streaming. + Connect to LiteLLM {endpointName} using HTTP transport. Compatible with any MCP client that supports HTTP streaming. @@ -450,7 +454,7 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] appropriate transport method. - + = ({ currentServerAccessGroups = []
- Connect to your MCP client - - Use tools directly from any MCP client with LiteLLM MCP. Enable your AI assistant to perform real-world - tasks through a simple, secure connection. + + {mode === "lazymcp" ? "Connect to your MCP client with LazyMCP" : "Connect to your MCP client"} + + + {mode === "lazymcp" + ? "LazyMCP connects MCP clients to LiteLLM while exposing only three gateway tools: mcp_describe, mcp_call, and mcp_status. This avoids sending every upstream MCP tool schema to the model on each request. LazyMCP respects virtual key, team, access group, toolset, and request-header permissions. The x-mcp-servers header works the same as standard MCP and server descriptions help the model choose the right server." + : "Use tools directly from any MCP client with LiteLLM MCP. Enable your AI assistant to perform real-world tasks through a simple, secure connection."}
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx index fc1d941c177..0c90cbee78c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx @@ -129,6 +129,37 @@ describe("MCPServers", () => { expect(networking.fetchMCPServers).toHaveBeenCalledWith("123", undefined); }); + it("should show LazyMCP Connect tab with lazymcp examples and keep Connect tab MCP URLs", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([]); + + const queryClient = createQueryClient(); + render( + + + , + ); + + await waitFor(() => { + expect(screen.getByRole("tab", { name: "Connect" })).toBeInTheDocument(); + }); + + await act(async () => { + fireEvent.click(screen.getByRole("tab", { name: "Connect" })); + }); + + expect(screen.getAllByText("http://localhost:4000/mcp").length).toBeGreaterThan(0); + + await act(async () => { + fireEvent.click(screen.getByRole("tab", { name: "LazyMCP Connect" })); + }); + + expect(screen.getByText("Connect to your MCP client with LazyMCP")).toBeInTheDocument(); + expect(screen.getByText(/mcp_describe, mcp_call, and mcp_status/)).toBeInTheDocument(); + expect(screen.getAllByText("http://localhost:4000/lazymcp").length).toBeGreaterThan(0); + expect(screen.getByText(/litellm_proxy\/lazymcp/)).toBeInTheDocument(); + }); + it("should fetch and merge health status for servers", async () => { // Mock MCP servers data without health status const mockServers = [ diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx index 72d5e4b5aa8..c14a41cc7c8 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx @@ -352,6 +352,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) All Servers Toolsets Connect + LazyMCP Connect Semantic Filter Network Settings {isAdminRole(userRole) && Submitted MCPs } @@ -435,6 +436,9 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) + + +