mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
test(mcp): cover scoped server name in the SSE initialize handler
This commit is contained in:
parent
64d381d41a
commit
d102ab36c4
1 changed files with 79 additions and 0 deletions
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue