mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(mcp): bound outbound tool-call concurrency per MCP server (#31641)
Add an optional per-server max_concurrent_requests that caps how many tool calls LiteLLM sends to one MCP server at once, so batch-processing backends are not overwhelmed by unbounded parallel dispatch. Excess calls queue on a per-server asyncio.Semaphore instead of being rejected. Unset or non-positive means unlimited, preserving existing behavior. Resolves LIT-2749
This commit is contained in:
parent
8e6098adc3
commit
58de920921
10 changed files with 234 additions and 3 deletions
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "max_concurrent_requests" INTEGER;
|
||||
|
|
@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
max_concurrent_requests Int?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -93,6 +93,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
has_user_credential: Optional[bool] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
approval_status: Optional[str] = Field(
|
||||
default="active",
|
||||
description="Approval status: 'pending_review', 'active', 'rejected'",
|
||||
|
|
|
|||
|
|
@ -13,7 +13,8 @@ import json
|
|||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import anyio
|
||||
|
|
@ -550,6 +551,11 @@ class MCPServerManager:
|
|||
]
|
||||
"""
|
||||
|
||||
# Per-server outbound tool-call concurrency limiters, lazily created from
|
||||
# each server's max_concurrent_requests. Keyed by server_id so the cap
|
||||
# survives the registry atomic-swap on config reload; a missing key means
|
||||
# the server has no configured limit.
|
||||
self._server_call_semaphores: Dict[str, asyncio.Semaphore] = {}
|
||||
self.tool_name_to_mcp_server_name_mapping: Dict[str, str] = {}
|
||||
"""
|
||||
{
|
||||
|
|
@ -790,6 +796,7 @@ class MCPServerManager:
|
|||
allow_sampling=bool(server_config.get("allow_sampling", False)),
|
||||
allow_elicitation=bool(server_config.get("allow_elicitation", False)),
|
||||
timeout=server_config.get("timeout", None),
|
||||
max_concurrent_requests=server_config.get("max_concurrent_requests", None),
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
|
||||
|
|
@ -1164,6 +1171,7 @@ class MCPServerManager:
|
|||
subject_token_type=(credentials_dict.get("subject_token_type") if credentials_dict else None)
|
||||
or "urn:ietf:params:oauth:token-type:access_token",
|
||||
timeout=getattr(mcp_server, "timeout", None),
|
||||
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
|
||||
)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="database")
|
||||
return new_server
|
||||
|
|
@ -3211,6 +3219,25 @@ class MCPServerManager:
|
|||
)
|
||||
)
|
||||
|
||||
def _get_call_semaphore(self, mcp_server: MCPServer) -> Optional[asyncio.Semaphore]:
|
||||
limit = mcp_server.max_concurrent_requests
|
||||
if limit is None or limit <= 0:
|
||||
return None
|
||||
semaphore = self._server_call_semaphores.get(mcp_server.server_id)
|
||||
if semaphore is None:
|
||||
semaphore = asyncio.Semaphore(limit)
|
||||
self._server_call_semaphores[mcp_server.server_id] = semaphore
|
||||
return semaphore
|
||||
|
||||
@asynccontextmanager
|
||||
async def _limit_outbound_concurrency(self, mcp_server: MCPServer) -> AsyncIterator[None]:
|
||||
semaphore = self._get_call_semaphore(mcp_server)
|
||||
if semaphore is None:
|
||||
yield
|
||||
return
|
||||
async with semaphore:
|
||||
yield
|
||||
|
||||
async def _call_regular_mcp_tool(
|
||||
self,
|
||||
mcp_server: MCPServer,
|
||||
|
|
@ -3370,7 +3397,8 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
async def _call_tool_via_client(client, params):
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await client.call_tool(params, host_progress_callback=host_progress_callback)
|
||||
|
||||
tasks.append(asyncio.create_task(_call_tool_via_client(client, call_tool_params)))
|
||||
|
||||
|
|
@ -3592,7 +3620,12 @@ class MCPServerManager:
|
|||
"transport to enable hook header injection.",
|
||||
server_name,
|
||||
)
|
||||
tasks.append(asyncio.create_task(self._call_openapi_tool_handler(mcp_server, name, arguments)))
|
||||
|
||||
async def _call_openapi_via_handler():
|
||||
async with self._limit_outbound_concurrency(mcp_server):
|
||||
return await self._call_openapi_tool_handler(mcp_server, name, arguments)
|
||||
|
||||
tasks.append(asyncio.create_task(_call_openapi_via_handler()))
|
||||
else:
|
||||
return await self._call_regular_mcp_tool(
|
||||
mcp_server=mcp_server,
|
||||
|
|
@ -4149,6 +4182,7 @@ class MCPServerManager:
|
|||
allow_all_keys=server.allow_all_keys,
|
||||
instructions=server.instructions,
|
||||
timeout=server.timeout,
|
||||
max_concurrent_requests=server.max_concurrent_requests,
|
||||
)
|
||||
|
||||
async def get_all_mcp_servers_with_health_and_teams(
|
||||
|
|
@ -4256,6 +4290,7 @@ class MCPServerManager:
|
|||
source_url=server.source_url,
|
||||
instructions=server.instructions,
|
||||
timeout=server.timeout,
|
||||
max_concurrent_requests=server.max_concurrent_requests,
|
||||
)
|
||||
|
||||
async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]:
|
||||
|
|
|
|||
|
|
@ -1261,6 +1261,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
byok_api_key_help_url: Optional[str] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
# BYOM submission fields — set by the endpoint, not by the caller.
|
||||
# Any caller-provided values are silently overridden before persistence.
|
||||
approval_status: Optional[str] = Field(
|
||||
|
|
@ -1346,6 +1347,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
byok_api_key_help_url: Optional[str] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -654,6 +654,7 @@ if MCP_AVAILABLE:
|
|||
allow_all_keys=payload.allow_all_keys,
|
||||
available_on_public_internet=payload.available_on_public_internet,
|
||||
timeout=payload.timeout,
|
||||
max_concurrent_requests=payload.max_concurrent_requests,
|
||||
)
|
||||
|
||||
def get_prisma_client_or_throw(message: str):
|
||||
|
|
|
|||
|
|
@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
max_concurrent_requests Int?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -118,6 +118,9 @@ class MCPServer(BaseModel):
|
|||
# MCP_PER_USER_TOKEN_DEFAULT_TTL when expires_in is absent.
|
||||
token_storage_ttl_seconds: Optional[int] = None
|
||||
timeout: Optional[float] = None
|
||||
# Max concurrent outbound tool calls to this server; excess calls queue.
|
||||
# None or a value <= 0 means unlimited.
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
# Resolved short-ID tool prefix when LITELLM_USE_SHORT_MCP_TOOL_PREFIX is
|
||||
# enabled. Set by ``MCPServerManager._assign_unique_short_prefix`` at
|
||||
# registration time so that natural-hash collisions between two
|
||||
|
|
|
|||
|
|
@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
max_concurrent_requests Int?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -0,0 +1,184 @@
|
|||
import asyncio
|
||||
from typing import Dict, Optional
|
||||
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
HOLD_SECONDS = 0.1
|
||||
|
||||
|
||||
class _ConcurrencyTracker:
|
||||
"""Records how many call_tool invocations are simultaneously in flight."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.current_by_server: Dict[str, int] = {}
|
||||
self.peak_by_server: Dict[str, int] = {}
|
||||
self.global_current = 0
|
||||
self.global_peak = 0
|
||||
|
||||
def enter(self, server_id: str) -> None:
|
||||
self.current_by_server[server_id] = self.current_by_server.get(server_id, 0) + 1
|
||||
self.peak_by_server[server_id] = max(self.peak_by_server.get(server_id, 0), self.current_by_server[server_id])
|
||||
self.global_current += 1
|
||||
self.global_peak = max(self.global_peak, self.global_current)
|
||||
|
||||
def exit(self, server_id: str) -> None:
|
||||
self.current_by_server[server_id] -= 1
|
||||
self.global_current -= 1
|
||||
|
||||
|
||||
def _make_server(server_id: str, max_concurrent_requests: Optional[int]) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=server_id,
|
||||
name=server_id,
|
||||
server_name=server_id,
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
max_concurrent_requests=max_concurrent_requests,
|
||||
)
|
||||
|
||||
|
||||
def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyTracker):
|
||||
async def fake_create_mcp_client(server, **kwargs):
|
||||
class _ProbeClient:
|
||||
async def call_tool(self, params, host_progress_callback=None):
|
||||
tracker.enter(server.server_id)
|
||||
try:
|
||||
await asyncio.sleep(HOLD_SECONDS)
|
||||
return "ok"
|
||||
finally:
|
||||
tracker.exit(server.server_id)
|
||||
|
||||
return _ProbeClient()
|
||||
|
||||
return patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client)
|
||||
|
||||
|
||||
async def _fire(manager: MCPServerManager, server: MCPServer, n: int) -> None:
|
||||
await asyncio.gather(
|
||||
*[
|
||||
manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="tool",
|
||||
arguments={},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
for _ in range(n)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_concurrent_requests_caps_in_flight_tool_calls():
|
||||
"""A configured cap of 2 must never let more than 2 calls hit one server at once."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-limited", max_concurrent_requests=2)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await _fire(manager, server, n=8)
|
||||
|
||||
assert tracker.peak_by_server["srv-limited"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unset_limit_allows_unbounded_concurrency():
|
||||
"""With no cap, all calls run concurrently (backward-compatible default)."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-unbounded", max_concurrent_requests=None)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await _fire(manager, server, n=6)
|
||||
|
||||
assert tracker.peak_by_server["srv-unbounded"] == 6
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_positive_limit_is_treated_as_unlimited():
|
||||
"""A cap of 0 must not deadlock; it means unlimited, not a zero-permit semaphore."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-zero", max_concurrent_requests=0)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await asyncio.wait_for(_fire(manager, server, n=5), timeout=5)
|
||||
|
||||
assert tracker.peak_by_server["srv-zero"] == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_limit_is_scoped_per_server():
|
||||
"""Each server gets its own limiter; one server's cap must not throttle another."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server_a = _make_server("srv-a", max_concurrent_requests=1)
|
||||
server_b = _make_server("srv-b", max_concurrent_requests=1)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await asyncio.gather(
|
||||
_fire(manager, server_a, n=3),
|
||||
_fire(manager, server_b, n=3),
|
||||
)
|
||||
|
||||
assert tracker.peak_by_server["srv-a"] == 1
|
||||
assert tracker.peak_by_server["srv-b"] == 1
|
||||
assert tracker.global_peak == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openapi_backed_server_also_respects_the_cap():
|
||||
"""OpenAPI (spec_path) servers dispatch through a different handler; the cap
|
||||
must apply there too, not only on the regular MCP client path."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-openapi", max_concurrent_requests=2)
|
||||
server.spec_path = "/fake/openapi.json"
|
||||
|
||||
async def fake_openapi_handler(mcp_server, name, arguments):
|
||||
tracker.enter(mcp_server.server_id)
|
||||
try:
|
||||
await asyncio.sleep(HOLD_SECONDS)
|
||||
return "ok"
|
||||
finally:
|
||||
tracker.exit(mcp_server.server_id)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server),
|
||||
patch.object(manager, "_call_openapi_tool_handler", side_effect=fake_openapi_handler),
|
||||
):
|
||||
await asyncio.gather(
|
||||
*[manager.call_tool(server_name="srv-openapi", name="tool", arguments={}) for _ in range(6)]
|
||||
)
|
||||
|
||||
assert tracker.peak_by_server["srv-openapi"] == 2
|
||||
|
||||
|
||||
def test_semaphore_is_reused_per_server_and_distinct_across_servers():
|
||||
manager = MCPServerManager()
|
||||
server_a = _make_server("srv-a", max_concurrent_requests=3)
|
||||
server_b = _make_server("srv-b", max_concurrent_requests=3)
|
||||
|
||||
sem_a_first = manager._get_call_semaphore(server_a)
|
||||
sem_a_second = manager._get_call_semaphore(server_a)
|
||||
sem_b = manager._get_call_semaphore(server_b)
|
||||
|
||||
assert sem_a_first is sem_a_second
|
||||
assert sem_a_first is not sem_b
|
||||
|
||||
|
||||
def test_no_semaphore_created_when_limit_absent():
|
||||
manager = MCPServerManager()
|
||||
server = _make_server("srv-none", max_concurrent_requests=None)
|
||||
|
||||
assert manager._get_call_semaphore(server) is None
|
||||
Loading…
Add table
Reference in a new issue