mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(mcp): static headers win over forwarded headers in OpenAPI MCP
Match the existing MCP invariant in merge_mcp_headers and the managed MCP path: operator-configured static headers always override caller-forwarded headers on name conflict, with case-insensitive comparison so different casing cannot bypass the precedence. _request_auth_header (BYOK) still overrides Authorization last. Addresses Veria review on PR #27383. Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
0f5aad5718
commit
3d409e7395
2 changed files with 89 additions and 6 deletions
|
|
@ -307,14 +307,40 @@ def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]:
|
|||
def _merge_openapi_tool_request_headers(
|
||||
static_headers: Dict[str, str]
|
||||
) -> Dict[str, str]:
|
||||
"""Merge static closure headers with per-request ContextVar overrides."""
|
||||
effective_headers = dict(static_headers)
|
||||
request_extra = _request_extra_headers.get()
|
||||
if request_extra:
|
||||
effective_headers.update(request_extra)
|
||||
"""Merge static closure headers with per-request ContextVar overrides.
|
||||
|
||||
Precedence (highest to lowest):
|
||||
1. ``_request_auth_header`` — BYOK override of ``Authorization``
|
||||
2. ``static_headers`` — operator-configured headers baked into the
|
||||
tool closure at registration time
|
||||
3. ``_request_extra_headers`` — per-request headers forwarded from
|
||||
the MCP caller (allowlisted by ``MCPServer.extra_headers``)
|
||||
|
||||
This matches the existing MCP invariant in
|
||||
:func:`litellm.proxy._experimental.mcp_server.utils.merge_mcp_headers`
|
||||
and the managed MCP path, where ``static_headers`` always wins over
|
||||
caller-forwarded headers. Keeping the same precedence here prevents an
|
||||
authenticated caller from overriding an operator-configured value
|
||||
(e.g. a tenant id or upstream API key) by sending the same header name.
|
||||
|
||||
Header names are compared case-insensitively so different casing cannot
|
||||
bypass the precedence rules.
|
||||
"""
|
||||
request_extra = _request_extra_headers.get() or {}
|
||||
static = static_headers or {}
|
||||
|
||||
static_lower_names = {k.lower() for k in static}
|
||||
effective_headers: Dict[str, str] = {
|
||||
k: v for k, v in request_extra.items() if k.lower() not in static_lower_names
|
||||
}
|
||||
effective_headers.update(static)
|
||||
|
||||
override_auth = _request_auth_header.get()
|
||||
if override_auth:
|
||||
for existing in [k for k in effective_headers if k.lower() == "authorization"]:
|
||||
del effective_headers[existing]
|
||||
effective_headers["Authorization"] = override_auth
|
||||
|
||||
return effective_headers
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1070,7 +1070,7 @@ class TestRequestExtraHeaders:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extra_headers_merged_with_static_headers(self):
|
||||
"""Request extra headers are merged on top of static (baked-in) headers."""
|
||||
"""Forwarded headers are passed through alongside non-conflicting static headers."""
|
||||
operation = {}
|
||||
func = create_tool_function(
|
||||
path="/data",
|
||||
|
|
@ -1096,6 +1096,63 @@ class TestRequestExtraHeaders:
|
|||
assert headers_sent.get("X-Static") == "static-value"
|
||||
assert headers_sent.get("X-TOKEN") == "dynamic-value"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_static_headers_win_over_forwarded_on_conflict(self):
|
||||
"""Static (operator) headers must override forwarded (caller) headers on name conflict."""
|
||||
operation = {}
|
||||
func = create_tool_function(
|
||||
path="/data",
|
||||
method="get",
|
||||
operation=operation,
|
||||
base_url="https://api.example.com",
|
||||
headers={"X-Tenant": "operator-tenant"},
|
||||
)
|
||||
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "ok")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
token = _request_extra_headers.set({"X-Tenant": "caller-spoofed"})
|
||||
try:
|
||||
result = await func()
|
||||
finally:
|
||||
_request_extra_headers.reset(token)
|
||||
|
||||
assert result == "ok"
|
||||
call_args = async_client.get.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent.get("X-Tenant") == "operator-tenant"
|
||||
assert "caller-spoofed" not in headers_sent.values()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_static_headers_win_case_insensitively(self):
|
||||
"""Forwarded header with different casing must not bypass the static-wins rule."""
|
||||
operation = {}
|
||||
func = create_tool_function(
|
||||
path="/data",
|
||||
method="get",
|
||||
operation=operation,
|
||||
base_url="https://api.example.com",
|
||||
headers={"X-Tenant": "operator-tenant"},
|
||||
)
|
||||
|
||||
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
|
||||
async_client = _create_mock_client("get", "ok")
|
||||
mock_client.return_value = async_client
|
||||
|
||||
token = _request_extra_headers.set({"x-tenant": "caller-spoofed"})
|
||||
try:
|
||||
result = await func()
|
||||
finally:
|
||||
_request_extra_headers.reset(token)
|
||||
|
||||
assert result == "ok"
|
||||
call_args = async_client.get.call_args
|
||||
headers_sent = call_args[1]["headers"]
|
||||
assert headers_sent.get("X-Tenant") == "operator-tenant"
|
||||
assert "x-tenant" not in headers_sent
|
||||
assert "caller-spoofed" not in headers_sent.values()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_header_still_overrides_extra_headers(self):
|
||||
"""_request_auth_header takes precedence for Authorization over extra headers."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue