mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(mcp): keep tool-call logging non-blocking so it can't mask outcomes
Guard the success and failure logging added for the MCP tool-call span so a logging-layer exception cannot change the request outcome: a successful tool call still returns its result if success logging raises, and a failed call preserves the original error and its HTTP status if failure logging raises, rather than surfacing a generic 500. Addresses the Greptile review on the bare async_failure_handler call in the REST handler
This commit is contained in:
parent
97d3aae9cd
commit
6b1b115883
4 changed files with 181 additions and 15 deletions
|
|
@ -878,9 +878,16 @@ if MCP_AVAILABLE:
|
|||
except Exception as e:
|
||||
if logging_obj is not None:
|
||||
logging_obj.call_type = CallTypes.call_mcp_tool.value
|
||||
await logging_obj.async_failure_handler(
|
||||
e, traceback.format_exc(), start_time, datetime.now()
|
||||
)
|
||||
try:
|
||||
await logging_obj.async_failure_handler(
|
||||
e, traceback.format_exc(), start_time, datetime.now()
|
||||
)
|
||||
except Exception as logging_error:
|
||||
verbose_logger.exception(
|
||||
"MCP tool call failed and failure logging also raised; "
|
||||
"preserving original error: %s",
|
||||
logging_error,
|
||||
)
|
||||
raise
|
||||
return result
|
||||
except MCPMissingUserEnvVarsError as e:
|
||||
|
|
|
|||
|
|
@ -2732,18 +2732,24 @@ if MCP_AVAILABLE:
|
|||
response = CallToolResult(content=cast(Any, local_content), isError=False)
|
||||
|
||||
if litellm_logging_obj is not None:
|
||||
litellm_logging_obj.post_call(original_response=response)
|
||||
end_time = datetime.now()
|
||||
await litellm_logging_obj.async_post_mcp_tool_call_hook(
|
||||
kwargs=litellm_logging_obj.model_call_details,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value
|
||||
await litellm_logging_obj.async_success_handler(
|
||||
result=response, start_time=start_time, end_time=end_time
|
||||
)
|
||||
try:
|
||||
litellm_logging_obj.post_call(original_response=response)
|
||||
end_time = datetime.now()
|
||||
await litellm_logging_obj.async_post_mcp_tool_call_hook(
|
||||
kwargs=litellm_logging_obj.model_call_details,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
litellm_logging_obj.call_type = CallTypes.call_mcp_tool.value
|
||||
await litellm_logging_obj.async_success_handler(
|
||||
result=response, start_time=start_time, end_time=end_time
|
||||
)
|
||||
except Exception as logging_error:
|
||||
verbose_logger.exception(
|
||||
"MCP tool call succeeded but success logging raised: %s",
|
||||
logging_error,
|
||||
)
|
||||
return response
|
||||
|
||||
@client
|
||||
|
|
|
|||
|
|
@ -4366,6 +4366,66 @@ async def test_execute_mcp_tool_drives_success_logging():
|
|||
assert logging_obj.call_type == CallTypes.call_mcp_tool.value
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_returns_result_when_success_logging_fails():
|
||||
"""
|
||||
Logging is non-blocking: a tool call that succeeds must return its result
|
||||
even if async_success_handler raises (e.g. an exporter error), rather than
|
||||
surfacing the logging error as a failed tool call.
|
||||
"""
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
execute_mcp_tool,
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
server = MCPServer(
|
||||
server_id="server-123",
|
||||
name="test_server",
|
||||
alias="test_server",
|
||||
server_name="test_server",
|
||||
url="https://test-server.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
mcp_info={"server_name": "test_server"},
|
||||
)
|
||||
|
||||
tool_result = CallToolResult(
|
||||
content=[TextContent(type="text", text="ok")], isError=False
|
||||
)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
logging_obj.post_call = MagicMock()
|
||||
logging_obj.async_post_mcp_tool_call_hook = AsyncMock()
|
||||
logging_obj.async_success_handler = AsyncMock(
|
||||
side_effect=RuntimeError("otel exporter down")
|
||||
)
|
||||
|
||||
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool",
|
||||
new_callable=AsyncMock,
|
||||
return_value=tool_result,
|
||||
),
|
||||
patch.object(global_mcp_tool_registry, "get_tool", return_value=None),
|
||||
):
|
||||
result = await execute_mcp_tool(
|
||||
name="test_server-some_tool",
|
||||
arguments={"x": 1},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=user_auth,
|
||||
litellm_logging_obj=logging_obj,
|
||||
requested_server_id="server-123",
|
||||
)
|
||||
|
||||
assert result is tool_result
|
||||
logging_obj.async_success_handler.assert_awaited_once()
|
||||
|
||||
|
||||
def test_tool_name_matches_case_insensitive():
|
||||
"""Test that _tool_name_matches performs case-insensitive comparison.
|
||||
|
||||
|
|
|
|||
|
|
@ -1164,6 +1164,99 @@ class TestCallToolRestAPI:
|
|||
assert logging_obj.async_failure_handler.await_args.args[0] is tool_error
|
||||
assert logging_obj.call_type == CallTypes.call_mcp_tool.value
|
||||
|
||||
async def test_failure_logging_error_does_not_mask_original_error(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""
|
||||
Logging is non-blocking: if async_failure_handler itself raises, the
|
||||
original tool-call error and its HTTP status must still surface rather
|
||||
than being replaced by a generic 500 from the logging layer.
|
||||
"""
|
||||
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
||||
async def fake_get_allowed_mcp_servers(*args, **kwargs):
|
||||
return ["server-1"]
|
||||
|
||||
class StubServer:
|
||||
server_id = "server-1"
|
||||
alias = "server-1"
|
||||
server_name = "server-1"
|
||||
name = "stub"
|
||||
allowed_tools = None
|
||||
mcp_info = {"server_name": "stub"}
|
||||
available_on_public_internet = True
|
||||
auth_type = None
|
||||
|
||||
stub_server = StubServer()
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.async_failure_handler = AsyncMock(
|
||||
side_effect=RuntimeError("otel exporter down")
|
||||
)
|
||||
|
||||
async def fake_common_processing(self, **kwargs):
|
||||
return self.data, logging_obj
|
||||
|
||||
async def fake_execute_mcp_tool(**kwargs):
|
||||
raise HTTPException(status_code=403, detail="forbidden")
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"build_effective_auth_contexts",
|
||||
fake_contexts,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda server_id: stub_server if server_id == "server-1" else None,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_config",
|
||||
{},
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic",
|
||||
fake_common_processing,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"execute_mcp_tool",
|
||||
fake_execute_mcp_tool,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
request_payload = {
|
||||
"server_id": "server-1",
|
||||
"name": "demo-tool",
|
||||
"arguments": {"foo": "bar"},
|
||||
}
|
||||
request = _build_request(
|
||||
path="/mcp-rest/tools/call",
|
||||
method="POST",
|
||||
json_body=request_payload,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await rest_endpoints.call_tool_rest_api(
|
||||
request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
logging_obj.async_failure_handler.assert_awaited_once()
|
||||
|
||||
|
||||
class TestGetToolsForSingleServer:
|
||||
"""Test _get_tools_for_single_server with object_permission filtering"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue