mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
resolve
This commit is contained in:
parent
b0d0de00b1
commit
4218c08fd8
6 changed files with 122 additions and 35 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
92
tests/mcp_tests/test_server_coverage.py
Normal file
92
tests/mcp_tests/test_server_coverage.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue