fix mcp scoped server name

This commit is contained in:
Hari Krishna Kancharla 2026-06-06 18:57:16 -04:00
parent 68d67212cd
commit d15796cfc3
3 changed files with 78 additions and 6 deletions

View file

@ -19,3 +19,7 @@ _mcp_active_toolset_id: ContextVar[Optional[str]] = ContextVar(
_mcp_gateway_initialize_instructions: ContextVar[Optional[str]] = ContextVar(
"_mcp_gateway_initialize_instructions", default=None
)
_mcp_gateway_server_name: ContextVar[Optional[str]] = ContextVar(
"_mcp_gateway_server_name", default=None
)

View file

@ -47,6 +47,7 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
_mcp_gateway_server_name,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug
from litellm.proxy._experimental.mcp_server.utils import (
@ -323,10 +324,14 @@ if MCP_AVAILABLE:
notification_options=notification_options,
experimental_capabilities=experimental_capabilities or {},
)
updates: Dict[str, Any] = {}
merged = _mcp_gateway_initialize_instructions.get()
if merged is not None:
return opts.model_copy(update={"instructions": merged})
return opts
updates["instructions"] = merged
scoped_server_name = _mcp_gateway_server_name.get()
if scoped_server_name is not None:
updates["server_name"] = scoped_server_name
return opts.model_copy(update=updates) if updates else opts
########################################################
############ Initialize the MCP Server #################
@ -1544,6 +1549,7 @@ if MCP_AVAILABLE:
user_api_key_auth: Optional[UserAPIKeyAuth],
mcp_servers: Optional[List[str]],
client_ip: Optional[str],
scoped_server_endpoint: bool = False,
) -> AsyncIterator[None]:
allowed = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
@ -1565,11 +1571,22 @@ if MCP_AVAILABLE:
return_exceptions=True,
)
merged = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed)
tok = _mcp_gateway_initialize_instructions.set(merged)
scoped_server_name = None
if scoped_server_endpoint and len(allowed) == 1:
scoped_server = allowed[0]
scoped_server_name = (
scoped_server.alias
or scoped_server.server_name
or scoped_server.name
or scoped_server.server_id
)
instructions_token = _mcp_gateway_initialize_instructions.set(merged)
server_name_token = _mcp_gateway_server_name.set(scoped_server_name)
try:
yield
finally:
_mcp_gateway_initialize_instructions.reset(tok)
_mcp_gateway_initialize_instructions.reset(instructions_token)
_mcp_gateway_server_name.reset(server_name_token)
async def _get_tools_from_mcp_servers( # noqa: PLR0915
user_api_key_auth: Optional[UserAPIKeyAuth],
@ -3620,6 +3637,7 @@ if MCP_AVAILABLE:
oauth2_headers,
raw_headers,
) = await extract_mcp_auth_context(scope, path)
scoped_server_endpoint = len(_get_mcp_servers_in_path(path) or []) == 1
# Extract client IP for MCP access control
_client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope))
@ -3896,6 +3914,7 @@ if MCP_AVAILABLE:
user_api_key_auth,
mcp_servers,
_client_ip,
scoped_server_endpoint=scoped_server_endpoint,
):
await target_manager.handle_request(scope, receive, local_send)
if use_stateful and session_id and scope.get("method") == "DELETE":
@ -3980,6 +3999,7 @@ if MCP_AVAILABLE:
oauth2_headers,
raw_headers,
) = await extract_mcp_auth_context(scope, path)
scoped_server_endpoint = len(_get_mcp_servers_in_path(path) or []) == 1
# Extract client IP for MCP access control
_sse_client_ip = IPAddressUtils.get_mcp_client_ip(StarletteRequest(scope))
@ -4052,6 +4072,7 @@ if MCP_AVAILABLE:
user_api_key_auth,
mcp_servers,
_sse_client_ip,
scoped_server_endpoint=scoped_server_endpoint,
):
await sse_session_manager.handle_request(scope, receive, send)
except MCPUpstreamAuthError as e:

View file

@ -4970,17 +4970,64 @@ class TestGatewayCreateInitializationOptions:
try:
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_gateway_initialize_instructions,
_mcp_gateway_server_name,
)
from litellm.proxy._experimental.mcp_server.server import server
except ImportError:
pytest.skip("MCP server not available")
tok = _mcp_gateway_initialize_instructions.set(None)
instructions_token = _mcp_gateway_initialize_instructions.set(None)
server_name_token = _mcp_gateway_server_name.set(None)
try:
opts = server.create_initialization_options()
assert getattr(opts, "instructions", None) is None
assert opts.server_name == "litellm-mcp-server"
finally:
_mcp_gateway_initialize_instructions.reset(tok)
_mcp_gateway_initialize_instructions.reset(instructions_token)
_mcp_gateway_server_name.reset(server_name_token)
@pytest.mark.asyncio
async def test_scoped_request_uses_configured_server_alias(self):
try:
from litellm.proxy._experimental.mcp_server.server import (
_gateway_initialize_instructions_request_scope,
global_mcp_server_manager,
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",
)
with (
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,
),
):
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth=None,
mcp_servers=["grafana"],
client_ip=None,
scoped_server_endpoint=True,
):
assert server.create_initialization_options().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."""