From 38bc921dd2ba47e61b00feeb11462dde62b4ce3d Mon Sep 17 00:00:00 2001 From: Milan Date: Thu, 7 May 2026 13:46:48 +0300 Subject: [PATCH 1/3] fix(mcp): forward extra_headers for OpenAPI MCP tools OpenAPI-generated tools only applied static closure headers and BYOK Authorization via ContextVar. Copy MCPServer.extra_headers from the incoming MCP request into _request_extra_headers (set in server.py before local tool dispatch), merge in openapi_to_mcp_generator via a small helper. OAuth2 M2M: do not forward caller Authorization from raw_headers (same rule as _prepare_mcp_server_headers for managed MCP). Adds TestRequestExtraHeaders and clarifies mcp_server_manager registration comment. Fixes #26794 Co-authored-by: Cursor --- .../mcp_server/mcp_server_manager.py | 3 +- .../mcp_server/openapi_to_mcp_generator.py | 30 +++- .../proxy/_experimental/mcp_server/server.py | 32 ++++ .../test_openapi_to_mcp_generator.py | 139 ++++++++++++++++++ 4 files changed, 195 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 9923c3ce4bf..50dd7efda13 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -497,7 +497,8 @@ class MCPServerManager: # Add any static headers from server config. # # Note: `extra_headers` on MCPServer is a List[str] of header names to forward - # from the client request (not available in this OpenAPI tool generation step). + # from each client MCP request; values are applied at call time via + # `_request_extra_headers` in server.py (not baked in here). # `static_headers` is a dict of concrete headers to always send. headers = ( merge_mcp_headers( 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 3b2fa097b70..a650c7e40d9 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -31,6 +31,13 @@ _request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.Contex "_request_auth_header", default=None ) +# Per-request extra headers forwarded from the client request. +# Populated from MCPServer.extra_headers names matched against raw request +# headers in server.py before dispatching to a local/OpenAPI tool handler. +_request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = ( + contextvars.ContextVar("_request_extra_headers", default=None) +) + def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str: """Ensure path params cannot introduce directory traversal.""" @@ -273,6 +280,20 @@ 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) + override_auth = _request_auth_header.get() + if override_auth: + effective_headers["Authorization"] = override_auth + return effective_headers + + def create_tool_function( path: str, method: str, @@ -310,14 +331,7 @@ def create_tool_function( The function safely handles parameter names that aren't valid Python identifiers by using **kwargs instead of named parameters. """ - # Allow per-request auth override (e.g. BYOK credential set via ContextVar). - # The ContextVar holds the full Authorization header value, including the - # correct prefix (Bearer / ApiKey / Basic) formatted by the caller in - # server.py based on the server's configured auth_type. - effective_headers = dict(headers) - override_auth = _request_auth_header.get() - if override_auth: - effective_headers["Authorization"] = override_auth + effective_headers = _merge_openapi_tool_request_headers(headers) # Build URL from base_url and path url = base_url + path diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 54d9bbe6e28..b0dc3eff189 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -158,6 +158,7 @@ if MCP_AVAILABLE: ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_auth_header, + _request_extra_headers, ) from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport from litellm.proxy._experimental.mcp_server.tool_registry import ( @@ -2195,11 +2196,42 @@ if MCP_AVAILABLE: auth_header_value = f"Basic {mcp_auth_header}" else: auth_header_value = f"Bearer {mcp_auth_header}" + + # Forward named client headers to OpenAPI tool upstream requests. + # MCPServer.extra_headers lists header names to copy from raw_headers. + # OAuth2 M2M: never take Authorization from the caller (matches + # _prepare_mcp_server_headers for managed MCP). + forwarded_headers: Optional[Dict[str, str]] = None + if mcp_server and mcp_server.extra_headers and raw_headers: + normalized_raw = { + str(k).lower(): v + for k, v in raw_headers.items() + if isinstance(k, str) + } + skip_caller_authorization = bool( + getattr(mcp_server, "has_client_credentials", False) + ) + for header_name in mcp_server.extra_headers: + if not isinstance(header_name, str): + continue + if ( + skip_caller_authorization + and header_name.lower() == "authorization" + ): + continue + value = normalized_raw.get(header_name.lower()) + if value is not None: + if forwarded_headers is None: + forwarded_headers = {} + forwarded_headers[header_name] = value + _auth_token = _request_auth_header.set(auth_header_value) + _extra_token = _request_extra_headers.set(forwarded_headers) try: local_content = await _handle_local_mcp_tool(name, arguments) finally: _request_auth_header.reset(_auth_token) + _request_extra_headers.reset(_extra_token) response = CallToolResult(content=cast(Any, local_content), isError=False) # Try managed MCP server tool (pass the full prefixed name) 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 efe100a11dc..4ff810bcd61 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 @@ -15,6 +15,8 @@ from unittest.mock import AsyncMock, patch import pytest from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_auth_header, + _request_extra_headers, _resolve_param_list, _resolve_ref, build_input_schema, @@ -868,3 +870,140 @@ class TestResolveOperationParams: assert "per_page" in names assert "sha" in names assert len(names) == 4 # no duplicates + + +class TestRequestExtraHeaders: + """Tests for _request_extra_headers ContextVar forwarding in tool_function.""" + + @pytest.mark.asyncio + async def test_extra_headers_forwarded_to_upstream(self): + """Extra headers set via ContextVar are included in the upstream request.""" + operation = {} + func = create_tool_function( + path="/data", + method="get", + operation=operation, + base_url="https://api.example.com", + ) + + 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-TOKEN": "secret-value"}) + 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-TOKEN") == "secret-value" + + @pytest.mark.asyncio + async def test_no_extra_headers_by_default(self): + """Without setting _request_extra_headers, no extra headers are injected.""" + operation = {} + func = create_tool_function( + path="/data", + method="get", + operation=operation, + base_url="https://api.example.com", + headers={"X-Static": "static-value"}, + ) + + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("get", "ok") + mock_client.return_value = async_client + + result = await func() + + assert result == "ok" + call_args = async_client.get.call_args + headers_sent = call_args[1]["headers"] + assert headers_sent == {"X-Static": "static-value"} + assert "X-TOKEN" not in headers_sent + + @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.""" + operation = {} + func = create_tool_function( + path="/data", + method="post", + operation=operation, + base_url="https://api.example.com", + headers={"X-Static": "static-value"}, + ) + + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("post", "created") + mock_client.return_value = async_client + + token = _request_extra_headers.set({"X-TOKEN": "dynamic-value"}) + try: + result = await func() + finally: + _request_extra_headers.reset(token) + + assert result == "created" + call_args = async_client.post.call_args + headers_sent = call_args[1]["headers"] + assert headers_sent.get("X-Static") == "static-value" + assert headers_sent.get("X-TOKEN") == "dynamic-value" + + @pytest.mark.asyncio + async def test_auth_header_still_overrides_extra_headers(self): + """_request_auth_header takes precedence for Authorization over extra headers.""" + operation = {} + func = create_tool_function( + path="/secure", + method="get", + operation=operation, + base_url="https://api.example.com", + ) + + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("get", "secure-data") + mock_client.return_value = async_client + + extra_token = _request_extra_headers.set( + {"Authorization": "Bearer extra", "X-TOKEN": "token-value"} + ) + auth_token = _request_auth_header.set("Bearer byok-credential") + try: + result = await func() + finally: + _request_auth_header.reset(auth_token) + _request_extra_headers.reset(extra_token) + + assert result == "secure-data" + call_args = async_client.get.call_args + headers_sent = call_args[1]["headers"] + assert headers_sent.get("Authorization") == "Bearer byok-credential" + assert headers_sent.get("X-TOKEN") == "token-value" + + @pytest.mark.asyncio + async def test_extra_headers_not_leaked_between_calls(self): + """After resetting the ContextVar, subsequent calls do not see the headers.""" + operation = {} + func = create_tool_function( + path="/data", + method="get", + operation=operation, + base_url="https://api.example.com", + ) + + 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-TOKEN": "first-call"}) + _request_extra_headers.reset(token) + + await func() + + call_args = async_client.get.call_args + headers_sent = call_args[1]["headers"] + assert "X-TOKEN" not in headers_sent From c31ead87d7d3922f8dd62543007d1375eb5a4393 Mon Sep 17 00:00:00 2001 From: Milan Date: Thu, 7 May 2026 14:17:05 +0300 Subject: [PATCH 2/3] refactor(mcp): access has_client_credentials on MCPServer directly Greptile: getattr default was redundant; property exists on MCPServer and mcp_server is non-None inside the extra_headers forwarding block. Co-authored-by: Cursor --- litellm/proxy/_experimental/mcp_server/server.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index b0dc3eff189..276a6e8a3bb 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2208,9 +2208,7 @@ if MCP_AVAILABLE: for k, v in raw_headers.items() if isinstance(k, str) } - skip_caller_authorization = bool( - getattr(mcp_server, "has_client_credentials", False) - ) + skip_caller_authorization = bool(mcp_server.has_client_credentials) for header_name in mcp_server.extra_headers: if not isinstance(header_name, str): continue From a1d1906025c21fe706694920e6364cecba7181ce Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Sat, 9 May 2026 18:40:15 +0000 Subject: [PATCH 3/3] 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 a650c7e40d9..3ed8835a410 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -283,14 +283,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 4ff810bcd61..59818e7b7cc 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 @@ -927,7 +927,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", @@ -953,6 +953,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."""