mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #18324 from BerriAI/litellm_feat_dynamic_env_propagation_for_stdio_MCP_server
feat: support MCP stdio header env overrides
This commit is contained in:
commit
55bfb24ef8
8 changed files with 295 additions and 41 deletions
|
|
@ -746,8 +746,33 @@ curl --location 'http://localhost:4000/github_mcp/mcp' \
|
|||
3. **Header Forwarding**: LiteLLM automatically forwards matching headers to the backend MCP server
|
||||
4. **Authentication**: The backend MCP server receives both the configured auth headers and the custom headers
|
||||
|
||||
---
|
||||
|
||||
### Passing Request Headers to STDIO env Vars
|
||||
|
||||
If your stdio MCP server needs per-request credentials, you can map HTTP headers from the client request directly into the environment for the launched stdio process. Reference the header name in the env value using the `${X-HEADER_NAME}` syntax. LiteLLM will read that header from the incoming request and set the env var before starting the command.
|
||||
|
||||
```json title="Forward X-GITHUB_PERSONAL_ACCESS_TOKEN header to stdio env" showLineNumbers
|
||||
{
|
||||
"mcpServers": {
|
||||
"github": {
|
||||
"command": "docker",
|
||||
"args": [
|
||||
"run",
|
||||
"-i",
|
||||
"--rm",
|
||||
"-e",
|
||||
"GITHUB_PERSONAL_ACCESS_TOKEN",
|
||||
"ghcr.io/github/github-mcp-server"
|
||||
],
|
||||
"env": {
|
||||
"GITHUB_PERSONAL_ACCESS_TOKEN": "${X-GITHUB_PERSONAL_ACCESS_TOKEN}"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
In this example, when a client makes a request with the `X-GITHUB_PERSONAL_ACCESS_TOKEN` header, the proxy forwards that value into the stdio process as the `GITHUB_PERSONAL_ACCESS_TOKEN` environment variable.
|
||||
|
||||
## Using your MCP with client side credentials
|
||||
|
||||
|
|
|
|||
|
|
@ -84,6 +84,8 @@ def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]:
|
|||
|
||||
|
||||
class MCPServerManager:
|
||||
_STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$")
|
||||
|
||||
def __init__(self):
|
||||
self.registry: Dict[str, MCPServer] = {}
|
||||
self.config_mcp_servers: Dict[str, MCPServer] = {}
|
||||
|
|
@ -671,11 +673,39 @@ class MCPServerManager:
|
|||
#########################################################
|
||||
# Methods that call the upstream MCP servers
|
||||
#########################################################
|
||||
def _build_stdio_env(
|
||||
self,
|
||||
server: MCPServer,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""Resolve stdio env values, supporting header-driven placeholders."""
|
||||
|
||||
if server.transport != MCPTransport.stdio or not server.env:
|
||||
return None
|
||||
|
||||
resolved_env: Dict[str, str] = {}
|
||||
normalized_headers = {k.lower(): v for k, v in (raw_headers or {}).items()}
|
||||
|
||||
for env_key, env_value in server.env.items():
|
||||
stripped_value = env_value.strip()
|
||||
match = self._STDIO_ENV_TEMPLATE_PATTERN.match(stripped_value)
|
||||
if match:
|
||||
header_name = match.group(1)
|
||||
header_value = normalized_headers.get(header_name.lower())
|
||||
if header_value is None:
|
||||
continue
|
||||
resolved_env[env_key] = header_value
|
||||
else:
|
||||
resolved_env[env_key] = env_value
|
||||
|
||||
return resolved_env
|
||||
|
||||
def _create_mcp_client(
|
||||
self,
|
||||
server: MCPServer,
|
||||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
stdio_env: Optional[Dict[str, str]] = None,
|
||||
) -> MCPClient:
|
||||
"""
|
||||
Create an MCPClient instance for the given server.
|
||||
|
|
@ -692,10 +722,13 @@ class MCPServerManager:
|
|||
# Handle stdio transport
|
||||
if transport == MCPTransport.stdio:
|
||||
# For stdio, we need to get the stdio config from the server
|
||||
resolved_env = stdio_env if stdio_env is not None else server.env or {}
|
||||
stdio_config: Optional[MCPStdioConfig] = None
|
||||
if server.command and server.args is not None:
|
||||
stdio_config = MCPStdioConfig(
|
||||
command=server.command, args=server.args, env=server.env or {}
|
||||
command=server.command,
|
||||
args=server.args,
|
||||
env=resolved_env,
|
||||
)
|
||||
|
||||
return MCPClient(
|
||||
|
|
@ -725,6 +758,7 @@ class MCPServerManager:
|
|||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
add_prefix: bool = True,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
) -> List[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
|
@ -751,10 +785,13 @@ class MCPServerManager:
|
|||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
## HANDLE OPENAPI TOOLS
|
||||
|
|
@ -784,6 +821,7 @@ class MCPServerManager:
|
|||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
add_prefix: bool = True,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
) -> List[Prompt]:
|
||||
"""
|
||||
Helper method to get prompts from a single MCP server with prefixed names.
|
||||
|
|
@ -807,10 +845,13 @@ class MCPServerManager:
|
|||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
prompts = await client.list_prompts()
|
||||
|
|
@ -833,6 +874,7 @@ class MCPServerManager:
|
|||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
add_prefix: bool = True,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
) -> List[Resource]:
|
||||
"""Fetch available resources from a single MCP server."""
|
||||
|
||||
|
|
@ -847,10 +889,13 @@ class MCPServerManager:
|
|||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
resources = await client.list_resources()
|
||||
|
|
@ -873,6 +918,7 @@ class MCPServerManager:
|
|||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
add_prefix: bool = True,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
) -> List[ResourceTemplate]:
|
||||
"""Fetch available resource templates from a single MCP server."""
|
||||
|
||||
|
|
@ -887,10 +933,13 @@ class MCPServerManager:
|
|||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
resource_templates = await client.list_resource_templates()
|
||||
|
|
@ -913,6 +962,7 @@ class MCPServerManager:
|
|||
url: AnyUrl,
|
||||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
) -> ReadResourceResult:
|
||||
"""Read resource contents from a specific MCP server."""
|
||||
|
||||
|
|
@ -924,10 +974,13 @@ class MCPServerManager:
|
|||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
return await client.read_resource(url)
|
||||
|
|
@ -939,6 +992,7 @@ class MCPServerManager:
|
|||
arguments: Optional[Dict[str, Any]] = None,
|
||||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
) -> GetPromptResult:
|
||||
"""Fetch a specific prompt definition from a single MCP server."""
|
||||
|
||||
|
|
@ -950,10 +1004,13 @@ class MCPServerManager:
|
|||
extra_headers = {}
|
||||
extra_headers.update(server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(server, raw_headers)
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
get_prompt_request_params = GetPromptRequestParams(
|
||||
|
|
@ -1742,10 +1799,13 @@ class MCPServerManager:
|
|||
extra_headers = {}
|
||||
extra_headers.update(mcp_server.static_headers)
|
||||
|
||||
stdio_env = self._build_stdio_env(mcp_server, raw_headers)
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=mcp_server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
call_tool_params = MCPCallToolRequestParams(
|
||||
|
|
|
|||
|
|
@ -71,12 +71,17 @@ if MCP_AVAILABLE:
|
|||
for tool in tools
|
||||
]
|
||||
|
||||
async def _get_tools_for_single_server(server, server_auth_header):
|
||||
async def _get_tools_for_single_server(
|
||||
server,
|
||||
server_auth_header,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
"""Helper function to get tools for a single server."""
|
||||
tools = await global_mcp_server_manager._get_tools_from_server(
|
||||
server=server,
|
||||
mcp_auth_header=server_auth_header,
|
||||
add_prefix=False,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
# Filter tools based on allowed_tools configuration
|
||||
|
|
@ -122,6 +127,7 @@ if MCP_AVAILABLE:
|
|||
try:
|
||||
# Extract auth headers from request
|
||||
headers = request.headers
|
||||
raw_headers_from_request = dict(headers)
|
||||
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(
|
||||
headers
|
||||
)
|
||||
|
|
@ -148,7 +154,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
try:
|
||||
list_tools_result = await _get_tools_for_single_server(
|
||||
server, server_auth_header
|
||||
server, server_auth_header, raw_headers_from_request
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
|
|
@ -169,7 +175,7 @@ if MCP_AVAILABLE:
|
|||
|
||||
try:
|
||||
tools_result = await _get_tools_for_single_server(
|
||||
server, server_auth_header
|
||||
server, server_auth_header, raw_headers_from_request
|
||||
)
|
||||
list_tools_result.extend(tools_result)
|
||||
except Exception as e:
|
||||
|
|
@ -232,13 +238,13 @@ if MCP_AVAILABLE:
|
|||
# but they weren't being extracted and passed to call_mcp_tool.
|
||||
# This fix ensures auth headers are properly extracted from the HTTP request
|
||||
# and passed through to the MCP server for authentication.
|
||||
headers = request.headers
|
||||
raw_headers_from_request = dict(headers)
|
||||
mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers(
|
||||
request.headers
|
||||
headers
|
||||
)
|
||||
mcp_server_auth_headers = (
|
||||
MCPRequestHandler._get_mcp_server_auth_headers_from_headers(
|
||||
request.headers
|
||||
)
|
||||
MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
|
||||
)
|
||||
|
||||
# Add extracted headers to data dict to pass to call_mcp_tool
|
||||
|
|
@ -246,6 +252,7 @@ if MCP_AVAILABLE:
|
|||
data["mcp_auth_header"] = mcp_auth_header
|
||||
if mcp_server_auth_headers:
|
||||
data["mcp_server_auth_headers"] = mcp_server_auth_headers
|
||||
data["raw_headers"] = raw_headers_from_request
|
||||
|
||||
result = await call_mcp_tool(**data)
|
||||
return result
|
||||
|
|
@ -300,6 +307,7 @@ if MCP_AVAILABLE:
|
|||
operation,
|
||||
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
|
||||
oauth2_headers: Optional[Dict[str, str]] = None,
|
||||
raw_headers: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
"""
|
||||
Common helper to create MCP client, execute operation, and ensure proper cleanup.
|
||||
|
|
@ -312,17 +320,27 @@ if MCP_AVAILABLE:
|
|||
Operation result or error response
|
||||
"""
|
||||
try:
|
||||
server_model = MCPServer(
|
||||
server_id=request.server_id or "",
|
||||
name=request.alias or request.server_name or "",
|
||||
url=request.url,
|
||||
transport=request.transport,
|
||||
auth_type=request.auth_type,
|
||||
mcp_info=request.mcp_info,
|
||||
command=request.command,
|
||||
args=request.args,
|
||||
env=request.env,
|
||||
)
|
||||
|
||||
stdio_env = global_mcp_server_manager._build_stdio_env(
|
||||
server_model, raw_headers
|
||||
)
|
||||
|
||||
client = global_mcp_server_manager._create_mcp_client(
|
||||
server=MCPServer(
|
||||
server_id=request.server_id or "",
|
||||
name=request.alias or request.server_name or "",
|
||||
url=request.url,
|
||||
transport=request.transport,
|
||||
auth_type=request.auth_type,
|
||||
mcp_info=request.mcp_info,
|
||||
),
|
||||
server=server_model,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=oauth2_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
return await operation(client)
|
||||
|
|
@ -338,7 +356,8 @@ if MCP_AVAILABLE:
|
|||
|
||||
@router.post("/test/connection")
|
||||
async def test_connection(
|
||||
request: NewMCPServerRequest,
|
||||
request: Request,
|
||||
new_mcp_server_request: NewMCPServerRequest,
|
||||
):
|
||||
"""
|
||||
Test if we can connect to the provided MCP server before adding it
|
||||
|
|
@ -351,7 +370,11 @@ if MCP_AVAILABLE:
|
|||
await client.run_with_session(_noop)
|
||||
return {"status": "ok"}
|
||||
|
||||
return await _execute_with_mcp_client(request, _test_connection_operation)
|
||||
return await _execute_with_mcp_client(
|
||||
new_mcp_server_request,
|
||||
_test_connection_operation,
|
||||
raw_headers=dict(request.headers),
|
||||
)
|
||||
|
||||
@router.post("/test/tools/list")
|
||||
async def test_tools_list(
|
||||
|
|
@ -405,4 +428,5 @@ if MCP_AVAILABLE:
|
|||
_list_tools_operation,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=dict(request.headers),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -775,6 +775,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
filtered_tools = filter_tools_by_allowed_tools(tools, server)
|
||||
|
|
@ -854,6 +855,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
all_prompts.extend(prompts)
|
||||
|
|
@ -912,6 +914,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
all_resources.extend(resources)
|
||||
|
||||
|
|
@ -969,6 +972,7 @@ if MCP_AVAILABLE:
|
|||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
add_prefix=add_prefix,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
)
|
||||
all_resource_templates.extend(resource_templates)
|
||||
|
|
@ -1392,6 +1396,7 @@ if MCP_AVAILABLE:
|
|||
arguments=arguments,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
async def mcp_read_resource(
|
||||
|
|
@ -1440,6 +1445,7 @@ if MCP_AVAILABLE:
|
|||
url=url,
|
||||
mcp_auth_header=server_auth_header,
|
||||
extra_headers=extra_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
def _get_standard_logging_mcp_tool_call(
|
||||
|
|
|
|||
|
|
@ -812,10 +812,19 @@ async def test_get_tools_from_mcp_servers():
|
|||
return_value=["server1_id", "server2_id"]
|
||||
)
|
||||
mock_manager_2.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2
|
||||
async def mock_get_tools_side_effect(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
):
|
||||
if server.server_id == "server1_id":
|
||||
return [mock_tool_1]
|
||||
return [mock_tool_2]
|
||||
|
||||
mock_manager_2._get_tools_from_server = AsyncMock(
|
||||
side_effect=lambda server, mcp_auth_header=None, extra_headers=None, add_prefix=False: (
|
||||
[mock_tool_1] if server.server_id == "server1_id" else [mock_tool_2]
|
||||
)
|
||||
side_effect=mock_get_tools_side_effect
|
||||
)
|
||||
|
||||
with patch(
|
||||
|
|
@ -1693,6 +1702,7 @@ async def test_get_tools_for_single_server():
|
|||
server=mock_server,
|
||||
mcp_auth_header="Bearer test_token",
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
|
|
|
|||
|
|
@ -294,6 +294,7 @@ async def test_mcp_get_prompt_success():
|
|||
arguments={"foo": "bar"},
|
||||
mcp_auth_header={"Authorization": "token"},
|
||||
extra_headers={"X-Test": "1"},
|
||||
raw_headers=None,
|
||||
)
|
||||
assert result is prompt_result
|
||||
|
||||
|
|
@ -349,6 +350,7 @@ async def test_mcp_read_resource_success():
|
|||
url="https://example.com/resource",
|
||||
mcp_auth_header={"Authorization": "token"},
|
||||
extra_headers={"X-Test": "1"},
|
||||
raw_headers=None,
|
||||
)
|
||||
assert result is read_result
|
||||
|
||||
|
|
@ -428,7 +430,11 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails():
|
|||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=True
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
):
|
||||
if server.name == "working_server":
|
||||
# Working server returns tools
|
||||
|
|
@ -524,7 +530,11 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
|
|||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=True
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
):
|
||||
# All servers fail
|
||||
raise Exception(f"Server {server.name} connection failed")
|
||||
|
|
@ -839,13 +849,19 @@ async def test_oauth2_headers_passed_to_mcp_client():
|
|||
# This will capture the arguments passed to _create_mcp_client
|
||||
captured_client_args = {}
|
||||
|
||||
def mock_create_mcp_client(server, mcp_auth_header=None, extra_headers=None):
|
||||
def mock_create_mcp_client(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
):
|
||||
# Capture the arguments for verification
|
||||
captured_client_args.update(
|
||||
{
|
||||
"server": server,
|
||||
"mcp_auth_header": mcp_auth_header,
|
||||
"extra_headers": extra_headers,
|
||||
"stdio_env": stdio_env,
|
||||
}
|
||||
)
|
||||
# Return a mock client that doesn't actually connect
|
||||
|
|
@ -934,7 +950,11 @@ async def test_list_tools_single_server_unprefixed_names():
|
|||
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
):
|
||||
tool = MagicMock()
|
||||
tool.name = f"{server.alias}-toolA" if add_prefix else "toolA"
|
||||
|
|
@ -1006,7 +1026,11 @@ async def test_list_tools_multiple_servers_prefixed_names():
|
|||
)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=True
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
):
|
||||
tool = MagicMock()
|
||||
# When multiple servers, add_prefix should be True -> prefixed names
|
||||
|
|
@ -1147,7 +1171,11 @@ async def test_list_tools_filters_by_key_team_permissions():
|
|||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
):
|
||||
# Return 4 tools, but only 2 should be allowed
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -1248,7 +1276,11 @@ async def test_list_tools_with_team_tool_permissions_inheritance():
|
|||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
):
|
||||
# Return 4 tools
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -1334,7 +1366,11 @@ async def test_list_tools_with_no_tool_permissions_shows_all():
|
|||
mock_manager.get_mcp_server_by_id = lambda server_id: server
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=False
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=False,
|
||||
raw_headers=None,
|
||||
):
|
||||
# Return 3 tools
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -1423,7 +1459,11 @@ async def test_list_tools_strips_prefix_when_matching_permissions():
|
|||
mock_manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, extra_headers=None, add_prefix=True
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
):
|
||||
# Return tools WITH prefix (as they come from MCP server)
|
||||
tool1 = MagicMock()
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from fastapi import HTTPException
|
|||
# Add the parent directory to the path so we can import litellm
|
||||
sys.path.insert(0, "../../../../../")
|
||||
|
||||
|
||||
import httpx
|
||||
from mcp import ReadResourceResult, Resource
|
||||
from mcp.types import (
|
||||
|
|
@ -99,6 +100,53 @@ class TestMCPServerManager:
|
|||
assert client.stdio_config["args"] == ["server.js"]
|
||||
assert client.stdio_config["env"] == {"NODE_ENV": "test"}
|
||||
|
||||
def test_build_stdio_env_only_accepts_x_prefixed_placeholders(self):
|
||||
"""Ensure only ${X-*} placeholders are substituted from headers."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="stdio-server-env",
|
||||
name="stdio_env",
|
||||
transport=MCPTransport.stdio,
|
||||
command="node",
|
||||
args=["server.js"],
|
||||
env={
|
||||
"PASSTHROUGH": "${X-Test-Header}",
|
||||
"STATIC": "value",
|
||||
"IGNORED": "${Not-Allowed}",
|
||||
},
|
||||
)
|
||||
|
||||
env = manager._build_stdio_env(
|
||||
server,
|
||||
raw_headers={
|
||||
"x-test-header": "resolved-value",
|
||||
"x-not-used": "other",
|
||||
},
|
||||
)
|
||||
|
||||
assert env == {
|
||||
"PASSTHROUGH": "resolved-value",
|
||||
"STATIC": "value",
|
||||
"IGNORED": "${Not-Allowed}",
|
||||
}
|
||||
|
||||
def test_build_stdio_env_missing_header_skips_entry(self):
|
||||
"""Ensure missing headers drop the placeholder from the resolved env."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="stdio-server-env-miss",
|
||||
name="stdio_env_miss",
|
||||
transport=MCPTransport.stdio,
|
||||
command="node",
|
||||
args=["server.js"],
|
||||
env={"EXPECTED": "${X-Missing}"},
|
||||
)
|
||||
|
||||
env = manager._build_stdio_env(server, raw_headers={})
|
||||
|
||||
# When the header isn't provided, the key is omitted entirely
|
||||
assert env == {}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_with_server_specific_auth_headers(self):
|
||||
"""Test list_tools method with server-specific auth headers"""
|
||||
|
|
@ -123,7 +171,10 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock _get_tools_from_server to return different results
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, mcp_protocol_version=None
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
if server.name == "github":
|
||||
tool1 = MagicMock()
|
||||
|
|
@ -174,7 +225,10 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock _get_tools_from_server
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, mcp_protocol_version=None
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
assert mcp_auth_header == "legacy-token" # Should use legacy header
|
||||
tool = MagicMock()
|
||||
|
|
@ -209,7 +263,10 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock _get_tools_from_server
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, mcp_protocol_version=None
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
|
|
@ -373,6 +430,7 @@ class TestMCPServerManager:
|
|||
server=server,
|
||||
mcp_auth_header="auth",
|
||||
extra_headers=None,
|
||||
stdio_env=None,
|
||||
)
|
||||
mock_client.list_resource_templates.assert_awaited_once()
|
||||
mock_prefix.assert_called_once_with(mock_templates, server, add_prefix=False)
|
||||
|
|
@ -554,7 +612,10 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock _get_tools_from_server
|
||||
async def mock_get_tools_from_server(
|
||||
server, mcp_auth_header=None, mcp_protocol_version=None
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
mcp_protocol_version=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
assert (
|
||||
mcp_auth_header == "server-specific-token"
|
||||
|
|
@ -587,7 +648,11 @@ class TestMCPServerManager:
|
|||
manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
# Mock successful _get_tools_from_server
|
||||
async def mock_get_tools_from_server(server, mcp_auth_header=None):
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
tool1 = MagicMock()
|
||||
tool1.name = "tool1"
|
||||
tool2 = MagicMock()
|
||||
|
|
@ -621,7 +686,11 @@ class TestMCPServerManager:
|
|||
manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
# Mock failed _get_tools_from_server
|
||||
async def mock_get_tools_from_server(server, mcp_auth_header=None):
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
raise Exception("Connection timeout")
|
||||
|
||||
manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
|
@ -683,7 +752,11 @@ class TestMCPServerManager:
|
|||
manager.get_mcp_server_by_id = mock_get_server_by_id
|
||||
|
||||
# Mock _get_tools_from_server with different results
|
||||
async def mock_get_tools_from_server(server, mcp_auth_header=None):
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
if server.server_id == "server1":
|
||||
tool = MagicMock()
|
||||
tool.name = "tool1"
|
||||
|
|
@ -724,7 +797,11 @@ class TestMCPServerManager:
|
|||
manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
# Mock _get_tools_from_server to verify auth header is passed
|
||||
async def mock_get_tools_from_server(server, mcp_auth_header=None):
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
assert mcp_auth_header == "test-token"
|
||||
tool = MagicMock()
|
||||
tool.name = "tool1"
|
||||
|
|
|
|||
|
|
@ -37,7 +37,13 @@ async def test_test_tools_list_forwards_mcp_auth_header(monkeypatch):
|
|||
|
||||
captured: dict = {}
|
||||
|
||||
async def fake_execute(request, operation, mcp_auth_header=None, oauth2_headers=None):
|
||||
async def fake_execute(
|
||||
request,
|
||||
operation,
|
||||
mcp_auth_header=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
captured["oauth2_headers"] = oauth2_headers
|
||||
return {
|
||||
|
|
@ -87,7 +93,13 @@ async def test_test_tools_list_extracts_oauth2_headers(monkeypatch):
|
|||
|
||||
captured: dict = {}
|
||||
|
||||
async def fake_execute(request, operation, mcp_auth_header=None, oauth2_headers=None):
|
||||
async def fake_execute(
|
||||
request,
|
||||
operation,
|
||||
mcp_auth_header=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
captured["oauth2_headers"] = oauth2_headers
|
||||
return {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue