mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(mcp): name the agent veto on scoped tool calls too
This commit is contained in:
parent
5e2d98d49b
commit
d6b8cb93b3
4 changed files with 110 additions and 6 deletions
|
|
@ -821,11 +821,7 @@ if MCP_AVAILABLE:
|
|||
from mcp.shared.exceptions import McpError
|
||||
from mcp.types import INVALID_REQUEST, ErrorData
|
||||
|
||||
detail: Final = e.detail
|
||||
message: Final = (
|
||||
str(detail.get("error")) if isinstance(detail, dict) and detail.get("error") else str(detail)
|
||||
)
|
||||
raise McpError(ErrorData(code=INVALID_REQUEST, message=message)) from e
|
||||
raise McpError(ErrorData(code=INVALID_REQUEST, message=_http_detail_message(e.detail))) from e
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error in list_tools endpoint: %s", e)
|
||||
# Return empty list instead of failing completely
|
||||
|
|
@ -1096,6 +1092,7 @@ if MCP_AVAILABLE:
|
|||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
client_ip=_client_ip,
|
||||
host_progress_callback=host_progress_callback,
|
||||
**data, # for logging
|
||||
)
|
||||
|
|
@ -1129,7 +1126,7 @@ if MCP_AVAILABLE:
|
|||
except HTTPException as e:
|
||||
verbose_logger.error("HTTPException in MCP tool call: %s", e)
|
||||
return CallToolResult(
|
||||
content=[TextContent(text=f"Error: {e.detail}", type="text")],
|
||||
content=[TextContent(text=f"Error: {_http_detail_message(e.detail)}", type="text")],
|
||||
isError=True,
|
||||
)
|
||||
except MCPUpstreamAuthError as e:
|
||||
|
|
@ -1447,6 +1444,9 @@ if MCP_AVAILABLE:
|
|||
|
||||
return allowed_mcp_servers
|
||||
|
||||
def _http_detail_message(detail: object) -> str:
|
||||
return str(detail.get("error")) if isinstance(detail, dict) and detail.get("error") else str(detail)
|
||||
|
||||
def _server_answers_to(server: MCPServer, name: str) -> bool:
|
||||
requested: Final = name.lower()
|
||||
return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known)
|
||||
|
|
@ -3148,6 +3148,7 @@ if MCP_AVAILABLE:
|
|||
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
|
||||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
|
|
@ -3178,6 +3179,12 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
)
|
||||
if mcp_servers and not allowed_mcp_servers:
|
||||
await _raise_denied_scoped_mcp_access(
|
||||
requested_names=mcp_servers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
if not allowed_mcp_servers:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
|
|||
|
|
@ -307,6 +307,7 @@ async def handle_mcp_tool_call(
|
|||
) -> CallToolResult:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_allowed_mcp_servers,
|
||||
_raise_denied_scoped_mcp_access,
|
||||
execute_mcp_tool,
|
||||
)
|
||||
|
||||
|
|
@ -315,6 +316,12 @@ async def handle_mcp_tool_call(
|
|||
mcp_servers=mcp_servers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
if mcp_servers and not allowed_mcp_servers:
|
||||
await _raise_denied_scoped_mcp_access(
|
||||
requested_names=mcp_servers,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
|
||||
# Reject before dispatch when the key has no accessible servers; otherwise an
|
||||
# unprefixed local tool name would fall through to the local registry in
|
||||
|
|
|
|||
|
|
@ -1593,6 +1593,32 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(
|
|||
assert exc_info.value.error.message == denial_message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_renders_denial_message_not_detail_dict():
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import mcp_server_tool_call
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
denial_message = "MCP server 'github' is not available to this key: the key is bound to agent 'agent-123'"
|
||||
denial = HTTPException(status_code=403, detail={"error": denial_message})
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
|
||||
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
|
||||
new=AsyncMock(return_value=(None, None, None, None, None, None, None)),
|
||||
),
|
||||
patch( # test-quality-ok: the tool-call helper is the handler's only collaborator; the suite's seam
|
||||
"litellm.proxy._experimental.mcp_server.server.call_mcp_tool",
|
||||
new=AsyncMock(side_effect=denial),
|
||||
),
|
||||
):
|
||||
result = await mcp_server_tool_call("github-search_issues", {})
|
||||
|
||||
assert result.isError is True
|
||||
assert result.content[0].text == f"Error: {denial_message}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_body_with_none_arguments():
|
||||
"""Test that proxy_server_request body handles None arguments correctly"""
|
||||
|
|
@ -3742,6 +3768,35 @@ async def test_call_mcp_tool_user_unauthorized_access():
|
|||
assert "User not allowed to call this tool" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_mcp_tool_scoped_denial_names_the_binding_agent():
|
||||
from litellm.proxy._experimental.mcp_server.server import call_mcp_tool
|
||||
|
||||
agent_bound_key = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="agent-123")
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: the server registry is a module-level singleton; the suite's only seam
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
|
||||
AsyncMock(return_value=[]),
|
||||
),
|
||||
patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
_scope_resolver({"github": "srv-github"}),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await call_mcp_tool(
|
||||
name="github-search_issues",
|
||||
arguments={},
|
||||
user_api_key_auth=agent_bound_key,
|
||||
mcp_servers=["github"],
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "MCP server 'github'" in exc_info.value.detail["error"]
|
||||
assert "agent 'agent-123'" in exc_info.value.detail["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_mcp_tool_unauthorized_403_does_not_leak_server_credentials():
|
||||
"""Regression for LIT-4703 / GH #29936.
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ Covers:
|
|||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -1229,3 +1230,37 @@ class TestMcpServerToolCallErrorHandling:
|
|||
|
||||
assert result.isError is True
|
||||
assert "User not allowed to call this tool" in result.content[0].text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.tool_search import handle_mcp_tool_call
|
||||
|
||||
agent_bound_key = UserAPIKeyAuth(api_key="test_key", agent_id="agent-123")
|
||||
|
||||
async def resolve(user_api_key_auth, mcp_servers, client_ip=None):
|
||||
if user_api_key_auth.agent_id:
|
||||
return []
|
||||
return [
|
||||
SimpleNamespace(
|
||||
server_id="srv-github", server_name="github", alias=None, short_prefix=None, access_groups=[]
|
||||
)
|
||||
]
|
||||
|
||||
with patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
new=AsyncMock(side_effect=resolve),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handle_mcp_tool_call(
|
||||
tool_name="github-create_issue",
|
||||
arguments={},
|
||||
user_api_key_dict=agent_bound_key,
|
||||
mcp_servers=["github"],
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "MCP server 'github'" in exc_info.value.detail["error"]
|
||||
assert "agent 'agent-123'" in exc_info.value.detail["error"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue