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:
YutaSaito 2025-12-24 06:29:53 +09:00 • committed by GitHub
commit 55bfb24ef8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 295 additions and 41 deletions

View file

@ -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

View file

@ -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(

View file

@ -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),
)

View file

@ -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(

View file

@ -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

View file

@ -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()

View file

@ -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"

View file

@ -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 {