refactor: reuse MCP OpenAPI header builder

This commit is contained in:
Genmin 2026-05-01 07:56:40 -07:00
parent ddcd071a9a
commit 4e91e69cf6
4 changed files with 85 additions and 44 deletions

View file

@ -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,

View file

@ -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],

View file

@ -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

View file

@ -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."""