mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
issues resolve
This commit is contained in:
parent
0cb5d9b1e4
commit
9b710c3502
8 changed files with 1060 additions and 386 deletions
File diff suppressed because it is too large
Load diff
|
|
@ -155,7 +155,11 @@ def _convert_single_content(content: Any) -> Dict[str, Any]:
|
|||
# ToolResultContent → represents tool results
|
||||
tool_content = getattr(content, "content", [])
|
||||
if isinstance(tool_content, list) and tool_content:
|
||||
texts = [getattr(c, "text", str(c)) for c in tool_content if getattr(c, "type", None) == "text"]
|
||||
texts = [
|
||||
getattr(c, "text", str(c))
|
||||
for c in tool_content
|
||||
if getattr(c, "type", None) == "text"
|
||||
]
|
||||
return {"type": "text", "text": "\n".join(texts) if texts else ""}
|
||||
return {"type": "text", "text": str(tool_content)}
|
||||
# Fallback: treat as text
|
||||
|
|
@ -242,7 +246,9 @@ def _extract_tool_calls(content: Any) -> List[Dict[str, Any]]:
|
|||
"type": "function",
|
||||
"function": {
|
||||
"name": getattr(item, "name", ""),
|
||||
"arguments": json.dumps(getattr(item, "input", {}), default=str),
|
||||
"arguments": json.dumps(
|
||||
getattr(item, "input", {}), default=str
|
||||
),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
|
@ -269,7 +275,11 @@ def _extract_tool_results(content: Any) -> List[Dict[str, Any]]:
|
|||
# Extract text from nested content
|
||||
nested_content = getattr(item, "content", [])
|
||||
if isinstance(nested_content, list):
|
||||
text_parts = [getattr(c, "text", str(c)) for c in nested_content if getattr(c, "type", None) == "text"]
|
||||
text_parts = [
|
||||
getattr(c, "text", str(c))
|
||||
for c in nested_content
|
||||
if getattr(c, "type", None) == "text"
|
||||
]
|
||||
result_text = "\n".join(text_parts) if text_parts else ""
|
||||
else:
|
||||
result_text = str(nested_content)
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -28,19 +28,25 @@ async def main():
|
|||
|
||||
return ElicitResult(action="accept", content=user_response)
|
||||
|
||||
async with sse_client("http://localhost:4000/mcp/sse", headers={"Authorization": "Bearer sk-1234"}) as (
|
||||
async with sse_client(
|
||||
"http://localhost:4000/mcp/sse", headers={"Authorization": "Bearer sk-1234"}
|
||||
) as (
|
||||
read_stream,
|
||||
write_stream,
|
||||
):
|
||||
logger.info("SSE connection established.")
|
||||
async with ClientSession(read_stream, write_stream, elicitation_callback=my_elicitation_callback) as session:
|
||||
async with ClientSession(
|
||||
read_stream, write_stream, elicitation_callback=my_elicitation_callback
|
||||
) as session:
|
||||
await session.initialize()
|
||||
logger.info("Initialized!")
|
||||
|
||||
logger.info("\n--- Testing Complex Pipeline (Elicitation + Sampling) ---")
|
||||
logger.info("Calling 'test_server-test_complex_pipeline'...")
|
||||
try:
|
||||
result = await session.call_tool("test_server-test_complex_pipeline", arguments={})
|
||||
result = await session.call_tool(
|
||||
"test_server-test_complex_pipeline", arguments={}
|
||||
)
|
||||
logger.info("\nFINAL TOOL RESULT:")
|
||||
logger.info("==================")
|
||||
logger.info(result.content[0].text)
|
||||
|
|
|
|||
|
|
@ -31,7 +31,11 @@ async def test_auth_context_persistence():
|
|||
@pytest.mark.asyncio
|
||||
async def test_get_or_extract_auth_context_fallback():
|
||||
"""Test get_or_extract_auth_context fallback to server object."""
|
||||
from litellm.proxy._experimental.mcp_server.server import server, MCPAuthenticatedUser, auth_context_var
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
server,
|
||||
MCPAuthenticatedUser,
|
||||
auth_context_var,
|
||||
)
|
||||
|
||||
auth_data = UserAPIKeyAuth(api_key="fallback-key")
|
||||
auth_user = MCPAuthenticatedUser(user_api_key_auth=auth_data)
|
||||
|
|
@ -61,7 +65,9 @@ async def test_extract_mcp_auth_context_with_key():
|
|||
|
||||
mock_user_auth = UserAPIKeyAuth(api_key="sk-123")
|
||||
|
||||
with patch("litellm.proxy.auth.auth_checks.common_checks", new_callable=AsyncMock) as mock_auth:
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.common_checks", new_callable=AsyncMock
|
||||
) as mock_auth:
|
||||
mock_auth.return_value = mock_user_auth
|
||||
|
||||
result = await extract_mcp_auth_context(mock_scope, "/mcp/sse")
|
||||
|
|
|
|||
|
|
@ -22,7 +22,10 @@ def test_convert_image_content():
|
|||
mock_image.mimeType = "image/jpeg"
|
||||
|
||||
result = _convert_single_content(mock_image)
|
||||
assert result == {"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,base64data"}}
|
||||
assert result == {
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/jpeg;base64,base64data"},
|
||||
}
|
||||
|
||||
|
||||
def test_convert_audio_content():
|
||||
|
|
@ -32,7 +35,10 @@ def test_convert_audio_content():
|
|||
mock_audio.mimeType = "audio/mp3"
|
||||
|
||||
result = _convert_single_content(mock_audio)
|
||||
assert result == {"type": "input_audio", "input_audio": {"data": "audiobase64", "format": "mp3"}}
|
||||
assert result == {
|
||||
"type": "input_audio",
|
||||
"input_audio": {"data": "audiobase64", "format": "mp3"},
|
||||
}
|
||||
|
||||
|
||||
def test_convert_list_content():
|
||||
|
|
@ -64,7 +70,10 @@ def test_resolve_model_from_hints():
|
|||
original_router = proxy_server.llm_router
|
||||
try:
|
||||
proxy_server.llm_router = MagicMock()
|
||||
proxy_server.llm_router.get_model_names.return_value = ["gpt-4", "claude-3-5-sonnet"]
|
||||
proxy_server.llm_router.get_model_names.return_value = [
|
||||
"gpt-4",
|
||||
"claude-3-5-sonnet",
|
||||
]
|
||||
result = _resolve_model_from_preferences(mock_prefs)
|
||||
assert result == "claude-3-5-sonnet"
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -78,7 +78,9 @@ def _make_params(
|
|||
return params
|
||||
|
||||
|
||||
def _make_completion_response(content="Hello!", model="gpt-4o-mini", finish_reason="stop", tool_calls=None):
|
||||
def _make_completion_response(
|
||||
content="Hello!", model="gpt-4o-mini", finish_reason="stop", tool_calls=None
|
||||
):
|
||||
"""Create a mock litellm completion response."""
|
||||
response = MagicMock()
|
||||
choice = MagicMock()
|
||||
|
|
@ -155,7 +157,9 @@ class TestConvertMCPMessagesToOpenAI:
|
|||
)
|
||||
|
||||
user_msg = _make_sampling_message("user", _make_text_content("Hi"))
|
||||
assistant_msg = _make_sampling_message("assistant", _make_text_content("Hello!"))
|
||||
assistant_msg = _make_sampling_message(
|
||||
"assistant", _make_text_content("Hello!")
|
||||
)
|
||||
result = _convert_mcp_messages_to_openai([user_msg, assistant_msg])
|
||||
assert len(result) == 2
|
||||
assert result[0]["role"] == "user"
|
||||
|
|
@ -173,7 +177,9 @@ class TestResolveModel:
|
|||
_resolve_model_from_preferences,
|
||||
)
|
||||
|
||||
result = _resolve_model_from_preferences(None, default_model="claude-3.5-sonnet")
|
||||
result = _resolve_model_from_preferences(
|
||||
None, default_model="claude-3.5-sonnet"
|
||||
)
|
||||
assert result == "claude-3.5-sonnet"
|
||||
|
||||
@patch("litellm.proxy.proxy_server.llm_router", None)
|
||||
|
|
@ -309,7 +315,9 @@ class TestConvertOpenAIResponseToMCPResult:
|
|||
_convert_openai_response_to_mcp_result,
|
||||
)
|
||||
|
||||
response = _make_completion_response(content="Partial...", finish_reason="length")
|
||||
response = _make_completion_response(
|
||||
content="Partial...", finish_reason="length"
|
||||
)
|
||||
result = _convert_openai_response_to_mcp_result(response, "gpt-4o")
|
||||
assert result.stopReason == "maxTokens"
|
||||
|
||||
|
|
@ -322,7 +330,9 @@ class TestConvertOpenAIResponseToMCPResult:
|
|||
tc.id = "call_123"
|
||||
tc.function.name = "get_weather"
|
||||
tc.function.arguments = '{"city": "NYC"}'
|
||||
response = _make_completion_response(content=None, finish_reason="tool_calls", tool_calls=[tc])
|
||||
response = _make_completion_response(
|
||||
content=None, finish_reason="tool_calls", tool_calls=[tc]
|
||||
)
|
||||
result = _convert_openai_response_to_mcp_result(response, "gpt-4o")
|
||||
assert result.stopReason == "toolUse"
|
||||
assert isinstance(result.content, list)
|
||||
|
|
|
|||
|
|
@ -35,14 +35,18 @@ from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPSer
|
|||
|
||||
def _reload_mcp_manager_module():
|
||||
utils_module = sys.modules["litellm.proxy._experimental.mcp_server.utils"]
|
||||
manager_module = sys.modules["litellm.proxy._experimental.mcp_server.mcp_server_manager"]
|
||||
manager_module = sys.modules[
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager"
|
||||
]
|
||||
importlib.reload(utils_module)
|
||||
reloaded = importlib.reload(manager_module)
|
||||
# After reload, server.py still holds a stale reference to the old
|
||||
# global_mcp_server_manager. Update it so tests that exercise server.py
|
||||
# functions (e.g. _get_tools_from_mcp_servers) use the fresh instance.
|
||||
server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server")
|
||||
if server_module is not None and hasattr(server_module, "global_mcp_server_manager"):
|
||||
if server_module is not None and hasattr(
|
||||
server_module, "global_mcp_server_manager"
|
||||
):
|
||||
server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager
|
||||
return reloaded
|
||||
|
||||
|
|
@ -303,7 +307,9 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock get_allowed_mcp_servers to return our test servers
|
||||
manager.get_allowed_mcp_servers = AsyncMock(return_value=["github", "zapier"])
|
||||
manager.get_mcp_server_by_id = MagicMock(side_effect=lambda x: server1 if x == "github" else server2)
|
||||
manager.get_mcp_server_by_id = MagicMock(
|
||||
side_effect=lambda x: server1 if x == "github" else server2
|
||||
)
|
||||
|
||||
# Mock _get_tools_from_server to return different results
|
||||
async def mock_get_tools_from_server(
|
||||
|
|
@ -331,7 +337,9 @@ class TestMCPServerManager:
|
|||
"zapier": "zapier-api-key",
|
||||
}
|
||||
|
||||
result = await manager.list_tools(mcp_server_auth_headers=mcp_server_auth_headers)
|
||||
result = await manager.list_tools(
|
||||
mcp_server_auth_headers=mcp_server_auth_headers
|
||||
)
|
||||
|
||||
# Verify that both servers were called with their specific auth headers
|
||||
assert len(result) == 3 # 2 from github + 1 from zapier
|
||||
|
|
@ -402,7 +410,9 @@ class TestMCPServerManager:
|
|||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
assert mcp_auth_header == "server-specific-token" # Should use server-specific header
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
) # Should use server-specific header
|
||||
tool = MagicMock()
|
||||
tool.name = "github_tool_1"
|
||||
return [tool]
|
||||
|
|
@ -433,7 +443,9 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=CallToolResult(content=[], isError=False))
|
||||
mock_client.call_tool = AsyncMock(
|
||||
return_value=CallToolResult(content=[], isError=False)
|
||||
)
|
||||
captured_extra_headers = None
|
||||
|
||||
async def capture_create_mcp_client(
|
||||
|
|
@ -547,7 +559,9 @@ class TestMCPServerManager:
|
|||
mock_client = AsyncMock()
|
||||
mock_resources = [Resource(name="file", uri="https://example.com/file")]
|
||||
mock_client.list_resources = AsyncMock(return_value=mock_resources)
|
||||
prefixed_resources = [Resource(name="alias-server-file", uri="https://example.com/file")]
|
||||
prefixed_resources = [
|
||||
Resource(name="alias-server-file", uri="https://example.com/file")
|
||||
]
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
|
|
@ -678,7 +692,9 @@ class TestMCPServerManager:
|
|||
mock_create_client.assert_called_once()
|
||||
called_kwargs = mock_create_client.call_args.kwargs
|
||||
assert called_kwargs["extra_headers"] == {"X-Test": "1", "X-Static": "1"}
|
||||
mock_client.read_resource.assert_awaited_once_with("https://example.com/resource")
|
||||
mock_client.read_resource.assert_awaited_once_with(
|
||||
"https://example.com/resource"
|
||||
)
|
||||
assert result is read_result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -725,7 +741,9 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
def raise_http_error():
|
||||
raise httpx.HTTPStatusError("unauthorized", request=request, response=response_obj)
|
||||
raise httpx.HTTPStatusError(
|
||||
"unauthorized", request=request, response=response_obj
|
||||
)
|
||||
|
||||
response_obj.raise_for_status = MagicMock(side_effect=raise_http_error)
|
||||
|
||||
|
|
@ -843,7 +861,9 @@ class TestMCPServerManager:
|
|||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
assert mcp_auth_header == "server-specific-token" # Should use server-specific header via server_name
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
) # Should use server-specific header via server_name
|
||||
tool = MagicMock()
|
||||
tool.name = "github_tool_1"
|
||||
return [tool]
|
||||
|
|
@ -910,7 +930,9 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock failed client.run_with_session
|
||||
mock_client = AsyncMock()
|
||||
mock_client.run_with_session = AsyncMock(side_effect=Exception("Connection timeout"))
|
||||
mock_client.run_with_session = AsyncMock(
|
||||
side_effect=Exception("Connection timeout")
|
||||
)
|
||||
manager._create_mcp_client = AsyncMock(return_value=mock_client)
|
||||
|
||||
# Perform health check
|
||||
|
|
@ -1032,7 +1054,9 @@ class TestMCPServerManager:
|
|||
# Capture the extra_headers passed to _create_mcp_client
|
||||
captured_extra_headers = None
|
||||
|
||||
async def capture_create_mcp_client(server, mcp_auth_header, extra_headers, stdio_env):
|
||||
async def capture_create_mcp_client(
|
||||
server, mcp_auth_header, extra_headers, stdio_env
|
||||
):
|
||||
nonlocal captured_extra_headers
|
||||
captured_extra_headers = extra_headers
|
||||
return mock_client
|
||||
|
|
@ -1374,7 +1398,9 @@ class TestMCPServerManager:
|
|||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
|
||||
|
||||
|
|
@ -1418,8 +1444,13 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Tool blocked_tool is not allowed for server test-server" in exc_info.value.detail["error"]
|
||||
assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
|
||||
assert (
|
||||
"Tool blocked_tool is not allowed for server test-server"
|
||||
in exc_info.value.detail["error"]
|
||||
)
|
||||
assert (
|
||||
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_tool_check_disallowed_tools_list_allows_tool(self):
|
||||
|
|
@ -1443,7 +1474,9 @@ class TestMCPServerManager:
|
|||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
|
||||
|
||||
|
|
@ -1487,8 +1520,13 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Tool banned_tool is not allowed for server test-server" in exc_info.value.detail["error"]
|
||||
assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
|
||||
assert (
|
||||
"Tool banned_tool is not allowed for server test-server"
|
||||
in exc_info.value.detail["error"]
|
||||
)
|
||||
assert (
|
||||
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_tool_check_no_restrictions_allows_any_tool(self):
|
||||
|
|
@ -1512,7 +1550,9 @@ class TestMCPServerManager:
|
|||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
|
||||
|
||||
|
|
@ -1549,7 +1589,9 @@ class TestMCPServerManager:
|
|||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
|
||||
|
||||
|
|
@ -1575,7 +1617,10 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Tool tool3 is not allowed for server test-server" in exc_info.value.detail["error"]
|
||||
assert (
|
||||
"Tool tool3 is not allowed for server test-server"
|
||||
in exc_info.value.detail["error"]
|
||||
)
|
||||
|
||||
async def test_get_tools_from_server_add_prefix(self):
|
||||
"""Verify _get_tools_from_server respects add_prefix True/False."""
|
||||
|
|
@ -1606,7 +1651,9 @@ class TestMCPServerManager:
|
|||
assert tools_prefixed[0].name == "zapier-send_email"
|
||||
|
||||
# Case 2: add_prefix=False (single-server) -> expect unprefixed
|
||||
tools_unprefixed = await manager._get_tools_from_server(server, add_prefix=False)
|
||||
tools_unprefixed = await manager._get_tools_from_server(
|
||||
server, add_prefix=False
|
||||
)
|
||||
assert len(tools_unprefixed) == 1
|
||||
assert tools_unprefixed[0].name == "send_email"
|
||||
|
||||
|
|
@ -1641,9 +1688,13 @@ class TestMCPServerManager:
|
|||
|
||||
# Mapping should include both original and prefixed names -> resolves calls either way
|
||||
assert manager.tool_name_to_mcp_server_name_mapping["create_issue"] == "jira"
|
||||
assert manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira"
|
||||
assert (
|
||||
manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira"
|
||||
)
|
||||
assert manager.tool_name_to_mcp_server_name_mapping["close_issue"] == "jira"
|
||||
assert manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira"
|
||||
assert (
|
||||
manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira"
|
||||
)
|
||||
|
||||
def test_get_mcp_server_from_tool_name_with_prefixed_and_unprefixed(self):
|
||||
"""After mapping is populated, manager resolves both prefixed and unprefixed tool names to the same server."""
|
||||
|
|
@ -1673,7 +1724,9 @@ class TestMCPServerManager:
|
|||
assert resolved_server_unpref.server_id == server.server_id
|
||||
|
||||
# Prefixed resolution
|
||||
resolved_server_pref = manager._get_mcp_server_from_tool_name("zapier-create_zap")
|
||||
resolved_server_pref = manager._get_mcp_server_from_tool_name(
|
||||
"zapier-create_zap"
|
||||
)
|
||||
assert resolved_server_pref is not None
|
||||
assert resolved_server_pref.server_id == server.server_id
|
||||
|
||||
|
|
@ -1718,7 +1771,9 @@ class TestMCPServerManager:
|
|||
new=AsyncMock(return_value=[tool1, tool2, tool3]),
|
||||
):
|
||||
# Call the REST endpoint helper
|
||||
filtered_response = await _get_tools_for_single_server(server, server_auth_header=None)
|
||||
filtered_response = await _get_tools_for_single_server(
|
||||
server, server_auth_header=None
|
||||
)
|
||||
|
||||
# Verify only allowed tools are in the response
|
||||
assert len(filtered_response) == 2
|
||||
|
|
@ -1768,7 +1823,9 @@ class TestMCPServerManager:
|
|||
new=AsyncMock(return_value=[tool1, tool2, tool3]),
|
||||
):
|
||||
# Call the REST endpoint helper
|
||||
all_tools_response = await _get_tools_for_single_server(server, server_auth_header=None)
|
||||
all_tools_response = await _get_tools_for_single_server(
|
||||
server, server_auth_header=None
|
||||
)
|
||||
|
||||
# Verify all tools are returned (no filtering)
|
||||
assert len(all_tools_response) == 3
|
||||
|
|
@ -1813,7 +1870,9 @@ class TestMCPServerManager:
|
|||
new=AsyncMock(return_value=[tool1, tool2]),
|
||||
):
|
||||
# Call the REST endpoint helper
|
||||
all_tools_response = await _get_tools_for_single_server(server, server_auth_header=None)
|
||||
all_tools_response = await _get_tools_for_single_server(
|
||||
server, server_auth_header=None
|
||||
)
|
||||
|
||||
# Verify all tools are returned (no filtering)
|
||||
assert len(all_tools_response) == 2
|
||||
|
|
@ -1883,7 +1942,9 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
|
|
@ -1926,7 +1987,9 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
|
||||
|
||||
|
|
@ -2071,7 +2134,9 @@ class TestMCPServerManager:
|
|||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
|
||||
|
||||
|
|
@ -2107,8 +2172,13 @@ class TestMCPServerManager:
|
|||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Tool deletepet is not allowed for server my_api_mcp" in exc_info.value.detail["error"]
|
||||
assert "Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
|
||||
assert (
|
||||
"Tool deletepet is not allowed for server my_api_mcp"
|
||||
in exc_info.value.detail["error"]
|
||||
)
|
||||
assert (
|
||||
"Contact proxy admin to allow this tool" in exc_info.value.detail["error"]
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_tool_without_broken_pipe_error(self):
|
||||
|
|
@ -2133,7 +2203,9 @@ class TestMCPServerManager:
|
|||
# Register the server and map a tool to it
|
||||
manager.registry = {"test-server": server}
|
||||
manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server"
|
||||
manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server"
|
||||
manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = (
|
||||
"test-server"
|
||||
)
|
||||
|
||||
# Create mock client that tracks call_tool usage
|
||||
mock_client = AsyncMock()
|
||||
|
|
@ -2157,7 +2229,9 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock proxy logging
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(return_value={})
|
||||
proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock(
|
||||
return_value={}
|
||||
)
|
||||
proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={})
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(return_value={})
|
||||
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
|
||||
|
|
@ -2222,7 +2296,9 @@ class TestMCPServerManager:
|
|||
# Verify MCPRequestHandler.get_allowed_mcp_servers was called with user_api_key_auth
|
||||
mock_get_allowed.assert_called_once()
|
||||
call_args = mock_get_allowed.call_args
|
||||
assert call_args[0][0] is user_api_key_auth # First positional arg should be user_api_key_auth
|
||||
assert (
|
||||
call_args[0][0] is user_api_key_auth
|
||||
) # First positional arg should be user_api_key_auth
|
||||
assert call_args[0][0].user_id == "user-123"
|
||||
assert call_args[0][0].object_permission_id == "perm_123"
|
||||
assert call_args[0][0].object_permission is not None
|
||||
|
|
@ -2467,7 +2543,10 @@ class TestMCPServerManagerUpstreamInstructionsCache:
|
|||
def test_get_returns_none_when_empty(self):
|
||||
"""Empty cache returns None for any key."""
|
||||
manager = MCPServerManager()
|
||||
assert manager._upstream_initialize_instructions_by_server_id.get("nonexistent") is None
|
||||
assert (
|
||||
manager._upstream_initialize_instructions_by_server_id.get("nonexistent")
|
||||
is None
|
||||
)
|
||||
|
||||
def test_remember_stores_stripped_value(self):
|
||||
"""_remember_upstream_initialize_instructions stores a stripped string."""
|
||||
|
|
@ -2475,7 +2554,9 @@ class TestMCPServerManagerUpstreamInstructionsCache:
|
|||
fake_server = MagicMock(server_id="srv")
|
||||
fake_client = MagicMock(_last_initialize_instructions=" hello \n")
|
||||
manager._remember_upstream_initialize_instructions(fake_server, fake_client)
|
||||
assert manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello"
|
||||
assert (
|
||||
manager._upstream_initialize_instructions_by_server_id.get("srv") == "hello"
|
||||
)
|
||||
|
||||
def test_remember_ignores_empty_string(self):
|
||||
"""Whitespace-only instructions are not stored."""
|
||||
|
|
@ -2553,7 +2634,9 @@ class TestMCPServerManagerExpandPermissionList:
|
|||
|
||||
def test_expands_server_name(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-usw1"] = self._make_server("id-usw1", server_name="a")
|
||||
manager.config_mcp_servers["id-usw1"] = self._make_server(
|
||||
"id-usw1", server_name="a"
|
||||
)
|
||||
|
||||
assert manager.expand_permission_list(["a"]) == ["id-usw1"]
|
||||
|
||||
|
|
@ -2577,7 +2660,9 @@ class TestMCPServerManagerExpandPermissionList:
|
|||
def test_name_collision_expands_to_all_matches(self):
|
||||
"""Two servers sharing a server_name both resolve — the documented behavior."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-config"] = self._make_server("id-config", server_name="shared")
|
||||
manager.config_mcp_servers["id-config"] = self._make_server(
|
||||
"id-config", server_name="shared"
|
||||
)
|
||||
manager.registry["id-db"] = self._make_server("id-db", server_name="shared")
|
||||
|
||||
assert sorted(manager.expand_permission_list(["shared"])) == [
|
||||
|
|
@ -2587,7 +2672,9 @@ class TestMCPServerManagerExpandPermissionList:
|
|||
|
||||
def test_searches_config_and_registry_union(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["cfg-id"] = self._make_server("cfg-id", server_name="a")
|
||||
manager.config_mcp_servers["cfg-id"] = self._make_server(
|
||||
"cfg-id", server_name="a"
|
||||
)
|
||||
manager.registry["reg-id"] = self._make_server("reg-id", server_name="b")
|
||||
|
||||
assert manager.expand_permission_list(["a"]) == ["cfg-id"]
|
||||
|
|
@ -2599,15 +2686,23 @@ class TestMCPServerManagerExpandPermissionList:
|
|||
servers whose server_name happens to equal that id.
|
||||
"""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["id-1"] = self._make_server("id-1", server_name="other_name")
|
||||
manager.config_mcp_servers["id-2"] = self._make_server("id-2", server_name="id-1")
|
||||
manager.config_mcp_servers["id-1"] = self._make_server(
|
||||
"id-1", server_name="other_name"
|
||||
)
|
||||
manager.config_mcp_servers["id-2"] = self._make_server(
|
||||
"id-2", server_name="id-1"
|
||||
)
|
||||
|
||||
assert manager.expand_permission_list(["id-1"]) == ["id-1"]
|
||||
|
||||
def test_mixed_ids_and_names_in_same_list(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="a")
|
||||
manager.config_mcp_servers["uuid-2"] = self._make_server("uuid-2", server_name="b")
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server(
|
||||
"uuid-1", server_name="a"
|
||||
)
|
||||
manager.config_mcp_servers["uuid-2"] = self._make_server(
|
||||
"uuid-2", server_name="b"
|
||||
)
|
||||
|
||||
# ["uuid-1", "b"] -> uuid-1 passes through, "b" resolves to uuid-2
|
||||
assert sorted(manager.expand_permission_list(["uuid-1", "b"])) == [
|
||||
|
|
@ -2618,7 +2713,9 @@ class TestMCPServerManagerExpandPermissionList:
|
|||
def test_deduplicates_overlapping_id_and_name_entries(self):
|
||||
"""If a list references the same server by both id and name, return it once."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="a")
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server(
|
||||
"uuid-1", server_name="a"
|
||||
)
|
||||
|
||||
assert manager.expand_permission_list(["uuid-1", "a"]) == ["uuid-1"]
|
||||
|
||||
|
|
@ -2628,10 +2725,14 @@ class TestMCPServerManagerExpandPermissionList:
|
|||
the cross-region portability the customer is asking for.
|
||||
"""
|
||||
usw1 = MCPServerManager()
|
||||
usw1.config_mcp_servers["hash-usw1"] = self._make_server("hash-usw1", server_name="a")
|
||||
usw1.config_mcp_servers["hash-usw1"] = self._make_server(
|
||||
"hash-usw1", server_name="a"
|
||||
)
|
||||
|
||||
usc1 = MCPServerManager()
|
||||
usc1.config_mcp_servers["hash-usc1"] = self._make_server("hash-usc1", server_name="a")
|
||||
usc1.config_mcp_servers["hash-usc1"] = self._make_server(
|
||||
"hash-usc1", server_name="a"
|
||||
)
|
||||
|
||||
assert usw1.expand_permission_list(["a"]) == ["hash-usw1"]
|
||||
assert usc1.expand_permission_list(["a"]) == ["hash-usc1"]
|
||||
|
|
@ -2660,14 +2761,18 @@ class TestMCPServerManagerExpandToolPermissions:
|
|||
concrete server_id, otherwise `.get(server_id)` misses and the tool
|
||||
restriction is silently dropped (caller treats None as allow-all)."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="my-alias")
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server(
|
||||
"uuid-a", server_name="my-alias"
|
||||
)
|
||||
|
||||
result = manager.expand_tool_permissions({"my-alias": ["read_file"]})
|
||||
assert result == {"uuid-a": ["read_file"]}
|
||||
|
||||
def test_passes_through_existing_server_id_key(self):
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha")
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server(
|
||||
"uuid-a", server_name="alpha"
|
||||
)
|
||||
|
||||
result = manager.expand_tool_permissions({"uuid-a": ["read_file"]})
|
||||
assert result == {"uuid-a": ["read_file"]}
|
||||
|
|
@ -2686,7 +2791,9 @@ class TestMCPServerManagerExpandToolPermissions:
|
|||
"""Two servers sharing a server_name both match; their tool lists get
|
||||
the restriction (matches the list-expansion collision semantics)."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server("uuid-1", server_name="shared")
|
||||
manager.config_mcp_servers["uuid-1"] = self._make_server(
|
||||
"uuid-1", server_name="shared"
|
||||
)
|
||||
manager.registry["uuid-2"] = self._make_server("uuid-2", server_name="shared")
|
||||
|
||||
result = manager.expand_tool_permissions({"shared": ["read_file"]})
|
||||
|
|
@ -2699,9 +2806,13 @@ class TestMCPServerManagerExpandToolPermissions:
|
|||
both refer to the same server, the tool lists are unioned rather
|
||||
than one overwriting the other."""
|
||||
manager = MCPServerManager()
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alias-a")
|
||||
manager.config_mcp_servers["uuid-a"] = self._make_server(
|
||||
"uuid-a", server_name="alias-a"
|
||||
)
|
||||
|
||||
result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["write_file"]})
|
||||
result = manager.expand_tool_permissions(
|
||||
{"uuid-a": ["read_file"], "alias-a": ["write_file"]}
|
||||
)
|
||||
assert sorted(result["uuid-a"]) == ["read_file", "write_file"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -2740,7 +2851,9 @@ class TestMCPServerManagerExpandToolPermissions:
|
|||
"litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_handle,
|
||||
patch("litellm.proxy._experimental.mcp_server.server.get_active_mcp_session") as mock_get_session,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.get_active_mcp_session"
|
||||
) as mock_get_session,
|
||||
):
|
||||
mock_session = MagicMock()
|
||||
mock_session.capabilities = "test_capabilities"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue