mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
refactor: reuse MCP OpenAPI header builder
This commit is contained in:
parent
ddcd071a9a
commit
4e91e69cf6
4 changed files with 85 additions and 44 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue