diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index c4754ef6117..330d11e3a9c 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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? diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 0bc81ece5f0..aed00c060ca 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0d2008cdade..887fe610403 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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]: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 6c36c813e76..41eedecbb04 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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", diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 4d67df16b0b..05cfc674497 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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): diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index c4754ef6117..330d11e3a9c 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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? diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 2108fe8990d..4f8c9a0aa48 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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 diff --git a/schema.prisma b/schema.prisma index c4754ef6117..330d11e3a9c 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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? 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 e7d0ee6247b..7deff9d9b38 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 @@ -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