test(mcp): cover scoped server name in the SSE initialize handler

This commit is contained in:
Hari Krishna Kancharla 2026-06-06 19:44:10 -04:00
parent 64d381d41a
commit d102ab36c4

View file

@ -5029,6 +5029,85 @@ class TestGatewayCreateInitializationOptions:
server.create_initialization_options().server_name == "litellm-mcp-server"
)
@pytest.mark.asyncio
async def test_sse_handler_scopes_server_name_from_single_server_path(self):
try:
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import (
global_mcp_server_manager,
handle_sse_mcp,
server,
)
except ImportError:
pytest.skip("MCP server not available")
scoped_server = MCPServer(
server_id="server-123",
name="upstream-server",
alias="grafana",
transport=MCPTransport.http,
url="https://example.com/mcp",
)
captured = {}
async def record_request(scope, receive, send):
captured["server_name"] = server.create_initialization_options().server_name
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/grafana",
"headers": [],
}
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(
UserAPIKeyAuth(api_key="sk-test"),
None,
["grafana"],
None,
None,
None,
),
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[scoped_server],
),
patch.object(
global_mcp_server_manager,
"_ensure_upstream_initialize_instructions_cached",
new_callable=AsyncMock,
),
patch(
"litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers",
new_callable=AsyncMock,
),
patch(
"litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth",
new_callable=AsyncMock,
),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch.object(
mcp_server.sse_session_manager,
"handle_request",
side_effect=record_request,
),
):
await handle_sse_mcp(scope, AsyncMock(), AsyncMock())
assert captured["server_name"] == "grafana"
assert (
server.create_initialization_options().server_name == "litellm-mcp-server"
)
def test_contextvar_set_injects_instructions(self):
"""When ContextVar has a value, it appears in InitializationOptions."""
try: