mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
efbdb6901a
commit
59eeae374c
15 changed files with 483 additions and 124 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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={})
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue