mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): add defensive None checks to prevent NoneType errors in back-to-back MCP calls
Fixed the 'argument of type NoneType is not iterable' error that occurred when making back-to-back MCP Responses API calls. The issue was caused by missing defensive checks for None values in several places: 1. _get_allowed_mcp_servers_from_mcp_server_names: Added None check for allowed_mcp_servers parameter before iteration 2. _get_mcp_tools_from_manager: Added defensive checks to ensure allowed_mcp_server_ids and allowed_mcp_servers are always lists 3. _deduplicate_mcp_tools: Added None checks for both mcp_tools and allowed_mcp_servers parameters 4. _filter_mcp_tools_by_allowed_tools: Added None checks for both mcp_tools and mcp_tools_with_litellm_proxy parameters 5. _extract_mcp_headers_from_params in streaming iterator: Added try-catch and hasattr check to safely iterate over tools 6. _create_initial_response_iterator: Improved error handling and added validation for tools parameter Added test test_mcp_handler_none_defensive_checks to verify the fix. Co-authored-by: ishaan <ishaan@berri.ai>
This commit is contained in:
parent
19141e180e
commit
f9412d6811
4 changed files with 145 additions and 43 deletions
|
|
@ -559,11 +559,14 @@ if MCP_AVAILABLE:
|
|||
|
||||
async def _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers: Optional[List[str]],
|
||||
allowed_mcp_servers: List[MCPServer],
|
||||
allowed_mcp_servers: Optional[List[MCPServer]],
|
||||
) -> List[MCPServer]:
|
||||
"""
|
||||
Get the filtered MCP servers from the MCP server names
|
||||
"""
|
||||
# Defensive check: ensure allowed_mcp_servers is not None
|
||||
if allowed_mcp_servers is None:
|
||||
allowed_mcp_servers = []
|
||||
|
||||
filtered_server: dict[str, MCPServer] = {}
|
||||
# Filter servers based on mcp_servers parameter if provided
|
||||
|
|
|
|||
|
|
@ -140,14 +140,24 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
allowed_mcp_server_ids = (
|
||||
await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth)
|
||||
)
|
||||
# Defensive check: ensure allowed_mcp_server_ids is a list
|
||||
if allowed_mcp_server_ids is None:
|
||||
allowed_mcp_server_ids = []
|
||||
|
||||
allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( # type: ignore[attr-defined]
|
||||
allowed_mcp_server_ids
|
||||
)
|
||||
# Defensive check: ensure allowed_mcp_servers is a list
|
||||
if allowed_mcp_servers is None:
|
||||
allowed_mcp_servers = []
|
||||
|
||||
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
|
||||
mcp_servers=mcp_servers,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
)
|
||||
# Defensive check: ensure the result is a list
|
||||
if allowed_mcp_servers is None:
|
||||
allowed_mcp_servers = []
|
||||
|
||||
server_names: List[str] = []
|
||||
for server in allowed_mcp_servers:
|
||||
|
|
@ -165,7 +175,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _deduplicate_mcp_tools(
|
||||
mcp_tools: List[MCPTool], allowed_mcp_servers: List[str]
|
||||
mcp_tools: Optional[List[MCPTool]], allowed_mcp_servers: Optional[List[str]]
|
||||
) -> tuple[List[MCPTool], dict[str, str]]:
|
||||
"""
|
||||
Deduplicate MCP tools by name, keeping the first occurrence of each tool.
|
||||
|
|
@ -177,6 +187,12 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
List of deduplicated MCP tools
|
||||
The returned dictionary maps each tool_name to the server_name
|
||||
"""
|
||||
# Defensive checks for None inputs
|
||||
if mcp_tools is None:
|
||||
mcp_tools = []
|
||||
if allowed_mcp_servers is None:
|
||||
allowed_mcp_servers = []
|
||||
|
||||
seen_names = set()
|
||||
deduplicated_tools = []
|
||||
tool_server_map: dict[str, str] = {}
|
||||
|
|
@ -201,9 +217,15 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
def _filter_mcp_tools_by_allowed_tools(
|
||||
mcp_tools: List[MCPTool], mcp_tools_with_litellm_proxy: List[ToolParam]
|
||||
mcp_tools: Optional[List[MCPTool]], mcp_tools_with_litellm_proxy: Optional[List[ToolParam]]
|
||||
) -> List[MCPTool]:
|
||||
"""Filter MCP tools based on allowed_tools parameter from the original tool configs."""
|
||||
# Defensive checks for None inputs
|
||||
if mcp_tools is None:
|
||||
return []
|
||||
if mcp_tools_with_litellm_proxy is None:
|
||||
return list(mcp_tools)
|
||||
|
||||
# Collect all allowed tool names from all MCP tool configs
|
||||
allowed_tool_names = set()
|
||||
for tool_config in mcp_tools_with_litellm_proxy:
|
||||
|
|
@ -239,7 +261,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
async def _process_mcp_tools_to_openai_format(
|
||||
user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam]
|
||||
user_api_key_auth: Any, mcp_tools_with_litellm_proxy: Optional[List[ToolParam]]
|
||||
) -> tuple[List[Any], dict[str, str]]:
|
||||
"""
|
||||
Centralized method to process MCP tools through the complete pipeline.
|
||||
|
|
@ -268,7 +290,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
|
||||
@staticmethod
|
||||
async def _process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam]
|
||||
user_api_key_auth: Any, mcp_tools_with_litellm_proxy: Optional[List[ToolParam]]
|
||||
) -> tuple[List[Any], dict[str, str]]:
|
||||
"""
|
||||
Process MCP tools through filtering and deduplication pipeline without OpenAI transformation.
|
||||
|
|
|
|||
|
|
@ -340,37 +340,41 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
# Also check if headers are provided in tools array (from request body)
|
||||
tools = self.original_request_params.get("tools")
|
||||
if tools:
|
||||
for tool in tools:
|
||||
if isinstance(tool, dict) and tool.get("type") == "mcp":
|
||||
tool_headers = tool.get("headers", {})
|
||||
if tool_headers and isinstance(tool_headers, dict):
|
||||
# Merge tool headers into mcp_server_auth_headers
|
||||
headers_obj_from_tool = Headers(tool_headers)
|
||||
tool_mcp_server_auth_headers = (
|
||||
MCPRequestHandler._get_mcp_server_auth_headers_from_headers(
|
||||
headers_obj_from_tool
|
||||
)
|
||||
)
|
||||
|
||||
if tool_mcp_server_auth_headers:
|
||||
if self.mcp_server_auth_headers is None:
|
||||
self.mcp_server_auth_headers = {}
|
||||
# Merge the headers from tool into existing headers
|
||||
for (
|
||||
server_alias,
|
||||
headers_dict,
|
||||
) in tool_mcp_server_auth_headers.items():
|
||||
if server_alias not in self.mcp_server_auth_headers:
|
||||
self.mcp_server_auth_headers[server_alias] = {}
|
||||
self.mcp_server_auth_headers[server_alias].update(
|
||||
headers_dict
|
||||
# Defensive check: ensure tools is iterable
|
||||
if tools and hasattr(tools, '__iter__'):
|
||||
try:
|
||||
for tool in tools:
|
||||
if isinstance(tool, dict) and tool.get("type") == "mcp":
|
||||
tool_headers = tool.get("headers", {})
|
||||
if tool_headers and isinstance(tool_headers, dict):
|
||||
# Merge tool headers into mcp_server_auth_headers
|
||||
headers_obj_from_tool = Headers(tool_headers)
|
||||
tool_mcp_server_auth_headers = (
|
||||
MCPRequestHandler._get_mcp_server_auth_headers_from_headers(
|
||||
headers_obj_from_tool
|
||||
)
|
||||
)
|
||||
|
||||
# Also merge raw headers
|
||||
if self.raw_headers is None:
|
||||
self.raw_headers = {}
|
||||
self.raw_headers.update(tool_headers)
|
||||
if tool_mcp_server_auth_headers:
|
||||
if self.mcp_server_auth_headers is None:
|
||||
self.mcp_server_auth_headers = {}
|
||||
# Merge the headers from tool into existing headers
|
||||
for (
|
||||
server_alias,
|
||||
headers_dict,
|
||||
) in tool_mcp_server_auth_headers.items():
|
||||
if server_alias not in self.mcp_server_auth_headers:
|
||||
self.mcp_server_auth_headers[server_alias] = {}
|
||||
self.mcp_server_auth_headers[server_alias].update(
|
||||
headers_dict
|
||||
)
|
||||
|
||||
# Also merge raw headers
|
||||
if self.raw_headers is None:
|
||||
self.raw_headers = {}
|
||||
self.raw_headers.update(tool_headers)
|
||||
except (TypeError, AttributeError) as e:
|
||||
verbose_logger.debug(f"Error iterating over tools in _extract_mcp_headers_from_params: {e}")
|
||||
|
||||
def _should_auto_execute_tools(self) -> bool:
|
||||
"""Check if tools should be auto-executed"""
|
||||
|
|
@ -498,22 +502,24 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
from litellm.responses.main import aresponses
|
||||
|
||||
# Make the initial response API call - but avoid the MCP wrapper
|
||||
params = self.original_request_params.copy()
|
||||
params = self.original_request_params.copy() if self.original_request_params else {}
|
||||
params["stream"] = True # Ensure streaming
|
||||
|
||||
# Use the pre-fetched all_tools from original_request_params (no re-processing needed)
|
||||
params_for_llm = {}
|
||||
for key, value in params.items():
|
||||
params_for_llm[
|
||||
key
|
||||
] = value # Copy all params as-is since tools are already processed
|
||||
# Skip None values and ensure tools is a valid list
|
||||
if value is None:
|
||||
continue
|
||||
if key == "tools" and not isinstance(value, (list, tuple)):
|
||||
verbose_logger.warning(f"Skipping invalid tools value: {type(value)}")
|
||||
continue
|
||||
params_for_llm[key] = value
|
||||
|
||||
tools_count = (
|
||||
len(params_for_llm.get("tools", []))
|
||||
if params_for_llm.get("tools")
|
||||
else 0
|
||||
)
|
||||
tools = params_for_llm.get("tools")
|
||||
tools_count = len(tools) if tools and isinstance(tools, (list, tuple)) else 0
|
||||
verbose_logger.debug(f"Making LLM call with {tools_count} tools")
|
||||
|
||||
response = await aresponses(**params_for_llm)
|
||||
|
||||
# Set the base iterator
|
||||
|
|
|
|||
|
|
@ -1177,4 +1177,75 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e():
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_handler_none_defensive_checks():
|
||||
"""
|
||||
Test that MCP handler methods properly handle None inputs without raising
|
||||
'argument of type NoneType is not iterable' errors.
|
||||
|
||||
This test verifies the fix for the bug where back-to-back MCP Responses API calls
|
||||
could fail when certain parameters were None.
|
||||
"""
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import LiteLLM_Proxy_MCP_Handler
|
||||
|
||||
print("🧪 Testing MCP handler None defensive checks...")
|
||||
|
||||
# Test 1: _deduplicate_mcp_tools with None inputs
|
||||
result_tools, result_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
|
||||
mcp_tools=None,
|
||||
allowed_mcp_servers=None
|
||||
)
|
||||
assert result_tools == [], "Should return empty list for None mcp_tools"
|
||||
assert result_map == {}, "Should return empty dict for None allowed_mcp_servers"
|
||||
print("✅ _deduplicate_mcp_tools handles None inputs correctly")
|
||||
|
||||
# Test 2: _deduplicate_mcp_tools with None mcp_tools but valid servers
|
||||
result_tools, result_map = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(
|
||||
mcp_tools=None,
|
||||
allowed_mcp_servers=["server1", "server2"]
|
||||
)
|
||||
assert result_tools == [], "Should return empty list for None mcp_tools"
|
||||
print("✅ _deduplicate_mcp_tools handles None mcp_tools correctly")
|
||||
|
||||
# Test 3: _filter_mcp_tools_by_allowed_tools with None inputs
|
||||
result = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools(
|
||||
mcp_tools=None,
|
||||
mcp_tools_with_litellm_proxy=None
|
||||
)
|
||||
assert result == [], "Should return empty list for None inputs"
|
||||
print("✅ _filter_mcp_tools_by_allowed_tools handles None inputs correctly")
|
||||
|
||||
# Test 4: _filter_mcp_tools_by_allowed_tools with None mcp_tools_with_litellm_proxy
|
||||
mock_tools = [
|
||||
type('MCPTool', (), {
|
||||
'name': 'test_tool',
|
||||
'description': 'A test tool',
|
||||
'inputSchema': {}
|
||||
})()
|
||||
]
|
||||
result = LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools(
|
||||
mcp_tools=mock_tools,
|
||||
mcp_tools_with_litellm_proxy=None
|
||||
)
|
||||
assert len(result) == 1, "Should return all tools when mcp_tools_with_litellm_proxy is None"
|
||||
print("✅ _filter_mcp_tools_by_allowed_tools handles None mcp_tools_with_litellm_proxy correctly")
|
||||
|
||||
# Test 5: _process_mcp_tools_without_openai_transform with None input
|
||||
result_tools, result_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=None,
|
||||
mcp_tools_with_litellm_proxy=None
|
||||
)
|
||||
assert result_tools == [], "Should return empty list for None mcp_tools_with_litellm_proxy"
|
||||
assert result_map == {}, "Should return empty dict for None input"
|
||||
print("✅ _process_mcp_tools_without_openai_transform handles None inputs correctly")
|
||||
|
||||
# Test 6: _process_mcp_tools_without_openai_transform with empty list
|
||||
result_tools, result_map = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
|
||||
user_api_key_auth=None,
|
||||
mcp_tools_with_litellm_proxy=[]
|
||||
)
|
||||
assert result_tools == [], "Should return empty list for empty mcp_tools_with_litellm_proxy"
|
||||
assert result_map == {}, "Should return empty dict for empty input"
|
||||
print("✅ _process_mcp_tools_without_openai_transform handles empty list correctly")
|
||||
|
||||
print("🎉 All MCP handler None defensive checks passed!")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue