diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0ee74960293..b2cf1faf49b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -136,6 +136,9 @@ from litellm.proxy._types import ( from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, +) from litellm.proxy.utils import ProxyLogging, get_server_root_path from litellm.repositories.table_repositories import MCPServerRepository from litellm.types.llms.custom_http import httpxSpecialProvider @@ -621,6 +624,30 @@ def _openapi_forwarded_extra_headers( return forwarded or None +MCP_PRE_CALL_DATA_KEY = "pre_call_data" + + +async def release_parallel_request_slot( + proxy_logging_obj: ProxyLogging | None, + user_api_key_auth: UserAPIKeyAuth | None, + pre_call_data: dict[str, Any] | None, +) -> None: + """ + Release the ``max_parallel_requests`` slot the MCP pre-call hooks acquired. + + The pre-call hooks run against a synthetic request dict that never reaches + the litellm success/failure logging callbacks, which is where LLM requests + release their slot, so the MCP tool call path owns the release itself. A + no-op when no slot was acquired. + """ + if proxy_logging_obj is None or user_api_key_auth is None or pre_call_data is None: + return + limiter = proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + return + await limiter.async_release_max_parallel_requests_slot(user_api_key_auth, pre_call_data) + + async def _resolve_byok_mcp_auth_header( mcp_server: MCPServer, user_api_key_auth: Optional[UserAPIKeyAuth], @@ -4518,8 +4545,14 @@ class MCPServerManager: ) as e: # Re-raise guardrail exceptions to properly fail the MCP call verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {str(e)}") + await release_parallel_request_slot( + proxy_logging_obj=proxy_logging_obj, + user_api_key_auth=user_api_key_auth, + pre_call_data=synthetic_llm_data, + ) raise e + hook_result[MCP_PRE_CALL_DATA_KEY] = synthetic_llm_data return hook_result def _create_during_hook_task( @@ -5122,6 +5155,7 @@ class MCPServerManager: # Using standard pre_call_hook ######################################################### hook_result: dict[str, Any] = {} + pre_call_data: dict[str, Any] | None = None if proxy_logging_obj: hook_result = await self.pre_call_tool_check( name=name, @@ -5132,86 +5166,96 @@ class MCPServerManager: server=mcp_server, raw_headers=raw_headers, ) + pre_call_data = hook_result.pop(MCP_PRE_CALL_DATA_KEY, None) if "arguments" in hook_result: arguments = hook_result["arguments"] - # Prepare tasks for during hooks - tasks = [] - if proxy_logging_obj: - during_hook_task = self._create_during_hook_task( - name=name, - arguments=arguments, - server_name_from_prefix=server_name, - user_api_key_auth=user_api_key_auth, - proxy_logging_obj=proxy_logging_obj, - start_time=start_time, + try: + # Prepare tasks for during hooks + tasks = [] + if proxy_logging_obj: + during_hook_task = self._create_during_hook_task( + name=name, + arguments=arguments, + server_name_from_prefix=server_name, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + start_time=start_time, + ) + tasks.append(during_hook_task) + + caller_oauth2_headers = oauth2_headers + oauth2_headers = await self._resolve_oauth2_headers_for_tool_call( + mcp_server, oauth2_headers, user_api_key_auth ) - tasks.append(during_hook_task) - caller_oauth2_headers = oauth2_headers - oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, oauth2_headers, user_api_key_auth) + # For OpenAPI servers, call the tool handler directly instead of via MCP client + if mcp_server.spec_path: + verbose_logger.debug("Calling OpenAPI tool %s directly via HTTP handler", name) + if hook_result.get("extra_headers"): + verbose_logger.warning( + "pre_mcp_call hook returned extra_headers for OpenAPI-backed " + "MCP server '%s' — header injection is not supported for " + "OpenAPI servers; headers will be ignored. Use SSE/HTTP " + "transport to enable hook header injection.", + server_name, + ) - # For OpenAPI servers, call the tool handler directly instead of via MCP client - if mcp_server.spec_path: - verbose_logger.debug("Calling OpenAPI tool %s directly via HTTP handler", name) - if hook_result.get("extra_headers"): - verbose_logger.warning( - "pre_mcp_call hook returned extra_headers for OpenAPI-backed " - "MCP server '%s' — header injection is not supported for " - "OpenAPI servers; headers will be ignored. Use SSE/HTTP " - "transport to enable hook header injection.", - server_name, + auth_header_value = ( + _format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None + ) + resolved_auth_headers, forwarded_headers = await self.resolve_openapi_upstream_auth( + mcp_server=mcp_server, + oauth2_headers=caller_oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=mcp_auth_header, + user_api_key_auth=user_api_key_auth, + forwarded_headers=_openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth), ) - auth_header_value = ( - _format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None - ) - resolved_auth_headers, forwarded_headers = await self.resolve_openapi_upstream_auth( - mcp_server=mcp_server, - oauth2_headers=caller_oauth2_headers, - raw_headers=raw_headers, - mcp_auth_header=mcp_auth_header, - user_api_key_auth=user_api_key_auth, - forwarded_headers=_openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth), - ) + async def _call_openapi_via_handler(): + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_auth_header, + _request_extra_headers, + _request_resolved_auth_headers, + ) - async def _call_openapi_via_handler(): - from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, + auth_token = _request_auth_header.set(auth_header_value) + extra_token = _request_extra_headers.set(forwarded_headers) + resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers) + try: + async with self._limit_outbound_concurrency(mcp_server): + return await self._call_openapi_tool_handler(mcp_server, name, arguments) + finally: + _request_auth_header.reset(auth_token) + _request_extra_headers.reset(extra_token) + _request_resolved_auth_headers.reset(resolved_token) + + tasks.append(asyncio.create_task(_call_openapi_via_handler())) + else: + return await self._call_regular_mcp_tool( + mcp_server=mcp_server, + original_tool_name=name, + arguments=arguments, + tasks=tasks, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + proxy_logging_obj=proxy_logging_obj, + host_progress_callback=host_progress_callback, + hook_extra_headers=hook_result.get("extra_headers"), + user_api_key_auth=user_api_key_auth, ) - auth_token = _request_auth_header.set(auth_header_value) - extra_token = _request_extra_headers.set(forwarded_headers) - resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers) - try: - async with self._limit_outbound_concurrency(mcp_server): - return await self._call_openapi_tool_handler(mcp_server, name, arguments) - finally: - _request_auth_header.reset(auth_token) - _request_extra_headers.reset(extra_token) - _request_resolved_auth_headers.reset(resolved_token) - - tasks.append(asyncio.create_task(_call_openapi_via_handler())) - else: - return await self._call_regular_mcp_tool( - mcp_server=mcp_server, - original_tool_name=name, - arguments=arguments, - tasks=tasks, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, + return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj) + finally: + await release_parallel_request_slot( proxy_logging_obj=proxy_logging_obj, - host_progress_callback=host_progress_callback, - hook_extra_headers=hook_result.get("extra_headers"), user_api_key_auth=user_api_key_auth, + pre_call_data=pre_call_data, ) - return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj) - ######################################################### # End of Methods that call the upstream MCP servers ######################################################### diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4fca4406a6f..bd480f2487b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -347,12 +347,14 @@ if MCP_AVAILABLE: outcome_wire_value, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCP_PRE_CALL_DATA_KEY, MCPServerManager, _caller_authorization_fans_out, _client_forwarded_authorization_headers, _should_strip_caller_authorization, _without_authorization, global_mcp_server_manager, + release_parallel_request_slot, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_auth_header, @@ -2734,73 +2736,81 @@ if MCP_AVAILABLE: ) # `pre_call_tool_check` may return guardrail-modified # arguments; honor them on the local path too. - if isinstance(hook_result, dict) and "arguments" in hook_result: + pre_call_data = hook_result.pop(MCP_PRE_CALL_DATA_KEY, None) + if "arguments" in hook_result: arguments = hook_result["arguments"] - verbose_logger.debug(f"Executing local registry tool: {name}") - # For BYOK servers the credential must be injected via a ContextVar - # because the tool function has headers baked into its closure. - # Pre-format the full Authorization header value using the server's - # configured auth_type so the generator doesn't need to know the prefix. - auth_header_value: Optional[str] = None - if mcp_auth_header: - server_auth_type = getattr(mcp_server, "auth_type", None) if mcp_server else None - if server_auth_type == MCPAuth.api_key: - auth_header_value = f"ApiKey {mcp_auth_header}" - elif server_auth_type == MCPAuth.basic: - auth_header_value = f"Basic {mcp_auth_header}" - else: - auth_header_value = f"Bearer {mcp_auth_header}" - - # Forward named client headers to OpenAPI tool upstream requests. - # MCPServer.extra_headers lists header names to copy from raw_headers. - # The strip decision is centralized in _should_strip_caller_authorization so this - # OpenAPI/local path agrees with the managed paths: M2M and the resolver-owned modes - # (token_exchange's raw subject token, authorization_code's stored token) must never - # have the caller's Authorization forwarded verbatim upstream. - forwarded_headers: Optional[Dict[str, str]] = None - if mcp_server and mcp_server.extra_headers and raw_headers: - normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} - skip_caller_authorization = _should_strip_caller_authorization( - mcp_server=mcp_server, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - ) - for header_name in mcp_server.extra_headers: - if not isinstance(header_name, str): - continue - if skip_caller_authorization and header_name.lower() == "authorization": - continue - value = normalized_raw.get(header_name.lower()) - if value is not None: - if forwarded_headers is None: - forwarded_headers = {} - forwarded_headers[header_name] = value - - resolved_auth_headers: dict[str, str] | None = None - if mcp_server: - ( - resolved_auth_headers, - forwarded_headers, - ) = await global_mcp_server_manager.resolve_openapi_upstream_auth( - mcp_server=mcp_server, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - mcp_auth_header=mcp_auth_header, - user_api_key_auth=user_api_key_auth, - forwarded_headers=forwarded_headers, - ) - - _auth_token = _request_auth_header.set(auth_header_value) - _extra_token = _request_extra_headers.set(forwarded_headers) - _resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers) try: - local_content = await _handle_local_mcp_tool(name, arguments) + verbose_logger.debug(f"Executing local registry tool: {name}") + # For BYOK servers the credential must be injected via a ContextVar + # because the tool function has headers baked into its closure. + # Pre-format the full Authorization header value using the server's + # configured auth_type so the generator doesn't need to know the prefix. + auth_header_value: str | None = None + if mcp_auth_header: + server_auth_type = getattr(mcp_server, "auth_type", None) if mcp_server else None + if server_auth_type == MCPAuth.api_key: + auth_header_value = f"ApiKey {mcp_auth_header}" + elif server_auth_type == MCPAuth.basic: + auth_header_value = f"Basic {mcp_auth_header}" + else: + auth_header_value = f"Bearer {mcp_auth_header}" + + # Forward named client headers to OpenAPI tool upstream requests. + # MCPServer.extra_headers lists header names to copy from raw_headers. + # The strip decision is centralized in _should_strip_caller_authorization so this + # OpenAPI/local path agrees with the managed paths: M2M and the resolver-owned modes + # (token_exchange's raw subject token, authorization_code's stored token) must never + # have the caller's Authorization forwarded verbatim upstream. + forwarded_headers: dict[str, str] | None = None + if mcp_server and mcp_server.extra_headers and raw_headers: + normalized_raw = {str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)} + skip_caller_authorization = _should_strip_caller_authorization( + mcp_server=mcp_server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + for header_name in mcp_server.extra_headers: + if not isinstance(header_name, str): + continue + if skip_caller_authorization and header_name.lower() == "authorization": + continue + value = normalized_raw.get(header_name.lower()) + if value is not None: + if forwarded_headers is None: + forwarded_headers = {} + forwarded_headers[header_name] = value + + resolved_auth_headers: dict[str, str] | None = None + if mcp_server: + ( + resolved_auth_headers, + forwarded_headers, + ) = await global_mcp_server_manager.resolve_openapi_upstream_auth( + mcp_server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=mcp_auth_header, + user_api_key_auth=user_api_key_auth, + forwarded_headers=forwarded_headers, + ) + + _auth_token = _request_auth_header.set(auth_header_value) + _extra_token = _request_extra_headers.set(forwarded_headers) + _resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers) + try: + local_content = await _handle_local_mcp_tool(name, arguments) + finally: + _request_auth_header.reset(_auth_token) + _request_extra_headers.reset(_extra_token) + _request_resolved_auth_headers.reset(_resolved_token) + response = CallToolResult(content=cast(Any, local_content), isError=False) finally: - _request_auth_header.reset(_auth_token) - _request_extra_headers.reset(_extra_token) - _request_resolved_auth_headers.reset(_resolved_token) - response = CallToolResult(content=cast(Any, local_content), isError=False) + await release_parallel_request_slot( + proxy_logging_obj=proxy_logging_obj, + user_api_key_auth=user_api_key_auth, + pre_call_data=pre_call_data, + ) # Try managed MCP server tool (pass the full prefixed name) # Primary and recommended way to use external MCP servers diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b2216488db2..c06201562f6 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -3378,7 +3378,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): except Exception as e: verbose_proxy_logger.exception(f"Error in rate limit failure event: {str(e)}") - async def async_release_max_parallel_requests_on_disconnect( + async def async_release_max_parallel_requests_slot( self, user_api_key_dict: UserAPIKeyAuth, request_data: dict | None = None, @@ -3393,7 +3393,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): client cancels a stream mid-flight, the cancellation surfaces as ``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback runs, so without this the slot leaks per cancelled stream until its - TTL prunes it. ``request_data`` carries the stashed acquisition; + TTL prunes it. MCP tool calls are the other caller: they run the + pre-call hooks against a synthetic request dict that no logging + callback ever sees, so the tool call path releases its own slot. + ``request_data`` carries the stashed acquisition; its presence (not the key object's current max_parallel_requests configuration, which can change mid-request) decides whether there is anything to release. diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5b81d1f2da3..33f0db48b7e 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2609,7 +2609,7 @@ class ProxyLogging: limiter = self.get_proxy_hook("parallel_request_limiter") if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): return - await limiter.async_release_max_parallel_requests_on_disconnect(user_api_key_dict, request_data) + await limiter.async_release_max_parallel_requests_slot(user_api_key_dict, request_data) def _init_response_taking_too_long_task(self, data: Optional[dict] = None): """ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index b56a12db5b1..a766762c9a0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -17,7 +17,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCP_PRE_CALL_DATA_KEY, + MCPServerManager, +) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging from litellm.types.mcp import MCPAuth, MCPTransport @@ -118,6 +121,7 @@ class TestPreCallToolCheckReturnsHeaders: server=server, ) + assert result.pop(MCP_PRE_CALL_DATA_KEY) == {"model": "fake"} assert result == {} @pytest.mark.asyncio @@ -179,6 +183,7 @@ class TestPreCallToolCheckReturnsHeaders: server=server, ) + assert result.pop(MCP_PRE_CALL_DATA_KEY) == {"model": "fake"} assert result == {} @pytest.mark.asyncio 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 dff1f1d87c7..065a3d0713a 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 @@ -7726,3 +7726,67 @@ class TestPreemptive401ModeAware: await self._run(delegate, None, has_stored_token=False) assert exc.value.status_code == 401 await self._run(delegate, self.LITELLM_KEY_HEADERS, has_stored_token=False) + + +@pytest.mark.asyncio +async def test_local_registry_tool_calls_release_max_parallel_requests_slot(monkeypatch): + """ + The local (OpenAPI-backed) dispatch path runs the same pre-call hooks as the + managed path, so it acquires a ``max_parallel_requests`` slot too. Nothing + downstream releases it, so ``execute_mcp_tool`` must; otherwise a key with + ``max_parallel_requests: N`` wedges after N sequential tool calls. + """ + import litellm + from litellm.caching.dual_cache import DualCache + from litellm.proxy._experimental.mcp_server import server as mcp_module + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) + from litellm.proxy.utils import InternalUsageCache, ProxyLogging + + cache = DualCache() + limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + proxy_logging_obj = ProxyLogging(user_api_key_cache=cache) + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter + monkeypatch.setattr(litellm, "callbacks", [limiter]) + + openapi_server = MCPServer( + server_id="openapi-server-id", + name="openapi_server", + server_name="openapi_server", + url="https://example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + spec_path="/fake/openapi.json", + ) + user_api_key_auth = UserAPIKeyAuth(api_key="hashed-key", max_parallel_requests=1) + + with ( + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj), + patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=MagicMock()), + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=openapi_server, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "resolve_openapi_upstream_auth", + new=AsyncMock(return_value=(None, None)), + ), + patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True), + patch.object( + mcp_module, + "_handle_local_mcp_tool", + new=AsyncMock(return_value=[TextContent(type="text", text="ok")]), + ), + ): + for _ in range(3): + result = await mcp_module.execute_mcp_tool( + name="openapi_server-tool", + arguments={}, + allowed_mcp_servers=[openapi_server], + start_time=datetime.now(), + user_api_key_auth=user_api_key_auth, + ) + assert result.isError is False diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 42b6cbee1c4..46891d960d5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -9240,3 +9240,164 @@ class TestDiscoveryFailureLogging: assert "typo_row" in caplog.text assert "authorization_url, token_url" in caplog.text assert "unresolved" in caplog.text + + +class TestMaxParallelRequestsSlotRelease: + """ + Every MCP tool call runs the proxy pre-call hooks, which acquire a + ``max_parallel_requests`` slot. Nothing in the MCP path fires the litellm + success/failure logging callbacks that release the slot for LLM requests, + so the tool call itself must release it; otherwise a key with + ``max_parallel_requests: N`` wedges after N strictly sequential tool calls + and only a restart (or the hour-long slot TTL) frees it. + """ + + @staticmethod + def _proxy_logging_with_limiter(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) + from litellm.proxy.utils import InternalUsageCache, ProxyLogging + + cache = DualCache() + limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache)) + proxy_logging_obj = ProxyLogging(user_api_key_cache=cache) + proxy_logging_obj.proxy_hook_mapping["parallel_request_limiter"] = limiter + return proxy_logging_obj, limiter + + @staticmethod + def _server() -> MCPServer: + return MCPServer( + server_id="slot-server", + name="slot_server", + server_name="slot_server", + url="https://example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + + @staticmethod + def _user_api_key_auth(max_parallel_requests: int): + from litellm.proxy._types import UserAPIKeyAuth + + return UserAPIKeyAuth(api_key="hashed-key", max_parallel_requests=max_parallel_requests) + + @staticmethod + def _patched_manager(manager: MCPServerManager, server: MCPServer, tool_call): + async def fake_create_mcp_client(server, **kwargs): + class _Client: + async def call_tool(self, params, host_progress_callback=None): + return await tool_call() + + return _Client() + + return ( + patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server), + patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client), + ) + + async def _call(self, manager, proxy_logging_obj, user_api_key_auth): + return await manager.call_tool( + server_name="slot_server", + name="tool", + arguments={}, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + ) + + @pytest.mark.asyncio + async def test_sequential_tool_calls_never_exhaust_the_limit(self, monkeypatch): + proxy_logging_obj, limiter = self._proxy_logging_with_limiter() + monkeypatch.setattr("litellm.callbacks", [limiter]) + manager = MCPServerManager() + server = self._server() + user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1) + + async def tool_call(): + return "ok" + + resolve_patch, client_patch = self._patched_manager(manager, server, tool_call) + with resolve_patch, client_patch: + for _ in range(5): + await self._call(manager, proxy_logging_obj, user_api_key_auth) + + @pytest.mark.asyncio + async def test_failed_tool_call_releases_its_slot(self, monkeypatch): + proxy_logging_obj, limiter = self._proxy_logging_with_limiter() + monkeypatch.setattr("litellm.callbacks", [limiter]) + manager = MCPServerManager() + server = self._server() + user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1) + + async def failing_tool_call(): + raise ValueError("upstream exploded") + + resolve_patch, client_patch = self._patched_manager(manager, server, failing_tool_call) + with resolve_patch, client_patch: + with pytest.raises(ValueError): + await self._call(manager, proxy_logging_obj, user_api_key_auth) + + async def tool_call(): + return "ok" + + resolve_patch, client_patch = self._patched_manager(manager, server, tool_call) + with resolve_patch, client_patch: + await self._call(manager, proxy_logging_obj, user_api_key_auth) + + @pytest.mark.asyncio + async def test_guardrail_blocked_tool_call_releases_its_slot(self, monkeypatch): + """A guardrail that rejects the call after the limiter admitted it must + not strand the acquired slot.""" + from litellm.integrations.custom_logger import CustomLogger + + class _BlockingGuardrail(CustomLogger): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + raise HTTPException(status_code=400, detail="blocked by guardrail") + + proxy_logging_obj, limiter = self._proxy_logging_with_limiter() + monkeypatch.setattr("litellm.callbacks", [limiter, _BlockingGuardrail()]) + manager = MCPServerManager() + server = self._server() + user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1) + + async def tool_call(): + return "ok" + + resolve_patch, client_patch = self._patched_manager(manager, server, tool_call) + with resolve_patch, client_patch: + with pytest.raises(HTTPException): + await self._call(manager, proxy_logging_obj, user_api_key_auth) + + monkeypatch.setattr("litellm.callbacks", [limiter]) + resolve_patch, client_patch = self._patched_manager(manager, server, tool_call) + with resolve_patch, client_patch: + await self._call(manager, proxy_logging_obj, user_api_key_auth) + + @pytest.mark.asyncio + async def test_concurrent_tool_calls_still_hit_the_limit(self, monkeypatch): + """Releasing after the call must not turn the limiter off: two tool + calls in flight against a limit of 1 still reject the second.""" + from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError + + proxy_logging_obj, limiter = self._proxy_logging_with_limiter() + monkeypatch.setattr("litellm.callbacks", [limiter]) + manager = MCPServerManager() + server = self._server() + user_api_key_auth = self._user_api_key_auth(max_parallel_requests=1) + + async def slow_tool_call(): + await asyncio.sleep(0.2) + return "ok" + + resolve_patch, client_patch = self._patched_manager(manager, server, slow_tool_call) + with resolve_patch, client_patch: + results = await asyncio.gather( + self._call(manager, proxy_logging_obj, user_api_key_auth), + self._call(manager, proxy_logging_obj, user_api_key_auth), + return_exceptions=True, + ) + + rejections = [r for r in results if isinstance(r, ProxyRateLimitError)] + assert len(rejections) == 1 + assert "max_parallel_requests" in str(rejections[0].detail) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index c76e1a60afd..c4b2abb76c6 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -3511,7 +3511,7 @@ async def test_release_max_parallel_requests_on_disconnect_v3(): await local_cache.async_get_cache(key=counter_key) ) == 1 - await handler.async_release_max_parallel_requests_on_disconnect( + await handler.async_release_max_parallel_requests_slot( user_api_key_dict, request_data={ "metadata": { @@ -3544,7 +3544,7 @@ async def test_release_on_disconnect_works_when_key_config_changed_v3(): counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" await _seed_max_parallel_requests_slots(local_cache, counter_key, [_TEST_SLOT_ID]) - await handler.async_release_max_parallel_requests_on_disconnect( + await handler.async_release_max_parallel_requests_slot( UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None), request_data={ "metadata": { @@ -3886,12 +3886,12 @@ async def test_release_max_parallel_requests_on_disconnect_noop_v3(): ) counter_key = f"{{api_key:{_api_key}}}:max_parallel_requests" - await handler.async_release_max_parallel_requests_on_disconnect( + await handler.async_release_max_parallel_requests_slot( UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=None) ) assert await local_cache.async_get_cache(key=counter_key) is None - await handler.async_release_max_parallel_requests_on_disconnect( + await handler.async_release_max_parallel_requests_slot( UserAPIKeyAuth(api_key=None, max_parallel_requests=5) ) assert await local_cache.async_get_cache(key=counter_key) is None