test: cover openapi mcp request context

This commit is contained in:
Genmin 2026-04-30 07:23:30 -07:00
parent 67fbb33689
commit 0aaae31f26

View file

@ -12,12 +12,21 @@ Validates that:
"""
import asyncio
from datetime import datetime
from typing import Any, Dict, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._experimental.mcp_server import server as mcp_server_module
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
)
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.mcp import MCPAuth, MCPTransport
@ -790,6 +799,144 @@ class TestHookHeaderMergePriority:
"X-Trace": "trace-from-request",
}
def test_openapi_request_headers_forwards_raw_headers_without_oauth(self):
"""Configured OpenAPI extra_headers are copied from raw request headers."""
manager = MCPServerManager()
server = self._make_server(extra_headers_config=["X-TOKEN", "X-Missing"])
server.__dict__["extra_headers"] = ["X-TOKEN", 123, "X-Missing"]
headers = manager._build_openapi_request_extra_headers(
mcp_server=server,
oauth2_headers=None,
raw_headers={"x-token": "request-token", 42: "ignored"},
hook_extra_headers=None,
)
assert headers == {"X-TOKEN": "request-token"}
class TestOpenAPIRequestContext:
"""Tests for request-scoped headers on OpenAPI-generated MCP tools."""
def _make_openapi_server(self, auth_type: MCPAuth = MCPAuth.bearer_token):
return MCPServer(
server_id="test-id",
name="openapi_server",
server_name="openapi_server",
url="https://example.com",
transport=MCPTransport.http,
auth_type=auth_type,
spec_path="/path/to/spec.yaml",
)
@pytest.mark.parametrize(
("auth_type", "expected_auth_header"),
[
(MCPAuth.api_key, "ApiKey secret"),
(MCPAuth.basic, "Basic secret"),
(MCPAuth.bearer_token, "Bearer secret"),
],
)
@pytest.mark.asyncio
async def test_direct_openapi_handler_sets_and_resets_request_context(
self, auth_type: MCPAuth, expected_auth_header: str
):
manager = MCPServerManager()
server = self._make_openapi_server(auth_type=auth_type)
tool_name = f"{server.name}-test_tool"
observed_context: Dict[str, Any] = {}
async def handler(**kwargs):
observed_context["arguments"] = kwargs
observed_context["auth"] = _request_auth_header.get()
observed_context["extra"] = _request_extra_headers.get()
return "ok"
previous_tool = global_mcp_tool_registry.tools.get(tool_name)
global_mcp_tool_registry.register_tool(
name=tool_name,
description="test tool",
input_schema={},
handler=handler,
)
try:
result = await manager._call_openapi_tool_handler(
server=server,
tool_name="test_tool",
arguments={"item": "1"},
mcp_auth_header="secret",
request_extra_headers={"X-TOKEN": "request-token"},
)
finally:
if previous_tool is None:
global_mcp_tool_registry.tools.pop(tool_name, None)
else:
global_mcp_tool_registry.tools[tool_name] = previous_tool
assert result.isError is False
assert result.content[0].text == "ok"
assert observed_context == {
"arguments": {"item": "1"},
"auth": expected_auth_header,
"extra": {"X-TOKEN": "request-token"},
}
assert _request_auth_header.get() is None
assert _request_extra_headers.get() is None
@pytest.mark.asyncio
async def test_local_openapi_tool_execution_sets_request_headers(self):
server = self._make_openapi_server(auth_type=MCPAuth.basic)
server.__dict__["extra_headers"] = ["X-TOKEN", "X-Missing", 123]
tool_name = f"{server.name}-test_tool"
observed_context: Dict[str, Any] = {}
async def handler(**kwargs):
observed_context["arguments"] = kwargs
observed_context["auth"] = _request_auth_header.get()
observed_context["extra"] = _request_extra_headers.get()
return "ok"
previous_tool = global_mcp_tool_registry.tools.get(tool_name)
global_mcp_tool_registry.register_tool(
name=tool_name,
description="test tool",
input_schema={},
handler=handler,
)
try:
with patch.object(
mcp_server_module.global_mcp_server_manager,
"_get_mcp_server_from_tool_name",
return_value=server,
):
result = await mcp_server_module.execute_mcp_tool(
name=tool_name,
arguments={"item": "1"},
allowed_mcp_servers=[server],
start_time=datetime.now(),
mcp_auth_header="secret",
oauth2_headers={"Authorization": "Bearer oauth-token"},
raw_headers={"x-token": "request-token"},
)
finally:
if previous_tool is None:
global_mcp_tool_registry.tools.pop(tool_name, None)
else:
global_mcp_tool_registry.tools[tool_name] = previous_tool
assert result.isError is False
assert result.content[0].text == "ok"
assert observed_context == {
"arguments": {"item": "1"},
"auth": "Basic secret",
"extra": {
"Authorization": "Bearer oauth-token",
"X-TOKEN": "request-token",
},
}
assert _request_auth_header.get() is None
assert _request_extra_headers.get() is None
class TestUserAPIKeyAuthJwtClaims:
"""Tests that UserAPIKeyAuth correctly carries jwt_claims."""