fix(mcp): log proxy tool dispatch exceptions (#40351)

This commit is contained in:
tin-berri 2026-09-08 19:56:21 -07:00 committed by GitHub
parent 24ef3ec63b
commit 902dd7b2b6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 138 additions and 13 deletions

View file

@ -1017,18 +1017,45 @@ if MCP_AVAILABLE:
if name == MCP_PROXY_CALL_TOOL_NAME
else None
)
proxy_result: Final = await handle_mcp_proxy_tool(
name=name,
arguments=arguments or {}, # mutable-ok: proxy handler payload
user_api_key_dict=user_api_key_auth,
client_ip=client_ip,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=proxy_logging_obj,
)
try:
proxy_result: Final = await handle_mcp_proxy_tool(
name=name,
arguments=arguments or {}, # mutable-ok: proxy handler payload
user_api_key_dict=user_api_key_auth,
client_ip=client_ip,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=proxy_logging_obj,
)
except Exception as exc:
if proxy_logging_obj is not None:
from litellm.proxy.proxy_server import proxy_logging_obj as request_logging_obj
failure_end: Final = datetime.now() # noqa: DTZ005 # matches the logging pipeline start time
failure_traceback: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG)
try:
proxy_logging_obj.failure_handler(exc, failure_traceback, proxy_call_start, failure_end)
await proxy_logging_obj.async_failure_handler(
exc, failure_traceback, proxy_call_start, failure_end
)
if not isinstance(exc, MCPUpstreamAuthError):
await request_logging_obj.post_call_failure_hook(
request_data={ # mutable-ok: failure hook mutates its request payload
"name": name,
"arguments": arguments,
"litellm_logging_obj": proxy_logging_obj,
},
original_exception=exc,
user_api_key_dict=user_api_key_auth,
route="/mcp/call_tool",
traceback_str=failure_traceback,
)
except Exception:
verbose_logger.exception("Error logging failed MCP proxy tool call")
raise
if proxy_logging_obj is not None:
return await _fire_mcp_tool_call_logging(
logging_obj=proxy_logging_obj,

View file

@ -471,6 +471,7 @@ class ProxyCallRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.events: queue.Queue[str] = queue.Queue()
self.failures: queue.Queue[str] = queue.Queue()
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
@ -479,6 +480,13 @@ class ProxyCallRecorder(CustomLogger):
if isinstance(payload, dict) and payload.get("call_type") == "call_mcp_tool":
self.events.put(json.dumps(payload, default=str))
async def async_log_failure_event(
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
payload = kwargs.get("standard_logging_object")
if isinstance(payload, dict) and payload.get("call_type") == "call_mcp_tool":
self.failures.put(json.dumps(payload, default=str))
proxy_call_recorder = ProxyCallRecorder()
@ -636,6 +644,28 @@ class TestProxyMcpAuthorizationScope:
assert payload["metadata"]["mcp_tool_call_metadata"]["name"] == "add"
assert payload["metadata"]["mcp_tool_call_metadata"]["namespaced_tool_name"] == "math_restricted/add"
@pytest.mark.asyncio
async def test_proxy_scope_exception_returns_iserror_and_emits_failure_log(self, proxy_server_url: str) -> None:
async with _scoped_session(
proxy_server_url,
"sk-none",
**{"x-mcp-servers": "math_restricted", "x-litellm-call-id": "proxy-scope-denial"},
) as session:
result = await session.call_tool("call_tool", {"tool_id": "denied-scope", "arguments": {}})
assert result.isError is True
assert result.content[0].text == (
"Error: The key is not allowed to access the requested MCP servers: math_restricted"
)
async with asyncio.timeout(10):
while True:
payload = json.loads(await asyncio.to_thread(proxy_call_recorder.failures.get, True, 5))
if payload["id"] == "proxy-scope-denial":
break
assert payload["call_type"] == "call_mcp_tool"
assert payload["status"] == "failure"
assert payload["response_cost"] == 0
assert "math_restricted" in payload["error_str"]
@pytest.mark.parametrize("arguments", ["wrong", False, None, [], 0])
def test_handler_rejects_non_object_arguments(
self, proxy_server_url: str, _proxy_server: ProxyRig, arguments: object

View file

@ -1,10 +1,16 @@
import json
from datetime import datetime
import pytest
from fastapi import HTTPException
from mcp.shared.exceptions import McpError
from pydantic import AnyUrl
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._experimental.mcp_server import server
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth
AUTH = UserAPIKeyAuth(api_key="key")
@ -48,3 +54,65 @@ async def test_proxy_rejects_non_tool_protocol_operations() -> None:
await server.list_resource_templates()
with pytest.raises(McpError):
await server.read_resource(AnyUrl("https://example.com/resource"))
class FailureRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.events: list[tuple[str, str]] = []
async def async_log_failure_event(
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.events.append(("failure", json.dumps(kwargs.get("standard_logging_object"), default=str)))
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
self.events.append(("success", json.dumps(kwargs.get("standard_logging_object"), default=str)))
async def async_post_call_failure_hook(
self,
request_data: dict[str, object],
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: str | None = None,
) -> None:
self.events.append(("post_failure", json.dumps(request_data, default=str)))
@pytest.mark.asyncio
@pytest.mark.usefixtures("proxy_mode")
async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.MonkeyPatch) -> None:
recorder = FailureRecorder()
monkeypatch.setattr(litellm, "callbacks", [recorder])
auth = UserAPIKeyAuth(
api_key="scope-denial-key-hash",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="denied", mcp_servers=["no-mcp-servers"]),
)
arguments = {"tool_id": "denied-scope", "arguments": {}}
with pytest.raises(HTTPException) as denied:
await server._dispatch_virtual_mcp_tool(
name="call_tool",
arguments=arguments,
user_api_key_auth=auth,
client_ip=None,
mcp_servers=["ungranted"],
raw_headers={"authorization": "Bearer raw-scope-secret", "x-litellm-call-id": "scope-denial"},
)
assert denied.value.status_code == 403
assert denied.value.detail == {"error": "The key is not allowed to access the requested MCP servers: ungranted"}
assert [kind for kind, _ in recorder.events] == ["failure", "post_failure"]
payload = json.loads(recorder.events[0][1])
assert payload["id"] == "scope-denial"
assert payload["call_type"] == "call_mcp_tool"
assert payload["status"] == "failure"
assert payload["response_cost"] == 0
assert "ungranted" in payload["error_str"]
hook_payload = json.loads(recorder.events[1][1])
assert hook_payload["standard_logging_object"] == payload
assert hook_payload["arguments"] == arguments
assert "raw_headers" not in hook_payload
assert "raw-scope-secret" not in recorder.events[1][1]