mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): challenge missing upstream OAuth during initialization
This commit is contained in:
parent
97a662f2b2
commit
cb6050b521
3 changed files with 226 additions and 11 deletions
|
|
@ -34,6 +34,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
_gateway_dcr_challenge,
|
||||
_is_mcp_admitted_user_subject,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.client_allowlist import (
|
||||
|
|
@ -1698,7 +1699,42 @@ if MCP_AVAILABLE:
|
|||
excludes a passthrough server is not pushed into an OAuth flow for
|
||||
a server it will be 403'd on immediately after authentication.
|
||||
"""
|
||||
for server_name in mcp_servers or []:
|
||||
if mcp_servers is None:
|
||||
allowed: Final = await operations._get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=None, client_ip=client_ip
|
||||
)
|
||||
eligible: Final = tuple(
|
||||
server for server in allowed if allowed_server_ids is None or server.server_id in allowed_server_ids
|
||||
)
|
||||
results: Final = await asyncio.gather(
|
||||
*(
|
||||
_raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope=scope,
|
||||
mcp_servers=[server.alias or server.server_name or server.name],
|
||||
oauth2_headers=oauth2_headers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip=client_ip,
|
||||
allowed_server_ids=allowed_server_ids,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
for server in eligible
|
||||
),
|
||||
return_exceptions=True,
|
||||
)
|
||||
for result in results:
|
||||
if isinstance(result, asyncio.CancelledError):
|
||||
raise result
|
||||
if results and all(isinstance(result, HTTPException) and result.status_code == 401 for result in results):
|
||||
if all(server.is_gateway_managed_oauth2 for server in eligible):
|
||||
raise _gateway_dcr_challenge(
|
||||
StarletteRequest(scope), get_route_relative_request_path(scope), None, invalid_token=False
|
||||
)
|
||||
first: Final = results[0]
|
||||
if isinstance(first, HTTPException):
|
||||
raise first
|
||||
return
|
||||
for server_name in mcp_servers:
|
||||
server = operations.global_mcp_server_manager.get_mcp_server_by_name(server_name, client_ip=client_ip)
|
||||
if server is not None and allowed_server_ids is not None and server.server_id not in allowed_server_ids:
|
||||
# Caller's narrowed scope excludes this server — skip the
|
||||
|
|
@ -2114,16 +2150,17 @@ if MCP_AVAILABLE:
|
|||
# from the fully-authorized server set: a passthrough server that
|
||||
# the active toolset excludes should not trigger an OAuth flow
|
||||
# for a server the caller will be 403'd on after authentication.
|
||||
await _raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope=scope,
|
||||
mcp_servers=mcp_servers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip=_client_ip,
|
||||
allowed_server_ids=toolset_allowed_server_ids,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
if mcp_servers is not None:
|
||||
await _raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope=scope,
|
||||
mcp_servers=mcp_servers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip=_client_ip,
|
||||
allowed_server_ids=toolset_allowed_server_ids,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
# Pre-flight auth check for pass-through servers. Must run after
|
||||
# toolset scoping so the probe list is derived from the fully-authorized
|
||||
|
|
@ -2200,6 +2237,18 @@ if MCP_AVAILABLE:
|
|||
consumed_messages, body = await _read_request_body_for_routing(receive)
|
||||
is_initialize = _is_initialize_request(body)
|
||||
|
||||
if is_initialize and mcp_servers is None:
|
||||
await _raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope=scope,
|
||||
mcp_servers=None,
|
||||
oauth2_headers=oauth2_headers,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip=_client_ip,
|
||||
allowed_server_ids=toolset_allowed_server_ids,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
use_stateful: Final = bool(session_id or is_initialize)
|
||||
target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless
|
||||
|
||||
|
|
|
|||
|
|
@ -854,3 +854,71 @@ async def test_tool_handler_preserves_unrelated_protocol_errors(
|
|||
await server.mcp_server_tool_call(_mcp_request_ctx(), CallToolRequestParams(name="example", arguments={}))
|
||||
assert caught.value is failure
|
||||
execute.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ("/mcp", "/github/mcp"))
|
||||
@pytest.mark.parametrize("has_token", (False, True))
|
||||
@pytest.mark.parametrize("healthy_companion", (False, True))
|
||||
async def test_initialize_challenges_missing_upstream_credentials_before_creating_session(
|
||||
monkeypatch: pytest.MonkeyPatch, path: str, has_token: bool, healthy_companion: bool
|
||||
) -> None:
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
github: Final = MCPServer(
|
||||
server_id="github-id", name="github", alias="github", server_name="github",
|
||||
url="https://github.example/mcp", transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code",
|
||||
authorization_url="https://github.example/authorize", token_url="https://github.example/token",
|
||||
client_id="registered-client",
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(api_key="test-owner", user_id="test-user")
|
||||
selected: Final = ["github"] if path != "/mcp" else None
|
||||
monkeypatch.setattr(
|
||||
server, "extract_mcp_auth_context", AsyncMock(return_value=(auth, None, selected, None, None, None))
|
||||
)
|
||||
manager: Final = server.operations.global_mcp_server_manager
|
||||
public: Final = MCPServer(
|
||||
server_id="public-id", name="public", alias="public", server_name="public",
|
||||
url="https://public.example/mcp", transport=MCPTransport.http, auth_type=MCPAuth.none,
|
||||
)
|
||||
eligible: Final = [github, public] if healthy_companion and selected is None else [github]
|
||||
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in eligible if s.alias == name))
|
||||
monkeypatch.setattr(manager, "has_user_oauth_token", AsyncMock(return_value=has_token))
|
||||
monkeypatch.setattr(manager, "_ensure_upstream_initialize_instructions_cached", AsyncMock())
|
||||
monkeypatch.setattr(server.operations, "_get_allowed_mcp_servers", AsyncMock(return_value=eligible))
|
||||
monkeypatch.setattr(server, "_check_passthrough_upstream_auth", AsyncMock())
|
||||
monkeypatch.setattr(
|
||||
server, "session_manager_stateful",
|
||||
StreamableHTTPSessionManager(app=server.server, stateless=False, json_response=True),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
server, "session_manager_stateless",
|
||||
StreamableHTTPSessionManager(app=server.server, stateless=True, json_response=True),
|
||||
)
|
||||
await server.initialize_session_managers()
|
||||
try:
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=server.app), base_url="http://gateway") as client:
|
||||
response: Final = await client.post(
|
||||
path, headers={"accept": "application/json, text/event-stream"},
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {
|
||||
"protocolVersion": "2025-06-18", "capabilities": {},
|
||||
"clientInfo": {"name": "test-client", "version": "1"},
|
||||
}},
|
||||
)
|
||||
can_initialize: Final = has_token or (healthy_companion and selected is None)
|
||||
assert response.status_code == (200 if can_initialize else 401), response.text
|
||||
if can_initialize:
|
||||
assert response.json()["result"]["serverInfo"]["name"]
|
||||
assert response.headers["mcp-session-id"]
|
||||
else:
|
||||
assert response.headers["www-authenticate"].startswith("Bearer ")
|
||||
if path == "/mcp":
|
||||
assert response.headers["www-authenticate"] == (
|
||||
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp"'
|
||||
)
|
||||
assert "mcp-session-id" not in response.headers
|
||||
finally:
|
||||
await server.shutdown_session_managers()
|
||||
|
|
|
|||
|
|
@ -10690,3 +10690,101 @@ async def test_discovery_adapter_preserves_authenticated_context(_mcp_request_ct
|
|||
context = dispatched.await_args.args[1]
|
||||
assert context.user_api_key_auth.user_id == "discover-caller"
|
||||
assert context.mcp_servers == ("allowed",)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"token_states, allowed_ids, expected_status",
|
||||
(
|
||||
((False,), None, 401),
|
||||
((False, False), None, 401),
|
||||
((False, True), None, None),
|
||||
((True, False), None, None),
|
||||
((False, False), {"server-1"}, 401),
|
||||
((False, True), {"server-1"}, None),
|
||||
((False,), set(), None),
|
||||
((), None, None),
|
||||
),
|
||||
)
|
||||
async def test_unified_preflight_challenges_only_when_all_authorized_servers_need_oauth(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
token_states: tuple[bool, ...],
|
||||
allowed_ids: set[str] | None,
|
||||
expected_status: int | None,
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
servers: Final = tuple(
|
||||
_make_oauth2_server(f"server-{index}").model_copy(update={"server_id": f"server-{index}"})
|
||||
for index in range(len(token_states))
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(api_key="test-key", user_id="reader")
|
||||
lookup: Final = AsyncMock(return_value=servers)
|
||||
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", lookup)
|
||||
manager: Final = mcp_operations.global_mcp_server_manager
|
||||
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda name, **kwargs: next(s for s in servers if s.alias == name))
|
||||
tokens: Final = AsyncMock(side_effect=lambda s, user: token_states[int(s.server_id.rsplit("-", 1)[1])])
|
||||
monkeypatch.setattr(manager, "has_user_oauth_token", tokens)
|
||||
scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp", "headers": [(b"host", b"gateway")]}
|
||||
request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope=scope, mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None,
|
||||
user_api_key_auth=auth, client_ip="127.0.0.1", allowed_server_ids=allowed_ids,
|
||||
)
|
||||
if expected_status is not None:
|
||||
with pytest.raises(HTTPException) as caught:
|
||||
await request
|
||||
assert caught.value.status_code == expected_status
|
||||
assert "www-authenticate" in {key.lower() for key in (caught.value.headers or {})}
|
||||
else:
|
||||
await request
|
||||
lookup.assert_awaited_once_with(user_api_key_auth=auth, mcp_servers=None, client_ip="127.0.0.1")
|
||||
assert {call.args[0].server_id for call in tokens.await_args_list} == {
|
||||
s.server_id for s in servers if allowed_ids is None or s.server_id in allowed_ids
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure", (HTTPException(status_code=503, detail="unavailable"), asyncio.CancelledError()))
|
||||
async def test_unified_preflight_does_not_misclassify_discovery_failure_as_oauth(
|
||||
monkeypatch: pytest.MonkeyPatch, failure: BaseException
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
upstream: Final = _make_oauth2_server("unavailable")
|
||||
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream]))
|
||||
manager: Final = mcp_operations.global_mcp_server_manager
|
||||
monkeypatch.setattr(manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream)
|
||||
discovery: Final = AsyncMock(side_effect=failure)
|
||||
monkeypatch.setattr(manager, "ensure_oauth_metadata_discovered", discovery)
|
||||
request: Final = server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "path": "/mcp", "method": "POST", "headers": []},
|
||||
mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(user_id="reader"), client_ip=None,
|
||||
)
|
||||
if isinstance(failure, asyncio.CancelledError):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await request
|
||||
else:
|
||||
await request
|
||||
discovery.assert_awaited_once_with(upstream)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unified_preflight_preserves_delegated_oauth_challenge(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
upstream: Final = _make_oauth2_server("delegated", delegate_auth_to_upstream=True)
|
||||
monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream]))
|
||||
monkeypatch.setattr(
|
||||
mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", lambda *args, **kwargs: upstream
|
||||
)
|
||||
with pytest.raises(HTTPException) as caught:
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "path": "/mcp", "method": "POST", "headers": [(b"host", b"gateway")]},
|
||||
mcp_servers=None, oauth2_headers=None, mcp_server_auth_headers=None,
|
||||
user_api_key_auth=UserAPIKeyAuth(user_id="reader"), client_ip=None,
|
||||
)
|
||||
assert caught.value.status_code == 401
|
||||
assert (caught.value.headers or {})["www-authenticate"] == (
|
||||
'Bearer resource_metadata="http://gateway/.well-known/oauth-protected-resource/mcp/delegated"'
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue