mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
feat(mcp): add per-server timeout configuration (#29672)
* feat(mcp): add per-server timeout configuration * fix(mcp): address timeout field review comments - use is not None guard instead of or for 0.0 edge case - copy timeout in both LiteLLM_MCPServerTable constructions (health check path + _build_mcp_server_table) - add timeout Float? column to all three schema.prisma files - extend round-trip test to cover _build_mcp_server_table direction - add test for zero timeout not treated as falsy * fix(mcp): forward timeout in _build_temporary_mcp_server_record * fix(mcp): return 504 instead of 500 when per-server timeout fires * test(mcp): add 504 timeout regression test; fix black formatting
This commit is contained in:
parent
17734eb621
commit
82a4f1543f
9 changed files with 146 additions and 4 deletions
|
|
@ -330,6 +330,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_description String[] @default([])
|
||||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -556,7 +556,9 @@ class MCPClient:
|
|||
)
|
||||
return tool_result
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning("MCP client tool call was cancelled")
|
||||
verbose_logger.warning(
|
||||
f"MCP client tool call timed out after {self.timeout}s for {self.server_url}"
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
|
|
|||
|
|
@ -684,6 +684,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),
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="config")
|
||||
|
|
@ -1096,6 +1097,7 @@ class MCPServerManager:
|
|||
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),
|
||||
)
|
||||
_warn_internal_delegate_pkce_if_applicable(new_server, source="database")
|
||||
return new_server
|
||||
|
|
@ -1662,7 +1664,9 @@ class MCPServerManager:
|
|||
transport_type=transport,
|
||||
auth_type=server.auth_type,
|
||||
auth_value=auth_value,
|
||||
timeout=MCP_CLIENT_TIMEOUT,
|
||||
timeout=(
|
||||
server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
),
|
||||
stdio_config=stdio_config,
|
||||
extra_headers=extra_headers,
|
||||
sampling_callback=sampling_cb,
|
||||
|
|
@ -1690,7 +1694,9 @@ class MCPServerManager:
|
|||
transport_type=transport,
|
||||
auth_type=server.auth_type,
|
||||
auth_value=auth_value,
|
||||
timeout=MCP_CLIENT_TIMEOUT,
|
||||
timeout=(
|
||||
server.timeout if server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
),
|
||||
extra_headers=extra_headers,
|
||||
aws_auth=aws_auth,
|
||||
sampling_callback=sampling_cb,
|
||||
|
|
@ -3160,12 +3166,24 @@ class MCPServerManager:
|
|||
|
||||
try:
|
||||
mcp_responses = await asyncio.gather(*tasks)
|
||||
except asyncio.CancelledError:
|
||||
timeout = (
|
||||
mcp_server.timeout
|
||||
if mcp_server.timeout is not None
|
||||
else MCP_CLIENT_TIMEOUT
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=504,
|
||||
detail={
|
||||
"error": "timeout",
|
||||
"message": f"MCP tool call timed out after {timeout}s",
|
||||
},
|
||||
)
|
||||
except (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
HTTPException,
|
||||
) as e:
|
||||
# Re-raise guardrail exceptions to properly fail the MCP call
|
||||
verbose_logger.error(
|
||||
f"Guardrail blocked MCP tool call during result check: {str(e)}"
|
||||
)
|
||||
|
|
@ -3953,6 +3971,7 @@ class MCPServerManager:
|
|||
registration_url=server.registration_url,
|
||||
allow_all_keys=server.allow_all_keys,
|
||||
instructions=server.instructions,
|
||||
timeout=server.timeout,
|
||||
)
|
||||
|
||||
async def get_all_mcp_servers_with_health_and_teams(
|
||||
|
|
@ -4052,6 +4071,7 @@ class MCPServerManager:
|
|||
byok_api_key_help_url=server.byok_api_key_help_url,
|
||||
source_url=server.source_url,
|
||||
instructions=server.instructions,
|
||||
timeout=server.timeout,
|
||||
)
|
||||
|
||||
async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]:
|
||||
|
|
|
|||
|
|
@ -1300,6 +1300,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
byok_description: List[str] = Field(default_factory=list)
|
||||
byok_api_key_help_url: Optional[str] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = 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(
|
||||
|
|
@ -1384,6 +1385,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
byok_description: List[str] = Field(default_factory=list)
|
||||
byok_api_key_help_url: Optional[str] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
@ -1458,6 +1460,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
byok_api_key_help_url: Optional[str] = None
|
||||
has_user_credential: Optional[bool] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
# BYOM submission fields
|
||||
approval_status: Optional[str] = Field(
|
||||
default="active",
|
||||
|
|
|
|||
|
|
@ -659,6 +659,7 @@ if MCP_AVAILABLE:
|
|||
registration_url=payload.registration_url,
|
||||
allow_all_keys=payload.allow_all_keys,
|
||||
available_on_public_internet=payload.available_on_public_internet,
|
||||
timeout=payload.timeout,
|
||||
)
|
||||
|
||||
def get_prisma_client_or_throw(message: str):
|
||||
|
|
|
|||
|
|
@ -330,6 +330,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_description String[] @default([])
|
||||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -109,6 +109,7 @@ class MCPServer(BaseModel):
|
|||
# Defaults to the token's expires_in minus the expiry buffer, or
|
||||
# MCP_PER_USER_TOKEN_DEFAULT_TTL when expires_in is absent.
|
||||
token_storage_ttl_seconds: Optional[int] = None
|
||||
timeout: Optional[float] = 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
|
||||
|
|
|
|||
|
|
@ -330,6 +330,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_description String[] @default([])
|
||||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import importlib
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
|
@ -3066,6 +3067,117 @@ class TestMCPServerTimestamps:
|
|||
rebuilt_table = manager._build_mcp_server_table(mcp_server)
|
||||
assert rebuilt_table.source_url == "https://github.com/org/mcp-server"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_round_trip_timeout_preserved(self):
|
||||
"""timeout survives the full round-trip: LiteLLM_MCPServerTable -> MCPServer -> LiteLLM_MCPServerTable."""
|
||||
manager = MCPServerManager()
|
||||
table_record = LiteLLM_MCPServerTable(
|
||||
server_id="timeout-server",
|
||||
server_name="timeout_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
timeout=120.0,
|
||||
)
|
||||
mcp_server = await manager.build_mcp_server_from_table(table_record)
|
||||
assert mcp_server.timeout == 120.0
|
||||
|
||||
rebuilt_table = manager._build_mcp_server_table(mcp_server)
|
||||
assert rebuilt_table.timeout == 120.0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_mcp_client_uses_server_timeout(self):
|
||||
"""_create_mcp_client must pass server.timeout to MCPClient when set."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="timeout-client-server",
|
||||
name="timeout_client_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
timeout=180.0,
|
||||
)
|
||||
client = await manager._create_mcp_client(server)
|
||||
assert client.timeout == 180.0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_mcp_client_falls_back_to_global_timeout(self):
|
||||
"""_create_mcp_client must fall back to MCP_CLIENT_TIMEOUT when server.timeout is None."""
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="default-timeout-server",
|
||||
name="default_timeout_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
client = await manager._create_mcp_client(server)
|
||||
assert client.timeout == MCP_CLIENT_TIMEOUT
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_mcp_client_zero_timeout_not_treated_as_falsy(self):
|
||||
"""server.timeout=0.0 must be passed through, not fall back to MCP_CLIENT_TIMEOUT."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="zero-timeout-server",
|
||||
name="zero_timeout_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
timeout=0.0,
|
||||
)
|
||||
client = await manager._create_mcp_client(server)
|
||||
assert client.timeout == 0.0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_servers_from_config_preserves_timeout(self):
|
||||
"""timeout from proxy config is loaded into MCPServer."""
|
||||
manager = MCPServerManager()
|
||||
config = {
|
||||
"my_server": {
|
||||
"url": "https://example.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
"timeout": 90.0,
|
||||
}
|
||||
}
|
||||
await manager.load_servers_from_config(config)
|
||||
servers = list(manager.config_mcp_servers.values())
|
||||
assert len(servers) == 1
|
||||
assert servers[0].timeout == 90.0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_regular_mcp_tool_timeout_returns_504(self):
|
||||
"""When the MCP client call is cancelled (timeout), _call_regular_mcp_tool raises HTTPException 504."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="timeout-tool-server",
|
||||
name="timeout_tool_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
timeout=2.0,
|
||||
)
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.call_tool = AsyncMock(side_effect=asyncio.CancelledError())
|
||||
|
||||
with patch.object(manager, "_create_mcp_client", return_value=mock_client):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="some_tool",
|
||||
arguments={},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 504
|
||||
assert exc_info.value.detail["error"] == "timeout"
|
||||
assert "2.0s" in exc_info.value.detail["message"]
|
||||
|
||||
|
||||
class TestInternalDelegatePkceWarningLog:
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue