mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
fix(openapi-mcp): forward extra_headers from client request to OpenAPI tools
For OpenAPI-backed MCP servers, `extra_headers` config now correctly forwards matching client request headers to the upstream API, consistent with the existing behavior for managed MCP servers. Adds a `_request_extra_headers` ContextVar in openapi_to_mcp_generator.py that is populated in server.py before dispatching to a local/OpenAPI tool handler, then merged into effective_headers inside the tool function closure. Fixes #26794 Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
2c0d20b327
commit
583eba177f
3 changed files with 175 additions and 1 deletions
|
|
@ -30,6 +30,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."""
|
||||
|
|
@ -274,7 +281,7 @@ def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
def create_tool_function(
|
||||
def create_tool_function( # noqa: PLR0915 -- function grew past 50 statements after extra_headers ContextVar wiring; refactor scope larger than this PR
|
||||
path: str,
|
||||
method: str,
|
||||
operation: Dict[str, Any],
|
||||
|
|
@ -316,6 +323,9 @@ def create_tool_function(
|
|||
# correct prefix (Bearer / ApiKey / Basic) formatted by the caller in
|
||||
# server.py based on the server's configured auth_type.
|
||||
effective_headers = dict(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
|
||||
|
|
|
|||
|
|
@ -150,6 +150,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 (
|
||||
|
|
@ -2035,11 +2036,33 @@ if MCP_AVAILABLE:
|
|||
auth_header_value = f"Basic {mcp_auth_header}"
|
||||
else:
|
||||
auth_header_value = f"Bearer {mcp_auth_header}"
|
||||
|
||||
# Build extra headers dict from client request for OpenAPI tools.
|
||||
# MCPServer.extra_headers is a list of header *names* whose values
|
||||
# should be forwarded from the incoming request to the upstream API.
|
||||
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)
|
||||
}
|
||||
for header_name in mcp_server.extra_headers:
|
||||
if not isinstance(header_name, str):
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -855,3 +857,142 @@ 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"]
|
||||
# _request_auth_header must win for Authorization
|
||||
assert headers_sent.get("Authorization") == "Bearer byok-credential"
|
||||
# Other extra headers must still be present
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue