mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
test: cover openapi mcp request context
This commit is contained in:
parent
67fbb33689
commit
0aaae31f26
1 changed files with 147 additions and 0 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue