mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(mcp): preserve edited settings in static connection previews
This commit is contained in:
parent
892d20d86f
commit
245369764e
2 changed files with 75 additions and 2 deletions
|
|
@ -1234,7 +1234,7 @@ if MCP_AVAILABLE:
|
|||
return client_id, client_secret, scopes
|
||||
|
||||
_STAGED_AUTH_VALUE_AUTH_TYPES: Final = frozenset(
|
||||
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.authorization)
|
||||
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token)
|
||||
)
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -1318,8 +1318,15 @@ if MCP_AVAILABLE:
|
|||
if _oauth2_flow == "client_credentials" and not request.token_url:
|
||||
_oauth2_flow = None
|
||||
|
||||
# Static previews inherit credentials before this step, but must not resolve back to
|
||||
# the saved record during client creation and discard the edited connection settings.
|
||||
preview_server_id: Final = (
|
||||
""
|
||||
if request.auth_type in _STAGED_AUTH_VALUE_AUTH_TYPES or request.auth_type in (None, MCPAuth.none)
|
||||
else request.server_id or ""
|
||||
)
|
||||
server_model: Final = MCPServer(
|
||||
server_id=request.server_id or "",
|
||||
server_id=preview_server_id,
|
||||
name=request.alias or request.server_name or "",
|
||||
url=request.url,
|
||||
transport=request.transport,
|
||||
|
|
|
|||
|
|
@ -87,6 +87,72 @@ def _route_has_dependency(route, dependency) -> bool:
|
|||
|
||||
|
||||
class TestExecuteWithMcpClient:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("auth_type", "auth_value", "expected_auth"),
|
||||
(
|
||||
(MCPAuth.none, None, {}),
|
||||
(MCPAuth.basic, "preview:correct", {"Authorization": "Basic cHJldmlldzpjb3JyZWN0"}),
|
||||
(MCPAuth.basic, None, {"Authorization": "Basic cHJldmlldzpzdG9yZWQ="}),
|
||||
(MCPAuth.bearer_token, "edited", {"Authorization": "Bearer edited"}),
|
||||
(MCPAuth.api_key, "edited", {"X-API-Key": "edited"}),
|
||||
(MCPAuth.token, "edited", {"Authorization": "token edited"}),
|
||||
(MCPAuth.authorization, "Custom edited", {"Authorization": "Custom edited"}),
|
||||
),
|
||||
)
|
||||
async def test_static_preview_uses_edited_connection_instead_of_registered_server(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
auth_type: MCPAuth,
|
||||
auth_value: str | None,
|
||||
expected_auth: dict[str, str],
|
||||
) -> None:
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy.management_endpoints import mcp_management_endpoints
|
||||
|
||||
saved: Final = MCPServer(
|
||||
server_id="saved-preview-server",
|
||||
name="saved",
|
||||
url="https://stored.example/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.basic,
|
||||
authentication_token="preview:stored",
|
||||
)
|
||||
manager: Final = MCPServerManager()
|
||||
manager.registry = {saved.server_id: saved}
|
||||
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager)
|
||||
payload: Final = NewMCPServerRequest(
|
||||
server_id=saved.server_id,
|
||||
server_name="edited",
|
||||
url="https://edited.example/mcp",
|
||||
transport=MCPTransport.sse,
|
||||
auth_type=auth_type,
|
||||
credentials={"auth_value": auth_value} if auth_value is not None else None,
|
||||
static_headers={"X-Preview": "edited"},
|
||||
)
|
||||
staged: Final = rest_endpoints._stage_server_test(payload, Headers())
|
||||
|
||||
async def inspect_connection(client: MCPClient) -> dict[str, object]:
|
||||
return {"url": client.server_url, "transport": client.transport_type, "headers": client._get_auth_headers()}
|
||||
|
||||
result: Final = await rest_endpoints._execute_with_mcp_client(
|
||||
staged.request,
|
||||
inspect_connection,
|
||||
mcp_auth_header=staged.mcp_auth_header,
|
||||
oauth2_headers=staged.oauth2_headers,
|
||||
)
|
||||
assert result == {
|
||||
"url": "https://edited.example/mcp",
|
||||
"transport": MCPTransport.sse,
|
||||
"headers": {"X-Preview": "edited", **expected_auth},
|
||||
}
|
||||
assert manager.get_mcp_server_by_id(saved.server_id) is saved
|
||||
assert saved.url == "https://stored.example/mcp"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redacts_stack_trace(self, monkeypatch):
|
||||
async def fake_create_client(*args, **kwargs):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue