fix: forward extra HEADERS FOR oPENapi mcp TOOLS

This commit is contained in:
volcano303 2026-04-29 23:31:02 +02:00
parent 602a6cff81
commit e75da48a1b
4 changed files with 143 additions and 4 deletions

View file

@ -30,6 +30,23 @@ HEADERS: Dict[str, str] = {}
_request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.ContextVar(
"_request_auth_header", default=None
)
_request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = (
contextvars.ContextVar("_request_extra_headers", default=None)
)
def _build_effective_headers(headers: Dict[str, str]) -> Dict[str, str]:
"""Merge static OpenAPI headers with request-scoped MCP header overrides."""
effective_headers = dict(headers)
request_extra_headers = _request_extra_headers.get()
if request_extra_headers:
effective_headers.update(request_extra_headers)
override_auth = _request_auth_header.get()
if override_auth:
effective_headers["Authorization"] = override_auth
return effective_headers
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
@ -314,10 +331,7 @@ def create_tool_function(
# 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 = _build_effective_headers(headers)
# Build URL from base_url and path
url = base_url + path

View file

@ -157,6 +157,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 (
@ -1124,6 +1125,29 @@ if MCP_AVAILABLE:
return server_auth_header, extra_headers
def _prepare_local_mcp_request_extra_headers(
server: Optional[MCPServer],
raw_headers: Optional[Dict[str, str]],
) -> Optional[Dict[str, str]]:
"""Build request-time extra headers for local OpenAPI-generated tools."""
if not server or not server.extra_headers or not raw_headers:
return None
normalized_raw_headers = {
str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)
}
extra_headers: Dict[str, 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 not None:
extra_headers[header] = header_value
return extra_headers or None
def _merge_gateway_initialize_instructions(
allowed_mcp_servers: List[MCPServer],
) -> Optional[str]:
@ -2120,11 +2144,17 @@ if MCP_AVAILABLE:
auth_header_value = f"Basic {mcp_auth_header}"
else:
auth_header_value = f"Bearer {mcp_auth_header}"
request_extra_headers = _prepare_local_mcp_request_extra_headers(
server=mcp_server,
raw_headers=raw_headers,
)
_auth_token = _request_auth_header.set(auth_header_value)
_extra_headers_token = _request_extra_headers.set(request_extra_headers)
try:
local_content = await _handle_local_mcp_tool(name, arguments)
finally:
_request_auth_header.reset(_auth_token)
_request_extra_headers.reset(_extra_headers_token)
response = CallToolResult(content=cast(Any, local_content), isError=False)
# Try managed MCP server tool (pass the full prefixed name)

View file

@ -135,6 +135,36 @@ def test_prepare_mcp_server_headers_case_insensitive_extra_headers():
assert extra_headers == {"Authorization": "Bearer token"}
def test_prepare_local_mcp_request_extra_headers_case_insensitive():
try:
from litellm.proxy._experimental.mcp_server.server import (
_prepare_local_mcp_request_extra_headers,
)
except ImportError:
pytest.skip("MCP server not available")
server = MCPServer(
server_id="server-case",
name="server",
transport=MCPTransport.http,
extra_headers=["X-TOKEN", "X-API-Key"],
)
extra_headers = _prepare_local_mcp_request_extra_headers(
server=server,
raw_headers={
"x-token": "request-token",
"X-API-KEY": "request-api-key",
"x-litellm-api-key": "litellm-key",
},
)
assert extra_headers == {
"X-TOKEN": "request-token",
"X-API-Key": "request-api-key",
}
@pytest.mark.asyncio
async def test_get_prompts_from_mcp_servers_success():
try:

View file

@ -17,6 +17,8 @@ import pytest
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_resolve_param_list,
_resolve_ref,
_request_auth_header,
_request_extra_headers,
build_input_schema,
create_tool_function,
extract_parameters,
@ -369,6 +371,69 @@ class TestCreateToolFunction:
# Should have no exec() calls
assert len(exec_calls) == 0, "create_tool_function should not use exec()"
@pytest.mark.asyncio
async def test_request_extra_headers_are_forwarded(self):
"""Test request-time extra headers are forwarded to OpenAPI requests."""
func = create_tool_function(
path="/protected",
method="get",
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
token = _request_extra_headers.set({"X-TOKEN": "request-token"})
try:
result = await func()
finally:
_request_extra_headers.reset(token)
assert result == "ok"
call_args = async_client.get.call_args
assert call_args.kwargs["headers"] == {
"X-Static": "static-value",
"X-TOKEN": "request-token",
}
@pytest.mark.asyncio
async def test_auth_override_wins_over_request_extra_headers(self):
"""Test x-mcp-auth Authorization override preserves existing precedence."""
func = create_tool_function(
path="/protected",
method="get",
operation={},
base_url="https://api.example.com",
headers={"Authorization": "Bearer static-token"},
)
with patch(GET_ASYNC_CLIENT_TARGET) as mock_client:
async_client = _create_mock_client("get", "ok")
mock_client.return_value = async_client
extra_token = _request_extra_headers.set(
{
"Authorization": "Bearer request-token",
"X-TOKEN": "request-token",
}
)
auth_token = _request_auth_header.set("Bearer mcp-auth-token")
try:
result = await func()
finally:
_request_auth_header.reset(auth_token)
_request_extra_headers.reset(extra_token)
assert result == "ok"
call_args = async_client.get.call_args
assert call_args.kwargs["headers"] == {
"Authorization": "Bearer mcp-auth-token",
"X-TOKEN": "request-token",
}
class TestBuildInputSchema:
"""Test that build_input_schema preserves original parameter names."""