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:
Cursor Agent 2026-01-23 19:20:02 +00:00
parent 19141e180e
commit f9412d6811
4 changed files with 145 additions and 43 deletions

View file

@ -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

View file

@ -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.

View file

@ -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

View file

@ -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!")