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:
onatozmenn 2026-08-04 14:59:44 +03:00
parent f6b9518ddb
commit cfab4f62d3
No known key found for this signature in database
2 changed files with 211 additions and 2 deletions

View file

@ -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": [

View file

@ -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():
"""