mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): honor x-litellm-tags on the MCP gateway's tools/list and tools/call
LLM routes read the header in add_litellm_data_to_request, so caller tags land in LiteLLM_SpendLogs.request_tags. Nothing read it on the MCP gateway: the tool-call handler hands that same helper a synthetic request carrying only a content type, so the header never reached the tag merge, and the list_tools spend log has a request_tags parameter no caller populates Tool calls now carry the caller's tags in the body, which that helper already reads, and list_tools falls back to the header when no tags were passed in. Per-application attribution behind a gateway that stamps the header now works the same for MCP traffic as it does for chat completions
This commit is contained in:
parent
f6b9518ddb
commit
cfab4f62d3
2 changed files with 211 additions and 2 deletions
|
|
@ -184,6 +184,17 @@ def _mcp_session_id_from_headers(
|
|||
return None
|
||||
|
||||
|
||||
def _request_tags_from_raw_headers(
|
||||
raw_headers: dict[str, str] | None,
|
||||
) -> list[str] | None:
|
||||
"""The caller's ``x-litellm-tags``, parsed by the same helper the LLM routes use so an
|
||||
MCP operation and a chat completion attribute an identical header identically."""
|
||||
if not raw_headers:
|
||||
return None
|
||||
headers = {key.lower(): value for key, value in raw_headers.items() if isinstance(key, str)}
|
||||
return LiteLLMProxyRequestSetup.add_request_tag_to_metadata(llm_router=None, headers=headers, data={})
|
||||
|
||||
|
||||
def _jsonrpc_text_has_top_level_method(text: str) -> bool:
|
||||
"""Whether a (possibly truncated) JSON-RPC envelope has a ``method`` key at
|
||||
the root object's top level.
|
||||
|
|
@ -1034,7 +1045,12 @@ if MCP_AVAILABLE:
|
|||
|
||||
host_progress_callback: Final = _capture_host_progress_callback(server)
|
||||
# Create a body date for logging
|
||||
body_data: Final = {"name": name, "arguments": arguments}
|
||||
request_tags: Final = _request_tags_from_raw_headers(raw_headers)
|
||||
body_data: Final = {
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
**({"tags": request_tags} if request_tags else {}),
|
||||
}
|
||||
# Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A)
|
||||
chain_id: Final = get_chain_id_from_headers(raw_headers)
|
||||
if chain_id:
|
||||
|
|
@ -1890,6 +1906,7 @@ if MCP_AVAILABLE:
|
|||
list_tools_call_id: Final = str(uuid.uuid4())
|
||||
# Derive trace_id from raw_headers when not explicitly passed (same as A2A / MCP call_tool)
|
||||
effective_litellm_trace_id: Final = litellm_trace_id or get_chain_id_from_headers(raw_headers)
|
||||
effective_request_tags: Final = request_tags or _request_tags_from_raw_headers(raw_headers)
|
||||
spend_logs_metadata: Final[dict[str, object]] = {
|
||||
"mcp_operation": "list_tools",
|
||||
}
|
||||
|
|
@ -1905,7 +1922,7 @@ if MCP_AVAILABLE:
|
|||
"litellm_trace_id": effective_litellm_trace_id,
|
||||
"metadata": {
|
||||
"spend_logs_metadata": spend_logs_metadata,
|
||||
**({"tags": request_tags} if request_tags else {}),
|
||||
**({"tags": effective_request_tags} if effective_request_tags else {}),
|
||||
},
|
||||
# Provide a small input payload for standard logging
|
||||
"input": [
|
||||
|
|
|
|||
|
|
@ -116,6 +116,48 @@ async def test_mcp_server_tool_call_body_contains_request_data():
|
|||
assert body["arguments"] == tool_arguments
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_carries_x_litellm_tags_header_into_request_data():
|
||||
"""The tool-call handler hands `add_litellm_data_to_request` a synthetic request that carries only
|
||||
a content type, so the caller's `x-litellm-tags` never reached the tag merge that runs there and
|
||||
the spend log for a tools/call had no tags. The tags travel in the body instead, which that same
|
||||
helper already reads, so the header attributes MCP traffic exactly as it does an LLM route."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
mcp_server_tool_call,
|
||||
set_auth_context,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
set_auth_context(
|
||||
UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
|
||||
raw_headers={"X-LiteLLM-Tags": "application:orders, service:checkout"},
|
||||
)
|
||||
|
||||
captured_data = {}
|
||||
|
||||
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config):
|
||||
captured_data.update(data)
|
||||
return data
|
||||
|
||||
async def mock_call_mcp_tool(*args, **kwargs):
|
||||
return [{"type": "text", "text": "mocked response"}]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request",
|
||||
mock_add_litellm_data_to_request,
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.call_mcp_tool",
|
||||
mock_call_mcp_tool,
|
||||
):
|
||||
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
|
||||
await mcp_server_tool_call("test_tool", {"param1": "value1"})
|
||||
|
||||
assert captured_data["tags"] == ["application:orders", "service:checkout"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror():
|
||||
"""The MCP session manager serializes handler exceptions as JSON-RPC errors, so a mid-session
|
||||
|
|
@ -4436,6 +4478,156 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
|
|||
assert spend_meta["per_server_list_outcomes"] == {"server_a": {"status": "ok", "tool_count": 1}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_from_mcp_servers_takes_list_tools_tags_from_x_litellm_tags_header():
|
||||
"""A gateway that stamps `x-litellm-tags` on proxied traffic gets per-application attribution on
|
||||
LLM routes; list_tools must read the same header so MCP usage is not stuck under the shared key.
|
||||
Nothing populated `request_tags`, so the header was the only source and it was being dropped."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from mcp.types import Tool as MCPTool
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
|
||||
|
||||
server_a = MagicMock(name="server_a_obj")
|
||||
server_a.name = "server_a"
|
||||
server_a.alias = "server_a"
|
||||
server_a.server_name = "server_a"
|
||||
server_a.server_id = "a"
|
||||
server_a.auth_type = None
|
||||
server_a.extra_headers = None
|
||||
|
||||
tool_1 = MCPTool(name="server_a-tool_1", description="test tool", inputSchema={"type": "object"})
|
||||
|
||||
dummy_logging_obj = MagicMock()
|
||||
dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}}
|
||||
dummy_logging_obj.async_success_handler = AsyncMock()
|
||||
function_setup_kwargs = {}
|
||||
|
||||
def _capture_function_setup(*_args, **kwargs):
|
||||
function_setup_kwargs.update(kwargs)
|
||||
return dummy_logging_obj, None
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
new=AsyncMock(return_value=[server_a]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
|
||||
side_effect=lambda tools, _server: tools,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
|
||||
new=AsyncMock(side_effect=lambda tools, **_: tools),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.function_setup",
|
||||
side_effect=_capture_function_setup,
|
||||
),
|
||||
):
|
||||
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
|
||||
|
||||
listing = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_auth,
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=["server_a"],
|
||||
mcp_server_auth_headers=None,
|
||||
raw_headers={"X-LiteLLM-Tags": "application:orders, service:checkout"},
|
||||
log_list_tools_to_spendlogs=True,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
)
|
||||
|
||||
assert listing.tools == [tool_1]
|
||||
assert function_setup_kwargs["metadata"]["tags"] == ["application:orders", "service:checkout"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_from_mcp_servers_prefers_explicit_request_tags_over_the_header():
|
||||
"""`request_tags` is the resolved value a caller passes in; a header must not override it."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from mcp.types import Tool as MCPTool
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
|
||||
|
||||
server_a = MagicMock(name="server_a_obj")
|
||||
server_a.name = "server_a"
|
||||
server_a.alias = "server_a"
|
||||
server_a.server_name = "server_a"
|
||||
server_a.server_id = "a"
|
||||
server_a.auth_type = None
|
||||
server_a.extra_headers = None
|
||||
|
||||
tool_1 = MCPTool(name="server_a-tool_1", description="test tool", inputSchema={"type": "object"})
|
||||
|
||||
dummy_logging_obj = MagicMock()
|
||||
dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}}
|
||||
dummy_logging_obj.async_success_handler = AsyncMock()
|
||||
function_setup_kwargs = {}
|
||||
|
||||
def _capture_function_setup(*_args, **kwargs):
|
||||
function_setup_kwargs.update(kwargs)
|
||||
return dummy_logging_obj, None
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
|
||||
new=AsyncMock(return_value=[server_a]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers",
|
||||
return_value=(None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
) as mock_manager,
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools",
|
||||
side_effect=lambda tools, _server: tools,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions",
|
||||
new=AsyncMock(side_effect=lambda tools, **_: tools),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.function_setup",
|
||||
side_effect=_capture_function_setup,
|
||||
),
|
||||
):
|
||||
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
|
||||
|
||||
await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_auth,
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=["server_a"],
|
||||
mcp_server_auth_headers=None,
|
||||
raw_headers={"x-litellm-tags": "from-header"},
|
||||
log_list_tools_to_spendlogs=True,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
request_tags=["explicit"],
|
||||
)
|
||||
|
||||
assert function_setup_kwargs["metadata"]["tags"] == ["explicit"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fails():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue