mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix: forward extra HEADERS FOR oPENapi mcp TOOLS
This commit is contained in:
parent
602a6cff81
commit
e75da48a1b
4 changed files with 143 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue