diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index a11a33b7bef..0166aeff4be 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -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."""