This commit is contained in:
Yug 2026-04-29 16:51:32 +05:30
parent b0d0de00b1
commit 4218c08fd8
6 changed files with 122 additions and 35 deletions

View file

@ -1282,6 +1282,7 @@ class MCPServerManager:
extra_headers: Optional[Dict[str, str]] = None,
add_prefix: bool = True,
raw_headers: Optional[Dict[str, str]] = None,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[MCPTool]:
"""
Helper method to get tools from a single MCP server with prefixed names.
@ -1315,6 +1316,7 @@ class MCPServerManager:
mcp_auth_header=mcp_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
user_api_key_auth=user_api_key_auth,
)
## HANDLE OPENAPI TOOLS

View file

@ -102,6 +102,13 @@ try:
Tool,
)
from mcp.server.session import ServerSession as _McpServerSession
import weakref
# Storage for session-specific auth context to avoid fragile monkey-patching
# and context-leakage between concurrent SSE sessions.
_session_auth_storage: "weakref.WeakKeyDictionary[Any, MCPAuthenticatedUser]" = (
weakref.WeakKeyDictionary()
)
active_mcp_session_var: contextvars.ContextVar[Optional[_McpServerSession]] = (
contextvars.ContextVar("active_mcp_session", default=None)
@ -243,20 +250,12 @@ if MCP_AVAILABLE:
json_response=False, # enables SSE streaming
stateless=True,
)
# Create SSE session manager
sse_session_manager = StreamableHTTPSessionManager(
app=server,
event_store=None,
json_response=False, # Use SSE responses for this endpoint
stateless=True,
)
# Context managers for proper lifecycle management
_session_manager_cm = None
_sse_session_manager_cm = None
async def initialize_session_managers():
"""Initialize the session managers. Can be called from main app lifespan."""
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm
# Use async lock to prevent concurrent initialization
async with _INITIALIZATION_LOCK:
if _SESSION_MANAGERS_INITIALIZED:
@ -264,10 +263,8 @@ if MCP_AVAILABLE:
verbose_logger.info("Initializing MCP session managers...")
# Start the session managers with context managers
_session_manager_cm = session_manager.run()
_sse_session_manager_cm = sse_session_manager.run()
# Enter the context managers
await _session_manager_cm.__aenter__()
await _sse_session_manager_cm.__aenter__()
_SESSION_MANAGERS_INITIALIZED = True
verbose_logger.info(
"MCP Server started with StreamableHTTP session manager and SSE transport!"
@ -275,18 +272,15 @@ if MCP_AVAILABLE:
async def shutdown_session_managers():
"""Shutdown the session managers."""
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _sse_session_manager_cm
global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm
if _SESSION_MANAGERS_INITIALIZED:
verbose_logger.info("Shutting down MCP session managers...")
try:
if _session_manager_cm:
await _session_manager_cm.__aexit__(None, None, None)
if _sse_session_manager_cm:
await _sse_session_manager_cm.__aexit__(None, None, None)
except Exception as e:
verbose_logger.exception(f"Error during session manager shutdown: {e}")
_session_manager_cm = None
_sse_session_manager_cm = None
_SESSION_MANAGERS_INITIALIZED = False
@contextlib.asynccontextmanager
@ -2688,10 +2682,11 @@ if MCP_AVAILABLE:
"SSE connection established, running server loop..."
)
# ContextVars can sometimes be lost when the MCP SDK spawns internal tasks
# ContextVars are lost when the MCP SDK spawns internal tasks
# (e.g. _receive_loop), so tool handlers can't read auth_context_var reliably.
# Storing it on the read_stream ensures it's isolated per SSE session.
streams[0]._litellm_auth_context = MCPAuthenticatedUser( # type: ignore
# Storing it in session-isolated storage ensures it's isolated per SSE session.
# We store it before calling server.run so it's available for the session's lifespan.
_session_auth_storage[streams[0]] = MCPAuthenticatedUser( # type: ignore
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
@ -2748,15 +2743,9 @@ if MCP_AVAILABLE:
raw_headers,
) = await extract_mcp_auth_context(scope, path)
_sse_client_ip = IPAddressUtils.get_mcp_client_ip(request)
set_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=mcp_servers,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_sse_client_ip,
)
# set_auth_context here is a no-op for actual tool execution since the SDK
# processes messages in background tasks that don't inherit this ContextVar.
# Authentication must be recovered from the session-auth-storage during execution.
except HTTPException:
raise
except Exception as e:
@ -2923,7 +2912,7 @@ if MCP_AVAILABLE:
session = request_ctx.get().session
read_stream = getattr(session, "_read_stream", None)
stored = getattr(read_stream, "_litellm_auth_context", None)
stored = _session_auth_storage.get(read_stream) if read_stream else None
except Exception as e:
verbose_logger.debug(
f"get_or_extract_auth_context FALLBACK failed: {e}"

View file

@ -84,7 +84,10 @@ async def test_get_or_extract_auth_context_fallback():
mock_user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-1")
from litellm.proxy._experimental.mcp_server.server import MCPAuthenticatedUser
mock_read_stream._litellm_auth_context = MCPAuthenticatedUser(
mock_session._read_stream = mock_read_stream
from litellm.proxy._experimental.mcp_server.server import _session_auth_storage
_session_auth_storage[mock_read_stream] = MCPAuthenticatedUser(
user_api_key_auth=mock_user_auth,
mcp_auth_header=None,
mcp_servers=None,
@ -93,7 +96,6 @@ async def test_get_or_extract_auth_context_fallback():
raw_headers=None,
client_ip=None
)
mock_session._read_stream = mock_read_stream
mock_request_ctx = MagicMock()
mock_request_ctx.get.return_value.session = mock_session

View file

@ -0,0 +1,92 @@
import pytest
from unittest.mock import MagicMock, AsyncMock, patch
from litellm.proxy._types import MCPTransportType
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.proxy._experimental.mcp_server.server import (
_get_prompts_from_mcp_servers,
_get_resources_from_mcp_servers,
_get_resource_templates_from_mcp_servers,
_get_tools_from_mcp_servers,
)
@pytest.mark.asyncio
async def test_get_prompts_from_mcp_servers_coverage():
server1 = MCPServer(
server_id="test-1", name="test1", transport="stdio", url="http://test1"
)
server2 = MCPServer(
server_id="test-2", name="test2", transport="stdio", url="http://test2"
)
mock_prompt = MagicMock()
mock_prompt.name = "test_prompt"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1, server2]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_prompts_from_server", new_callable=AsyncMock) as mock_get:
mock_get.side_effect = [[mock_prompt], Exception("Server error")]
result = await _get_prompts_from_mcp_servers(
user_api_key_auth=None,
mcp_auth_header=None,
mcp_servers=["test1", "test2"]
)
assert len(result) == 1
assert result[0] == mock_prompt
@pytest.mark.asyncio
async def test_get_resources_from_mcp_servers_coverage():
server1 = MCPServer(
server_id="test-1", name="test1", transport="stdio", url="http://test1"
)
mock_resource = MagicMock()
mock_resource.name = "test_resource"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resources_from_server", new_callable=AsyncMock) as mock_get:
mock_get.return_value = [mock_resource]
result = await _get_resources_from_mcp_servers(
user_api_key_auth=None,
mcp_auth_header=None,
mcp_servers=["test1"]
)
assert len(result) == 1
assert result[0] == mock_resource
@pytest.mark.asyncio
async def test_get_resource_templates_from_mcp_servers_coverage():
server1 = MCPServer(
server_id="test-1", name="test1", transport="stdio", url="http://test1"
)
mock_template = MagicMock()
mock_template.name = "test_template"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_resource_templates_from_server", new_callable=AsyncMock) as mock_get:
mock_get.return_value = [mock_template]
result = await _get_resource_templates_from_mcp_servers(
user_api_key_auth=None,
mcp_auth_header=None,
mcp_servers=["test1"]
)
assert len(result) == 1
assert result[0] == mock_template
@pytest.mark.asyncio
async def test_get_tools_from_mcp_servers_coverage():
server1 = MCPServer(
server_id="test-1", name="test1", transport="stdio", url="http://test1"
)
mock_tool = MagicMock()
mock_tool.name = "test_tool"
with patch("litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", return_value=[server1]):
with patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", new_callable=AsyncMock) as mock_get:
mock_get.return_value = [mock_tool]
# test with some tracking headers
result = await _get_tools_from_mcp_servers(
user_api_key_auth=None,
mcp_auth_header=None,
mcp_servers=["test1"],
log_list_tools_to_spendlogs=True,
litellm_trace_id="test-trace"
)
assert len(result) == 1
assert result[0] == mock_tool

View file

@ -40,12 +40,14 @@ async def test_get_or_extract_auth_context_fallback():
auth_data = UserAPIKeyAuth(api_key="fallback-key")
auth_user = MCPAuthenticatedUser(user_api_key_auth=auth_data)
# Mock request_ctx.get().session._read_stream._litellm_auth_context
# Mock request_ctx.get().session._read_stream
mock_session = MagicMock()
mock_read_stream = MagicMock()
mock_read_stream._litellm_auth_context = auth_user
mock_session._read_stream = mock_read_stream
from litellm.proxy._experimental.mcp_server.server import _session_auth_storage
_session_auth_storage[mock_read_stream] = auth_user
mock_request_ctx = MagicMock()
mock_request_ctx.get.return_value.session = mock_session

View file

@ -548,7 +548,7 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()
@ -588,7 +588,7 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()
@ -634,7 +634,7 @@ class TestHookHeaderMergePriority:
captured_extra_headers: Dict[str, Any] = {}
async def fake_create_mcp_client(
server, mcp_auth_header=None, extra_headers=None, stdio_env=None
server, mcp_auth_header=None, extra_headers=None, stdio_env=None, **kwargs
):
captured_extra_headers["value"] = extra_headers
mock_client = MagicMock()