From 4e91e69cf6ed53af74056c9d1d8971f76b648dc5 Mon Sep 17 00:00:00 2001 From: Genmin Date: Fri, 1 May 2026 07:56:40 -0700 Subject: [PATCH] refactor: reuse MCP OpenAPI header builder --- .../mcp_server/mcp_server_manager.py | 37 +++++----------- .../proxy/_experimental/mcp_server/server.py | 27 ++++-------- .../proxy/_experimental/mcp_server/utils.py | 43 +++++++++++++++++++ .../mcp_server/test_mcp_hook_extra_headers.py | 22 ++++++++++ 4 files changed, 85 insertions(+), 44 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 62cbd70eda0..104d69ced67 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -51,6 +51,7 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mc from litellm.proxy._experimental.mcp_server.utils import ( MCP_TOOL_PREFIX_SEPARATOR, add_server_prefix_to_name, + build_mcp_runtime_extra_headers, compute_short_server_prefix, get_server_prefix, is_short_mcp_tool_prefix_enabled, @@ -2522,34 +2523,16 @@ class MCPServerManager: server_auth_header: Optional[Union[Dict[str, str], str]], ) -> Optional[Dict[str, str]]: """Build outbound headers shared by MCP transports and OpenAPI handlers.""" - extra_headers: Dict[str, str] = {} - - if oauth2_headers: - extra_headers.update(oauth2_headers) - - if mcp_server.extra_headers and raw_headers: - normalized_raw_headers = { - str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) - } - for header in mcp_server.extra_headers: - if not isinstance(header, str): - continue - if ( - mcp_server.has_client_credentials - and header.lower() == "authorization" - ): - continue - header_value = normalized_raw_headers.get(header.lower()) - if header_value is None: - continue - extra_headers[header] = header_value - - if include_static_headers and mcp_server.static_headers: - extra_headers.update(mcp_server.static_headers) + extra_headers = build_mcp_runtime_extra_headers( + server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + include_static_headers=include_static_headers, + ) if hook_extra_headers: if "Authorization" in hook_extra_headers: - if "Authorization" in extra_headers: + if extra_headers and "Authorization" in extra_headers: verbose_logger.warning( "MCPServerManager: hook_extra_headers 'Authorization' will overwrite " "the existing Authorization header from static_headers or oauth2_headers. " @@ -2568,9 +2551,11 @@ class MCPServerManager: "the hook JWT to be the sole credential.", mcp_server.server_name or mcp_server.name, ) + if extra_headers is None: + extra_headers = {} extra_headers.update(hook_extra_headers) - return extra_headers or None + return extra_headers async def _call_regular_mcp_tool( # noqa: PLR0915 self, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e3a700fe67f..979a21f8fdc 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -48,6 +48,7 @@ from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_NAME, LITELLM_MCP_SERVER_VERSION, add_server_prefix_to_name, + build_mcp_runtime_extra_headers, get_server_prefix, iter_known_server_prefixes, ) @@ -1157,25 +1158,15 @@ if MCP_AVAILABLE: raw_headers: Optional[Dict[str, str]], ) -> Optional[Dict[str, str]]: """Build per-request headers for a local OpenAPI-generated MCP tool.""" - extra_headers = oauth2_headers.copy() if oauth2_headers else None + if server is None: + return oauth2_headers.copy() if oauth2_headers else None - if server and server.extra_headers and raw_headers: - if extra_headers is None: - extra_headers = {} - - normalized_raw_headers = { - str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) - } - - for header in server.extra_headers: - if not isinstance(header, str): - continue - header_value = normalized_raw_headers.get(header.lower()) - if header_value is None: - continue - extra_headers[header] = header_value - - return extra_headers + return build_mcp_runtime_extra_headers( + server=server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + include_static_headers=False, + ) def _merge_gateway_initialize_instructions( allowed_mcp_servers: List[MCPServer], diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index df5705c3425..9ad24694169 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -320,3 +320,46 @@ def merge_mcp_headers( merged.update({str(k): str(v) for k, v in static_headers.items()}) return merged or None + + +def build_mcp_runtime_extra_headers( + *, + server: Any, + oauth2_headers: Optional[Mapping[str, str]] = None, + raw_headers: Optional[Mapping[Any, str]] = None, + hook_extra_headers: Optional[Mapping[str, str]] = None, + include_static_headers: bool = False, +) -> Optional[Dict[str, str]]: + """Build per-request outbound headers for MCP and OpenAPI MCP calls.""" + extra_headers: Dict[str, str] = {} + + if oauth2_headers: + extra_headers.update({str(k): str(v) for k, v in oauth2_headers.items()}) + + configured_headers = getattr(server, "extra_headers", None) + if configured_headers and raw_headers: + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } + for header in configured_headers: + if not isinstance(header, str): + continue + if ( + getattr(server, "has_client_credentials", False) + and header.lower() == "authorization" + ): + continue + header_value = normalized_raw_headers.get(header.lower()) + if header_value is None: + continue + extra_headers[header] = str(header_value) + + if include_static_headers: + static_headers = getattr(server, "static_headers", None) + if static_headers: + extra_headers.update({str(k): str(v) for k, v in static_headers.items()}) + + if hook_extra_headers: + extra_headers.update({str(k): str(v) for k, v in hook_extra_headers.items()}) + + return extra_headers or None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 0166aeff4be..4bc633c41b3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -814,6 +814,28 @@ class TestHookHeaderMergePriority: assert headers == {"X-TOKEN": "request-token"} + def test_gateway_openapi_request_headers_match_manager_builder(self): + """Gateway OpenAPI execution uses the shared runtime header builder.""" + manager = MCPServerManager() + server = self._make_server(extra_headers_config=["X-TOKEN", "X-Trace"]) + + manager_headers = manager._build_openapi_request_extra_headers( + mcp_server=server, + oauth2_headers={"Authorization": "Bearer oauth-token"}, + raw_headers={"x-token": "request-token", "x-trace": "trace-from-request"}, + hook_extra_headers=None, + ) + gateway_headers = mcp_server_module._get_request_extra_headers_for_openapi_tool( + server=server, + oauth2_headers={"Authorization": "Bearer oauth-token"}, + raw_headers={ + "x-token": "request-token", + "x-trace": "trace-from-request", + }, + ) + + assert gateway_headers == manager_headers + class TestOpenAPIRequestContext: """Tests for request-scoped headers on OpenAPI-generated MCP tools."""