mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #38806 from BerriAI/litellm_fix_mcp_test_connection_oauth_bearer
fix(mcp): forward staged credentials on /mcp-rest/test/connection like /test/tools/list
This commit is contained in:
commit
0e9e2c01f3
2 changed files with 186 additions and 33 deletions
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import importlib
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
|
@ -8,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal
|
|||
import anyio
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT
|
||||
|
|
@ -1166,6 +1168,45 @@ if MCP_AVAILABLE:
|
|||
scopes: Final[list[str] | None] = scopes_raw if isinstance(scopes_raw, list) else None
|
||||
return client_id, client_secret, scopes
|
||||
|
||||
_STAGED_AUTH_VALUE_AUTH_TYPES: Final = frozenset(
|
||||
(MCPAuth.api_key, MCPAuth.bearer_token, MCPAuth.basic, MCPAuth.authorization)
|
||||
)
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _StagedServerTest:
|
||||
request: NewMCPServerRequest
|
||||
mcp_auth_header: str | None
|
||||
oauth2_headers: dict[str, str] | None
|
||||
|
||||
def _stage_server_test(new_mcp_server_request: NewMCPServerRequest, headers: Headers) -> _StagedServerTest:
|
||||
"""
|
||||
Resolve the credentials a not-yet-saved server config carries for a preview call.
|
||||
|
||||
Both preview endpoints (``/test/connection`` and ``/test/tools/list``) must hand the
|
||||
temporary client the same credentials, or a server that the saved connection reaches
|
||||
fine fails one of them.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
|
||||
request: Final = _inherit_credentials_from_existing_server(new_mcp_server_request)
|
||||
mcp_auth_header: Final = (
|
||||
request.credentials.get("auth_value")
|
||||
if request.auth_type in _STAGED_AUTH_VALUE_AUTH_TYPES and isinstance(request.credentials, dict)
|
||||
else None
|
||||
)
|
||||
# Authorization doubles as the admission fallback (LITELLM_API_KEY_HEADER_NAME_SECONDARY):
|
||||
# when the primary x-litellm-api-key header is absent, the Authorization value is the
|
||||
# caller's LiteLLM key, not an upstream token, and must never be forwarded upstream.
|
||||
oauth2_headers: Final = (
|
||||
MCPRequestHandler._get_oauth2_headers_from_headers(headers)
|
||||
if request.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES
|
||||
and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY)
|
||||
else None
|
||||
)
|
||||
return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers)
|
||||
|
||||
async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None:
|
||||
with anyio.move_on_after(deadline):
|
||||
return await client.list_tools(raise_on_error=True)
|
||||
|
|
@ -1387,6 +1428,8 @@ if MCP_AVAILABLE:
|
|||
},
|
||||
)
|
||||
|
||||
staged: Final = _stage_server_test(new_mcp_server_request, request.headers)
|
||||
|
||||
async def _test_connection_operation(client):
|
||||
async def _noop(session):
|
||||
return "ok"
|
||||
|
|
@ -1395,8 +1438,10 @@ if MCP_AVAILABLE:
|
|||
return {"status": "ok"}
|
||||
|
||||
return await _execute_with_mcp_client(
|
||||
new_mcp_server_request,
|
||||
staged.request,
|
||||
_test_connection_operation,
|
||||
mcp_auth_header=staged.mcp_auth_header,
|
||||
oauth2_headers=staged.oauth2_headers,
|
||||
raw_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
|
||||
|
|
@ -1417,37 +1462,11 @@ if MCP_AVAILABLE:
|
|||
},
|
||||
)
|
||||
|
||||
new_mcp_server_request = _inherit_credentials_from_existing_server(new_mcp_server_request)
|
||||
staged: Final = _stage_server_test(new_mcp_server_request, request.headers)
|
||||
|
||||
# For OpenAPI spec servers, generate tools from the spec directly
|
||||
if new_mcp_server_request.spec_path:
|
||||
return await _preview_openapi_tools(new_mcp_server_request.spec_path)
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
|
||||
headers: Final = request.headers
|
||||
|
||||
mcp_auth_header: str | None = None
|
||||
if new_mcp_server_request.auth_type in {
|
||||
MCPAuth.api_key,
|
||||
MCPAuth.bearer_token,
|
||||
MCPAuth.basic,
|
||||
MCPAuth.authorization,
|
||||
}:
|
||||
credentials: Final = getattr(new_mcp_server_request, "credentials", None)
|
||||
if isinstance(credentials, dict):
|
||||
mcp_auth_header = credentials.get("auth_value")
|
||||
|
||||
# Authorization doubles as the admission fallback (LITELLM_API_KEY_HEADER_NAME_SECONDARY):
|
||||
# when the primary x-litellm-api-key header is absent, the Authorization value is the
|
||||
# caller's LiteLLM key, not an upstream token, and must never be forwarded upstream.
|
||||
oauth2_headers: dict[str, str] | None = None
|
||||
if new_mcp_server_request.auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES and headers.get(
|
||||
MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY
|
||||
):
|
||||
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
|
||||
if staged.request.spec_path:
|
||||
return await _preview_openapi_tools(staged.request.spec_path)
|
||||
|
||||
async def _list_tools_operation(client):
|
||||
# Bound the whole pagination walk: without this the preview is limited only by the
|
||||
|
|
@ -1478,9 +1497,9 @@ if MCP_AVAILABLE:
|
|||
}
|
||||
|
||||
return await _execute_with_mcp_client(
|
||||
new_mcp_server_request,
|
||||
staged.request,
|
||||
_list_tools_operation,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
oauth2_headers=oauth2_headers,
|
||||
mcp_auth_header=staged.mcp_auth_header,
|
||||
oauth2_headers=staged.oauth2_headers,
|
||||
raw_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -602,6 +602,140 @@ class TestTestConnection:
|
|||
route = _get_route("/mcp-rest/test/connection", "POST")
|
||||
assert _route_has_dependency(route, user_api_key_auth)
|
||||
|
||||
@staticmethod
|
||||
def _capture_execute(monkeypatch) -> dict:
|
||||
captured: dict = {}
|
||||
|
||||
async def fake_execute(
|
||||
request,
|
||||
operation,
|
||||
mcp_auth_header=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
captured["request"] = request
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
captured["oauth2_headers"] = oauth2_headers
|
||||
return {"status": "ok"}
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False)
|
||||
return captured
|
||||
|
||||
@staticmethod
|
||||
def _oauth2_authorization_code_payload(**overrides) -> NewMCPServerRequest:
|
||||
return NewMCPServerRequest(
|
||||
server_name="github_mcp",
|
||||
url="https://api.githubcopilot.com/mcp/",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
oauth2_flow="authorization_code",
|
||||
authorization_url="https://github.com/login/oauth/authorize",
|
||||
token_url="https://github.com/login/oauth/access_token",
|
||||
**overrides,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_staged_oauth2_bearer(self, monkeypatch):
|
||||
"""The just-authorized upstream token rides the request's Authorization header, exactly
|
||||
as /test/tools/list receives it; dropping it makes every authorization_code server fail
|
||||
the connection test that its tools preview passes."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
captured = self._capture_execute(monkeypatch)
|
||||
request = _build_request(
|
||||
{"x-litellm-api-key": "sk-admin-session", "authorization": "Bearer upstream-oauth-token"},
|
||||
path="/mcp-rest/test/connection",
|
||||
)
|
||||
|
||||
result = await rest_endpoints.test_connection(
|
||||
request,
|
||||
self._oauth2_authorization_code_payload(),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result == {"status": "ok"}
|
||||
assert captured["oauth2_headers"] == {"Authorization": "Bearer upstream-oauth-token"}
|
||||
assert captured["mcp_auth_header"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_staged_auth_value(self, monkeypatch):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
captured = self._capture_execute(monkeypatch)
|
||||
request = _build_request({"x-litellm-api-key": "sk-admin-session"}, path="/mcp-rest/test/connection")
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://example.com/mcp",
|
||||
auth_type=MCPAuth.bearer_token,
|
||||
credentials={"auth_value": "upstream-static-token"},
|
||||
)
|
||||
|
||||
result = await rest_endpoints.test_connection(
|
||||
request,
|
||||
payload,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result == {"status": "ok"}
|
||||
assert captured["mcp_auth_header"] == "upstream-static-token"
|
||||
assert captured["oauth2_headers"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inherits_stored_credentials_of_saved_server(self, monkeypatch):
|
||||
"""The edit form resends a saved server without its masked credential; the stored one
|
||||
must be used, as /test/tools/list already does."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
captured = self._capture_execute(monkeypatch)
|
||||
saved = MCPServer(
|
||||
server_id="saved-server-id",
|
||||
name="example",
|
||||
url="https://example.com/mcp",
|
||||
transport="http",
|
||||
auth_type=MCPAuth.bearer_token,
|
||||
authentication_token="stored-upstream-token",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda server_id: saved if server_id == "saved-server-id" else None,
|
||||
)
|
||||
request = _build_request({"x-litellm-api-key": "sk-admin-session"}, path="/mcp-rest/test/connection")
|
||||
payload = NewMCPServerRequest(
|
||||
server_id="saved-server-id",
|
||||
server_name="example",
|
||||
url="https://example.com/mcp",
|
||||
auth_type=MCPAuth.bearer_token,
|
||||
)
|
||||
|
||||
result = await rest_endpoints.test_connection(
|
||||
request,
|
||||
payload,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result == {"status": "ok"}
|
||||
assert captured["mcp_auth_header"] == "stored-upstream-token"
|
||||
assert captured["request"].credentials == {"auth_value": "stored-upstream-token"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_does_not_forward_authorization_that_satisfied_admission(self, monkeypatch):
|
||||
"""With no x-litellm-api-key, the Authorization value is the caller's LiteLLM key and
|
||||
must never reach the upstream."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
captured = self._capture_execute(monkeypatch)
|
||||
request = _build_request({"authorization": "Bearer sk-litellm-admission-key"}, path="/mcp-rest/test/connection")
|
||||
|
||||
result = await rest_endpoints.test_connection(
|
||||
request,
|
||||
self._oauth2_authorization_code_payload(),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result == {"status": "ok"}
|
||||
assert captured["oauth2_headers"] is None
|
||||
|
||||
|
||||
class TestTestToolsList:
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue