fix(mcp): expose client HTTP headers to logging callbacks and hooks (#36724)

* fix(mcp): expose client HTTP headers to logging callbacks and hooks

MCP protocol tool calls built a synthetic Request with only content-type, so metadata.headers reaching logging callbacks and guardrails was empty while /mcp-rest/tools/call exposed the full set. Rebuild the synthetic request from the connection's raw headers (shared with the sampling path), and pass sanitized headers to the pre-call hook, the MCP to LLM guardrail bridge and the Responses API MCP bridge. Credential headers stay masked and proxy key headers stripped.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): strip custom proxy key and upstream MCP credential headers from logging copies

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(mcp): make client side auth header name accessor public

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): strip custom proxy key and client redaction opt-out from mcp headers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(mcp): drop custom proxy key header in the synthetic request builder

Strips general_settings.litellm_key_header_name in build_synthetic_mcp_request so every caller, including sampling, is covered, and reverts passing general_settings into add_litellm_data_to_request on the tool call path since that also switches on enforced_params.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: shivam <shivam@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-08-13 20:07:16 -07:00 • committed by GitHub
parent efbdb6901a
commit 59eeae374c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 483 additions and 124 deletions

View file

@ -1190,7 +1190,7 @@ class MCPRequestHandler:
DEPRECATED: This method is deprecated in favor of server-specific auth headers using the format x-mcp-{{server_alias}}-{{header_name}} instead.
"""
mcp_client_side_auth_header_name: Final[str] = MCPRequestHandler._get_mcp_client_side_auth_header_name()
mcp_client_side_auth_header_name: Final[str] = MCPRequestHandler.get_mcp_client_side_auth_header_name()
auth_header: Final = headers.get(mcp_client_side_auth_header_name)
if auth_header:
verbose_logger.warning(
@ -1265,7 +1265,7 @@ class MCPRequestHandler:
return oauth2_headers
@staticmethod
def _get_mcp_client_side_auth_header_name() -> str:
def get_mcp_client_side_auth_header_name() -> str:
"""
Get the header name used to pass the MCP auth header to the MCP server

View file

@ -118,6 +118,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
is_short_mcp_tool_prefix_enabled,
iter_known_server_prefixes,
iter_known_tool_name_spellings,
logging_safe_mcp_headers,
match_known_server_prefix,
match_known_tool_name,
merge_mcp_headers,
@ -4603,6 +4604,7 @@ class MCPServerManager:
),
"user_api_key_hash": (getattr(user_api_key_auth, "api_key_hash", None) if user_api_key_auth else None),
"incoming_bearer_token": incoming_bearer_token,
"headers": logging_safe_mcp_headers(raw_headers),
}
# Create MCP request object for processing

View file

@ -1042,100 +1042,15 @@ def _build_sampling_request(
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
) -> "Request":
"""Build a synthetic FastAPI Request for sampling sub-calls.
"""The synthetic FastAPI Request for sampling sub-calls, carrying the original
MCP connection's headers and client IP."""
from litellm.proxy._experimental.mcp_server.utils import build_synthetic_mcp_request
Converts the original MCP connection's HTTP headers into ASGI
scope format so that ``add_litellm_data_to_request`` can apply
header-dependent guardrails, tag-based routing, trace correlation,
and ``forward_llm_provider_auth_headers``.
Key fields populated:
- **headers**: All original HTTP headers are forwarded (except
hop-by-hop: content-length, transfer-encoding). This ensures
``traceparent``, ``authorization``, ``user-agent``, and
``x-litellm-api-key`` are visible to pre-call utils.
- **client**: The ASGI ``(host, port)`` tuple so that
``request.client.host`` returns the real client IP for
IP-based routing and guardrails.
- **server**: Derived from the running proxy's ``server_host``
/ ``server_port`` when available, avoiding the misleading
``127.0.0.1:0`` placeholder.
- **x-forwarded-for**: Injected from ``client_ip`` if the
original headers don't already carry it, as a fallback for
IP attribution.
"""
from fastapi import Request
# --- Build ASGI headers ---
_scope_headers: Final[list[tuple[bytes, bytes]]] = [(b"content-type", b"application/json")]
# Hop-by-hop headers that must NOT be forwarded into the
# synthetic request (they describe the original HTTP framing,
# not the logical request).
_HOP_BY_HOP: Final = frozenset(
{
"content-length",
"transfer-encoding",
"connection",
"keep-alive",
"upgrade",
"te",
"trailer",
}
return build_synthetic_mcp_request(
path="/mcp/sampling/createMessage",
raw_headers=raw_headers,
client_ip=client_ip,
)
if raw_headers:
for hdr_name, hdr_value in raw_headers.items():
_key = hdr_name.lower()
# Skip content-type (already set), x-forwarded-for (use resolved
# client_ip instead to prevent spoofing), and hop-by-hop headers
if _key in {"content-type", "x-forwarded-for"} or _key in _HOP_BY_HOP:
continue
_scope_headers.append(
(
_key.encode("latin-1", errors="replace"),
hdr_value.encode("utf-8"),
)
)
# Inject x-forwarded-for from captured client_ip if the
# original headers don't already carry it
if client_ip and not any(h[0] == b"x-forwarded-for" for h in _scope_headers):
_scope_headers.append((b"x-forwarded-for", client_ip.encode("utf-8")))
# --- Derive server (host, port) from the running proxy ---
_server_host = "127.0.0.1"
_server_port = 4000 # LiteLLM default
try:
from litellm.proxy import proxy_server
_proxy_host: Final[str | None] = getattr(proxy_server, "server_host", None)
_proxy_port: Final[str | int | None] = getattr(proxy_server, "server_port", None)
if _proxy_host:
_server_host = str(_proxy_host)
if _proxy_port:
_server_port = int(_proxy_port)
except (ImportError, AttributeError, TypeError, ValueError):
pass
# --- Build ASGI client tuple for request.client.host ---
_client_tuple = None
if client_ip:
_client_tuple = (client_ip, 0)
scope: Final[dict[str, object]] = {
"type": "http",
"method": "POST",
"path": "/mcp/sampling/createMessage",
"scheme": "http",
"server": (_server_host, _server_port),
"query_string": b"",
"root_path": "",
"headers": _scope_headers,
}
if _client_tuple is not None:
scope["client"] = _client_tuple
return Request(scope=scope)
async def _build_completion_kwargs(

View file

@ -58,9 +58,11 @@ from litellm.proxy._experimental.mcp_server.utils import (
LITELLM_MCP_SERVER_VERSION,
MCPMissingUserEnvVarsError,
add_server_prefix_to_name,
build_synthetic_mcp_request,
extract_mcp_tool_result_error_message,
get_server_prefix,
iter_known_server_prefixes,
logging_safe_mcp_headers,
match_known_tool_name,
)
from litellm.proxy._types import (
@ -860,11 +862,11 @@ if MCP_AVAILABLE:
name: str,
arguments: dict[str, object],
user_api_key_auth: UserAPIKeyAuth,
raw_headers: Mapping[str, str] | None = None,
client_ip: str | None = None,
) -> LiteLLMLoggingObj | None:
"""Run the pre-call pipeline (guardrails + logging setup) for a virtual
mcp_tool_call so the SSE path spend-logs like the REST path."""
from fastapi import Request
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
@ -874,13 +876,10 @@ if MCP_AVAILABLE:
proxy_logging_obj,
)
request: Final = Request(
scope={
"type": "http",
"method": "POST",
"path": "/mcp/tools/call",
"headers": [(b"content-type", b"application/json")],
}
request: Final = build_synthetic_mcp_request(
path="/mcp/tools/call",
raw_headers=raw_headers,
client_ip=client_ip,
)
_, virtual_logging_obj = await ProxyBaseLLMRequestProcessing(
data={"name": name, "arguments": arguments}
@ -952,7 +951,11 @@ if MCP_AVAILABLE:
assert user_api_key_auth is not None # guaranteed by the flag check above
virtual_logging_obj: Final = await _build_virtual_call_logging_obj(
name=name, arguments=args, user_api_key_auth=user_api_key_auth
name=name,
arguments=args,
user_api_key_auth=user_api_key_auth,
raw_headers=raw_headers,
client_ip=client_ip,
)
return await handle_mcp_tool_call(
tool_name=args.get("tool_name", ""),
@ -979,7 +982,6 @@ if MCP_AVAILABLE:
Raises:
HTTPException: If tool not found or arguments missing
"""
from fastapi import Request
from mcp.server.lowlevel.server import request_ctx
from mcp.types import CallToolResult
@ -1041,13 +1043,10 @@ if MCP_AVAILABLE:
body_data["litellm_trace_id"] = chain_id
body_data["litellm_session_id"] = chain_id
request: Final = Request(
scope={
"type": "http",
"method": "POST",
"path": "/mcp/tools/call",
"headers": [(b"content-type", b"application/json")],
}
request: Final = build_synthetic_mcp_request(
path="/mcp/tools/call",
raw_headers=raw_headers,
client_ip=_client_ip,
)
if user_api_key_auth is not None:
data = await add_litellm_data_to_request(
@ -1905,6 +1904,7 @@ if MCP_AVAILABLE:
"litellm_trace_id": effective_litellm_trace_id,
"metadata": {
"spend_logs_metadata": spend_logs_metadata,
"headers": logging_safe_mcp_headers(raw_headers),
**({"tags": request_tags} if request_tags else {}),
},
# Provide a small input payload for standard logging

View file

@ -7,6 +7,7 @@ import importlib
import json
import os
import re
import typing
from collections.abc import Iterable, Iterator, Mapping, MutableMapping, MutableSequence
from collections.abc import Set as AbstractSet
from typing import Any, Final, Protocol
@ -14,6 +15,9 @@ from urllib.parse import quote
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if typing.TYPE_CHECKING:
from fastapi import Request
class _McpServerLike(Protocol):
@property
@ -862,3 +866,146 @@ def set_mcp_tool_result_structured_content(result: object, value: object) -> boo
return True
except (AttributeError, TypeError, ValueError):
return False
_HOP_BY_HOP_HEADERS: Final = frozenset(
{
"content-length",
"transfer-encoding",
"connection",
"keep-alive",
"upgrade",
"te",
"trailer",
}
)
_SYNTHETIC_REQUEST_EXCLUDED_HEADERS: Final = _HOP_BY_HOP_HEADERS | frozenset({"content-type", "x-forwarded-for"})
_SYNTHETIC_REQUEST_SERVER: Final = ("127.0.0.1", 4000)
_MCP_SERVER_AUTH_HEADER_PREFIX: Final = "x-mcp-"
def _custom_litellm_key_header_name() -> str | None:
"""``general_settings.litellm_key_header_name``, the deployment's custom header name for
the proxy virtual key, so it is stripped from observability copies like the standard ones."""
try:
from litellm.proxy.proxy_server import general_settings
except ImportError:
return None
return general_settings.get("litellm_key_header_name") if general_settings else None
def _mcp_client_side_auth_header_name() -> str:
"""The header name the client passes the upstream MCP credential in, falling back to the
default when ``general_settings`` is unavailable (the SDK, outside a running proxy)."""
from .auth.user_api_key_auth_mcp import MCPRequestHandler
try:
return MCPRequestHandler.get_mcp_client_side_auth_header_name()
except ImportError:
return MCPRequestHandler.LITELLM_MCP_AUTH_HEADER_NAME
def _upstream_credential_headers(header_names: Iterable[str]) -> frozenset[str]:
"""Lowercased names of the headers in ``header_names`` that carry an upstream MCP
credential rather than request context: the configured client side auth header and
the per-server ``x-mcp-{alias}-{header}`` family. ``clean_headers`` only knows the
credential headers of the chat completions path, so these are dropped on top of it.
"""
from .auth.user_api_key_auth_mcp import MCPRequestHandler
non_credential: Final = frozenset(
{
MCPRequestHandler.LITELLM_MCP_SERVERS_HEADER_NAME.lower(),
MCPRequestHandler.LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME.lower(),
}
)
client_side_auth: Final = _mcp_client_side_auth_header_name().lower()
return frozenset(
name
for name in (raw_name.lower() for raw_name in header_names)
if name == client_side_auth or (name.startswith(_MCP_SERVER_AUTH_HEADER_PREFIX) and name not in non_credential)
)
def build_synthetic_mcp_request(
*,
path: str,
raw_headers: Mapping[str, str] | None = None,
client_ip: str | None = None,
) -> "Request":
"""A synthetic FastAPI ``Request`` carrying the MCP connection's HTTP headers.
The MCP protocol transports do not hand a per-call ``Request`` to the tool
handlers, so one is reconstructed from the connection's ``raw_headers``. That
lets ``add_litellm_data_to_request`` derive ``metadata.headers``,
``proxy_server_request``, header-based tags, guardrails and trace correlation
exactly as on the chat completions path. Hop-by-hop headers describe the
original HTTP framing rather than the logical request, so they are dropped, and
``x-forwarded-for`` comes from the resolved ``client_ip`` to avoid spoofing. Upstream
MCP credentials and the deployment's proxy key header, including a custom
``litellm_key_header_name``, are dropped so they cannot reach a callback or a guardrail
through the derived metadata even when a caller omits ``general_settings``.
"""
from fastapi import Request
custom_key_header: Final = _custom_litellm_key_header_name()
excluded: Final = (
_SYNTHETIC_REQUEST_EXCLUDED_HEADERS
| _upstream_credential_headers(raw_headers.keys() if raw_headers else ())
| (frozenset({custom_key_header.lower()}) if custom_key_header else frozenset())
)
forwarded: Final = tuple(
(
name.lower().encode("latin-1", errors="replace"),
value.encode("utf-8", errors="replace"),
)
for name, value in (raw_headers.items() if raw_headers else ())
if name.lower() not in excluded
)
xff: Final = ((b"x-forwarded-for", client_ip.encode("utf-8")),) if client_ip else ()
return Request(
scope={
"type": "http",
"method": "POST",
"path": path,
"scheme": "http",
"server": _SYNTHETIC_REQUEST_SERVER,
"query_string": b"",
"root_path": "",
"headers": ((b"content-type", b"application/json"), *forwarded, *xff),
**({"client": (client_ip, 0)} if client_ip else {}),
}
)
def logging_safe_mcp_headers(raw_headers: Mapping[str, str] | None) -> Mapping[str, str]:
"""The MCP request's client headers, sanitized the way the chat completions path
sanitizes them before they reach a logging callback or a guardrail: proxy key
headers stripped, including the custom key header name the deployment configured,
upstream MCP credentials dropped, and credential-bearing values masked.
Client-controlled behaviour flags (``litellm-disable-message-redaction``) are dropped
too: these headers are read back out of the metadata to change proxy behaviour, so
leaving one in place would let any MCP client turn off the redaction an admin
configured. This path carries no key or team object to authorize an opt-out with, so
it always strips them."""
from starlette.datastructures import Headers
from litellm.proxy.litellm_pre_call_utils import (
UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS,
clean_headers,
redact_credential_headers,
)
excluded: Final = (
_upstream_credential_headers(raw_headers.keys() if raw_headers else ())
| UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS
)
cleaned: Final = clean_headers(
Headers(raw_headers),
litellm_key_header_name=_custom_litellm_key_header_name(),
)
return redact_credential_headers({name: value for name, value in cleaned.items() if name.lower() not in excluded})

View file

@ -274,7 +274,7 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS: Final = (
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
)
_UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS: Final = frozenset(
UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS: Final = frozenset(
{
"litellm-disable-message-redaction",
}
@ -355,7 +355,7 @@ def _strip_untrusted_request_header_controls(
return
for header_name in list(headers.keys()):
if isinstance(header_name, str) and header_name.lower() in _UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS:
if isinstance(header_name, str) and header_name.lower() in UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS:
if allow_client_message_redaction_opt_out:
continue
headers.pop(header_name, None)

View file

@ -679,6 +679,7 @@ class ProxyLogging:
# (e.g. MCPJWTSigner) to independently verify the caller's identity
# before re-signing an outbound token (FR-5 verify+re-sign).
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
"metadata": {"headers": kwargs.get("headers") or {}},
}
return synthetic_data

View file

@ -11,6 +11,7 @@ from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._experimental.mcp_server.utils import (
logging_safe_mcp_headers,
split_server_prefix_from_name,
strip_known_server_prefix,
)
@ -653,6 +654,7 @@ class LiteLLM_Proxy_MCP_Handler:
tool_results: Final[list[MCPToolResult]] = []
tool_call_id: str | None = None
rules_obj: Final = Rules()
logging_safe_headers: Final = logging_safe_mcp_headers(raw_headers)
for tool_call in tool_calls:
logging_request_data: dict[str, object] = {}
tool_name: str | None = None
@ -697,6 +699,7 @@ class LiteLLM_Proxy_MCP_Handler:
"tool_call_id": tool_call_id,
"tool_name": sanitized_tool_name,
"server_name": server_name,
"headers": logging_safe_headers,
}
logging_request_data = {
"model": f"MCP: {tool_name}",
@ -708,7 +711,7 @@ class LiteLLM_Proxy_MCP_Handler:
"proxy_server_request": {
"url": "/mcp/tools/call",
"method": "POST",
"headers": {},
"headers": logging_safe_headers,
"body": {
"name": sanitized_tool_name,
"arguments": parsed_arguments,

View file

@ -2874,7 +2874,7 @@ class TestMCPCustomHeaderName:
mock_general_settings.get.return_value = general_setting
# Call the method
result = MCPRequestHandler._get_mcp_client_side_auth_header_name()
result = MCPRequestHandler.get_mcp_client_side_auth_header_name()
# Assert the result
assert result == expected_header_name
@ -2938,7 +2938,7 @@ class TestMCPCustomHeaderName:
# Mock the header name method
with patch.object(
MCPRequestHandler,
"_get_mcp_client_side_auth_header_name",
"get_mcp_client_side_auth_header_name",
return_value=custom_header_name,
):
# Create headers from the test data
@ -2963,7 +2963,7 @@ class TestMCPCustomHeaderName:
# Mock the custom header name
with patch.object(
MCPRequestHandler,
"_get_mcp_client_side_auth_header_name",
"get_mcp_client_side_auth_header_name",
return_value="custom-auth-header",
):
# Create ASGI scope with custom header

View file

@ -1196,3 +1196,45 @@ class TestOpenApiResolvedUpstreamAuth:
)
assert resolved is None
lookup.assert_not_awaited()
class TestPreCallToolCheckExposesClientHeaders:
"""The pre_mcp_call guardrail payload must carry the caller's sanitized HTTP headers."""
@pytest.mark.asyncio
async def test_sanitized_client_headers_reach_the_guardrail_payload(self):
manager = MCPServerManager()
server = MCPServer(
server_id="test-id",
name="test_server",
server_name="test_server",
url="https://example.com",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
)
captured: Dict[str, Any] = {}
def capture(request_obj, kwargs):
captured.update(kwargs)
return {"model": "fake"}
proxy_logging = MagicMock(spec=ProxyLogging)
proxy_logging._create_mcp_request_object_from_kwargs = MagicMock(return_value=MagicMock())
proxy_logging._convert_mcp_to_llm_format = MagicMock(side_effect=capture)
proxy_logging.pre_call_hook = AsyncMock(return_value=None)
with patch.object(manager, "check_allowed_or_banned_tools", return_value=True):
with patch.object(manager, "check_tool_permission_for_key_team", new_callable=AsyncMock):
with patch.object(manager, "validate_allowed_params"):
await manager.pre_call_tool_check(
name="test_tool",
arguments={"key": "val"},
server_name="test_server",
user_api_key_auth=None,
proxy_logging_obj=proxy_logging,
server=server,
raw_headers={"x-nuid": "nuid-1", "x-litellm-api-key": "sk-proxy"},
)
assert captured["headers"] == {"x-nuid": "nuid-1"}

View file

@ -77,7 +77,7 @@ async def test_mcp_server_tool_call_body_contains_request_data():
# Mock the add_litellm_data_to_request function to capture the data
captured_data = {}
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config):
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config, **kwargs):
captured_data.update(data)
# Simulate the proxy_server_request creation
captured_data["proxy_server_request"] = {
@ -116,6 +116,107 @@ 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_forwards_client_headers_to_logging():
"""The MCP protocol path must hand the connection's client headers to the pre-call
pipeline, so logging callbacks and guardrails see them the way the REST path does."""
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-nuid": "nuid-1",
"x-app-id": "app-1",
"content-length": "42",
"x-forwarded-for": "9.9.9.9",
},
client_ip="1.2.3.4",
)
captured_headers = {}
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config, **kwargs):
captured_headers.update(request.headers)
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", {"param": "value"})
assert captured_headers.get("x-nuid") == "nuid-1"
assert captured_headers.get("x-app-id") == "app-1"
assert "content-length" not in captured_headers
assert captured_headers.get("x-forwarded-for") == "1.2.3.4"
@pytest.mark.asyncio
async def test_mcp_server_tool_call_strips_custom_litellm_key_header():
"""The deployment can rename the proxy key header via general_settings.litellm_key_header_name.
The pre-call pipeline only knows that name if it is passed in, so without it the virtual key
reaches metadata.headers and proxy_server_request.headers in plaintext."""
try:
from litellm.proxy._experimental.mcp_server.server import (
mcp_server_tool_call,
set_auth_context,
)
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
except ImportError:
pytest.skip("MCP server not available")
set_auth_context(
UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
raw_headers={"x-company-key": "sk-proxy-secret", "x-nuid": "nuid-1"},
client_ip="1.2.3.4",
)
captured_data = {}
async def capturing_add_litellm_data_to_request(**kwargs):
data = await add_litellm_data_to_request(**kwargs)
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",
capturing_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()):
with patch.dict(
"litellm.proxy.proxy_server.general_settings",
{"litellm_key_header_name": "x-company-key"},
clear=False,
):
await mcp_server_tool_call("test_tool", {"param": "value"})
metadata_headers = captured_data["metadata"]["headers"]
assert metadata_headers.get("x-nuid") == "nuid-1"
assert "x-company-key" not in metadata_headers
assert "x-company-key" not in captured_data["proxy_server_request"]["headers"]
@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
@ -133,7 +234,7 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror():
set_auth_context(UserAPIKeyAuth(api_key="test_key", user_id="test_user"))
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config):
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config, **kwargs):
return data
async def mock_call_mcp_tool(*args, **kwargs):
@ -1245,7 +1346,7 @@ async def test_mcp_server_tool_call_body_with_none_arguments():
# Mock the add_litellm_data_to_request function to capture the data
captured_data = {}
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config):
async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config, **kwargs):
captured_data.update(data)
captured_data["proxy_server_request"] = {
"url": str(request.url),

View file

@ -1,7 +1,11 @@
from unittest.mock import patch
import pytest
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.utils import (
build_synthetic_mcp_request,
logging_safe_mcp_headers,
validate_and_normalize_mcp_server_payload,
validate_tool_display_names,
)
@ -47,3 +51,99 @@ class TestValidateAndNormalizeMcpServerPayload:
tool_name_to_display_name={"read_wiki_structure": "browse_repo_docs"},
)
validate_and_normalize_mcp_server_payload(payload)
class TestLoggingSafeMcpHeaders:
def test_returns_empty_for_missing_headers(self):
assert logging_safe_mcp_headers(None) == {}
assert logging_safe_mcp_headers({}) == {}
def test_exposes_custom_headers_and_masks_credentials(self):
safe = logging_safe_mcp_headers(
{
"x-nuid": "nuid-1",
"x-app-id": "app-1",
"x-litellm-api-key": "sk-proxy",
"cookie": "session=secret",
}
)
assert safe == {
"x-nuid": "nuid-1",
"x-app-id": "app-1",
"cookie": "***REDACTED***",
}
def test_strips_custom_litellm_key_header(self):
"""general_settings.litellm_key_header_name carries the proxy virtual key, so it must
never reach a callback or a guardrail even though clean_headers cannot know its name."""
with patch.dict(
"litellm.proxy.proxy_server.general_settings",
{"litellm_key_header_name": "x-company-key"},
clear=False,
):
safe = logging_safe_mcp_headers({"x-company-key": "sk-proxy", "x-nuid": "nuid-1"})
assert safe == {"x-nuid": "nuid-1"}
def test_strips_client_controlled_redaction_opt_out(self):
"""litellm-disable-message-redaction is read back out of the logged metadata to turn off
redaction, so leaving it in place lets any MCP client undo what an admin configured."""
safe = logging_safe_mcp_headers({"litellm-disable-message-redaction": "true", "x-nuid": "nuid-1"})
assert safe == {"x-nuid": "nuid-1"}
def test_strips_upstream_mcp_credentials(self):
safe = logging_safe_mcp_headers(
{
"x-mcp-auth": "Bearer upstream",
"x-mcp-github-authorization": "Bearer gh_token",
"x-mcp-zapier-x-api-key": "zapier-key",
"x-nuid": "nuid-1",
}
)
assert safe == {"x-nuid": "nuid-1"}
def test_strips_custom_mcp_client_side_auth_header(self):
with patch.dict(
"litellm.proxy.proxy_server.general_settings",
{"mcp_client_side_auth_header_name": "x-upstream-token"},
clear=False,
):
safe = logging_safe_mcp_headers({"x-upstream-token": "Bearer upstream", "x-nuid": "nuid-1"})
assert safe == {"x-nuid": "nuid-1"}
class TestBuildSyntheticMcpRequest:
def test_forwards_client_headers_without_upstream_credentials(self):
"""The synthetic request feeds add_litellm_data_to_request, which derives
metadata.headers, so upstream MCP credentials must not ride along."""
request = build_synthetic_mcp_request(
path="/mcp/tools/call",
raw_headers={
"x-nuid": "nuid-1",
"x-mcp-auth": "Bearer upstream",
"x-mcp-github-authorization": "Bearer gh_token",
},
)
assert request.headers.get("x-nuid") == "nuid-1"
assert "x-mcp-auth" not in request.headers
assert "x-mcp-github-authorization" not in request.headers
def test_drops_custom_litellm_key_header(self):
"""Callers such as the sampling flow build metadata off this request, so the
deployment's custom proxy key header must never be forwarded on it."""
with patch.dict(
"litellm.proxy.proxy_server.general_settings",
{"litellm_key_header_name": "x-company-key"},
clear=False,
):
request = build_synthetic_mcp_request(
path="/mcp/sampling/createMessage",
raw_headers={"x-company-key": "sk-proxy-secret", "x-nuid": "nuid-1"},
)
assert request.headers.get("x-nuid") == "nuid-1"
assert "x-company-key" not in request.headers

View file

@ -76,6 +76,23 @@ def test_convert_mcp_to_llm_format_defaults_model(proxy_logging, make_mcp_reques
}
def test_convert_mcp_to_llm_format_exposes_headers_on_metadata(proxy_logging, make_mcp_request_obj):
"""Guardrails read the caller's HTTP headers off ``metadata.headers`` on the chat
completions path, so the MCP bridge has to put them in the same place."""
req = make_mcp_request_obj()
out = proxy_logging._convert_mcp_to_llm_format(
request_obj=req,
kwargs={"headers": {"x-nuid": "nuid-1"}},
)
assert out["metadata"]["headers"] == {"x-nuid": "nuid-1"}
def test_convert_mcp_to_llm_format_defaults_headers_to_empty(proxy_logging, make_mcp_request_obj):
req = make_mcp_request_obj()
out = proxy_logging._convert_mcp_to_llm_format(request_obj=req, kwargs={})
assert out["metadata"]["headers"] == {}
def test_convert_mcp_to_llm_format_missing_request_obj_raises(proxy_logging):
with pytest.raises(AttributeError):
proxy_logging._convert_mcp_to_llm_format(request_obj=None, kwargs={})

View file

@ -536,6 +536,37 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
assert mock_get_tools.await_args.kwargs["request_tags"] == ["team-a"]
@pytest.mark.asyncio
async def test_execute_tool_calls_exposes_sanitized_client_headers_to_logging(monkeypatch):
"""The Responses API MCP bridge used to log an empty header dict, hiding the caller's
headers from logging callbacks and hooks."""
_setup_proxy_logging(monkeypatch)
_setup_mcp_call_environment(monkeypatch)
captured = {}
def fake_function_setup(*_args, **kwargs):
captured.update(kwargs)
return None, None
handler_module = importlib.import_module(
"litellm.responses.mcp.litellm_proxy_mcp_handler"
)
monkeypatch.setattr(handler_module, "function_setup", fake_function_setup)
tool_name = "deepwiki-read_wiki_structure"
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map={tool_name: "deepwiki"},
tool_calls=[{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}],
user_api_key_auth=None,
raw_headers={"x-nuid": "nuid-1", "x-litellm-api-key": "sk-proxy", "cookie": "s=1"},
)
expected = {"x-nuid": "nuid-1", "cookie": "***REDACTED***"}
assert captured["metadata"]["headers"] == expected
assert captured["proxy_server_request"]["headers"] == expected
@pytest.mark.asyncio
async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monkeypatch):
_setup_proxy_logging(monkeypatch)

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22943
"limit": 22941
},
"LIT002": {
"limit": 27141
"limit": 27139
},
"LIT003": {
"limit": 269
@ -27,7 +27,7 @@
"limit": 0
},
"LIT010": {
"limit": 16722
"limit": 16716
},
"LIT011": {
"limit": 5596