mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix mcp scoped server name
This commit is contained in:
parent
68d67212cd
commit
d15796cfc3
3 changed files with 78 additions and 6 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue