mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(mcp): recognize per-server auth header at connect and stop persisting browser-authorize tokens
Two correctness fixes for the client-forwarded token modes.
The preemptive-401 connect gate for true_passthrough and oauth_delegate
only inspected the request-wide Authorization, so a caller who bound the
upstream token via the per-server x-mcp-{alias}-authorization header (the
mandatory shape in a multi-server aggregate, where the request-wide
Authorization is withheld) was spuriously 401'd at connect even though
egress already honors that header. The gate now recognizes the per-server
header for both modes via a shared helper, mode-correctly: true_passthrough
treats any Authorization or the per-server header as the upstream token,
oauth_delegate keeps requiring a distinct x-litellm-api-key so a lone
Authorization consumed for admission is never mistaken for an upstream
token. The preemptive raise is also gated to single-server scopes so a
multi-server aggregate degrades gracefully (the listing absorbs a
per-server failure) instead of one missing token 401-ing the whole connect.
The browser-only Authorize flow was writing the upstream access and refresh
token to LiteLLM_MCPUserCredentials, contradicting the modes' persist-nothing
contract: the temp OAuth-relay server was cached with a hardcoded oauth2
auth_type, so needs_user_oauth_token was true and the token exchange stored
it. The create and edit forms now send the real auth_type for these modes,
so the temp server is not oauth2, needs_user_oauth_token is false, and the
exchange skips storage while still returning the token to the browser
session.
This commit is contained in:
parent
74a15c21ae
commit
98818df418
3 changed files with 87 additions and 2 deletions
|
|
@ -6,6 +6,8 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
|
||||
# Fixture to mock IP address check for all MCP tests
|
||||
# This prevents tests from failing due to IP-based access control
|
||||
|
|
@ -3403,6 +3405,83 @@ async def test_token_exchange_passes_through_upstream_expires_in():
|
|||
assert body["expires_in"] == 43200
|
||||
|
||||
|
||||
async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool:
|
||||
"""Run exchange_token_with_server for a server of ``auth_type`` and report whether it attempted
|
||||
to persist the exchanged token server-side. The client-forwarded token modes must not persist:
|
||||
their contract is that the upstream token stays browser-held, minted/stored/refreshed nowhere."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="t",
|
||||
name="t",
|
||||
server_name="t",
|
||||
alias="t",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=auth_type,
|
||||
client_id="cid",
|
||||
client_secret="cs",
|
||||
authorization_url="https://provider.com/oauth/authorize",
|
||||
token_url="https://provider.com/oauth/token",
|
||||
)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
fake_http_response = MagicMock()
|
||||
fake_http_response.json.return_value = {"access_token": "tok", "refresh_token": "r", "token_type": "Bearer"}
|
||||
fake_http_response.raise_for_status = MagicMock()
|
||||
fake_http_client = MagicMock()
|
||||
fake_http_client.post = AsyncMock(return_value=fake_http_response)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
|
||||
return_value=fake_http_client,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request",
|
||||
new_callable=AsyncMock,
|
||||
return_value="admin-user",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._store_per_user_token_server_side",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_store,
|
||||
):
|
||||
await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type="authorization_code",
|
||||
code="c",
|
||||
redirect_uri="http://127.0.0.1:3000/cb",
|
||||
client_id="cid",
|
||||
client_secret=None,
|
||||
code_verifier=None,
|
||||
)
|
||||
return mock_store.await_count > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
|
||||
async def test_token_exchange_does_not_persist_for_client_forwarded_modes(auth_type):
|
||||
"""The browser-only Authorize for true_passthrough / oauth_delegate must not write the upstream
|
||||
token to the DB: these modes forward a browser-held token and persist nothing server-side."""
|
||||
assert await _exchange_persistence_attempted_for_auth_type(auth_type) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_persists_for_oauth2():
|
||||
"""Guard the test's own discriminator: a genuine oauth2 (authorization_code) server DOES persist,
|
||||
so the passthrough no-persist assertion above is meaningful and not vacuously true."""
|
||||
assert await _exchange_persistence_attempted_for_auth_type(MCPAuth.oauth2) is True
|
||||
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# OBO (token_exchange) Protected Resource Metadata: discovery must name the
|
||||
# JWT-auth issuer the client SSOs with, not the gateway.
|
||||
|
|
|
|||
|
|
@ -184,7 +184,10 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
|
|||
description: values.description,
|
||||
url,
|
||||
transport: transport === TRANSPORT.OPENAPI ? "http" : transport,
|
||||
auth_type: AUTH_TYPE.OAUTH2,
|
||||
auth_type:
|
||||
values.auth_type === AUTH_TYPE.TRUE_PASSTHROUGH || values.auth_type === AUTH_TYPE.OAUTH_DELEGATE
|
||||
? values.auth_type
|
||||
: AUTH_TYPE.OAUTH2,
|
||||
credentials: values.credentials,
|
||||
authorization_url: values.authorization_url,
|
||||
token_url: values.token_url,
|
||||
|
|
|
|||
|
|
@ -163,7 +163,10 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
|
|||
description: values.description || mcpServer.description,
|
||||
url,
|
||||
transport,
|
||||
auth_type: AUTH_TYPE.OAUTH2,
|
||||
auth_type:
|
||||
mcpServer.auth_type === AUTH_TYPE.TRUE_PASSTHROUGH || mcpServer.auth_type === AUTH_TYPE.OAUTH_DELEGATE
|
||||
? mcpServer.auth_type
|
||||
: AUTH_TYPE.OAUTH2,
|
||||
credentials: values.credentials,
|
||||
mcp_access_groups: values.mcp_access_groups || mcpServer.mcp_access_groups,
|
||||
static_headers: staticHeaders,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue