From 3d409e7395bcc3029bdfe5e87ce4955b366fe3a3 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Sat, 9 May 2026 18:40:15 +0000 Subject: [PATCH] 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 --- .../mcp_server/openapi_to_mcp_generator.py | 36 +++++++++-- .../test_openapi_to_mcp_generator.py | 59 ++++++++++++++++++- 2 files changed, 89 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 973f6ee59df..271517bb1e6 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index b91501b8d41..39f3c767220 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -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."""