mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(mcp): preserve safe OAuth retries and recovery challenges
This commit is contained in:
parent
8136a67410
commit
a7024590f8
6 changed files with 159 additions and 46 deletions
|
|
@ -462,7 +462,8 @@ class MCPRequestHandler:
|
|||
|
||||
request_route: Final = get_request_route(request)
|
||||
# Only OAuth metadata routes registered under /.well-known/ are public.
|
||||
if request_route.startswith("/.well-known/"):
|
||||
is_public_metadata: Final = request_route.startswith("/.well-known/")
|
||||
if is_public_metadata:
|
||||
validated_user_api_key_auth = UserAPIKeyAuth()
|
||||
elif has_explicit_litellm_key:
|
||||
# An explicit x-litellm-api-key is always a LiteLLM credential, even
|
||||
|
|
@ -552,7 +553,7 @@ class MCPRequestHandler:
|
|||
|
||||
scope.pop(CONNECTION_SCOPE_KEY, None)
|
||||
connection_header: Final = headers.get("authorization")
|
||||
if is_connection_credential(connection_header):
|
||||
if is_connection_credential(connection_header) and not is_public_metadata:
|
||||
scope[CONNECTION_SCOPE_KEY] = await MCPRequestHandler._admit_connection_credential(
|
||||
request=request,
|
||||
request_route=request_route,
|
||||
|
|
@ -627,10 +628,14 @@ class MCPRequestHandler:
|
|||
resource=f"{get_request_base_url(request)}/mcp",
|
||||
)
|
||||
connection: Final = open_connection_credential(connection_header)
|
||||
if connection is None:
|
||||
if connection is None or connection.binding != expected_binding:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid or expired MCP connection credential",
|
||||
detail=(
|
||||
"Invalid or expired MCP connection credential"
|
||||
if connection is None
|
||||
else "Connection credential belongs to a different key or resource"
|
||||
),
|
||||
headers=MappingProxyType(
|
||||
{
|
||||
"www-authenticate": connection_challenge(request, expected_binding),
|
||||
|
|
@ -638,8 +643,6 @@ class MCPRequestHandler:
|
|||
}
|
||||
),
|
||||
)
|
||||
if connection.binding != expected_binding:
|
||||
raise HTTPException(status_code=401, detail="Connection credential belongs to a different key or resource")
|
||||
return connection
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1216,7 +1219,13 @@ class MCPRequestHandler:
|
|||
raise HTTPException(status_code=401, detail="Invalid or expired credential")
|
||||
|
||||
@staticmethod
|
||||
async def _enforce_admitted_live_policy(admitted: UserAPIKeyAuth, request: Request, route: str) -> None:
|
||||
async def _enforce_admitted_live_policy(
|
||||
admitted: UserAPIKeyAuth,
|
||||
request: Request,
|
||||
route: str,
|
||||
*,
|
||||
request_data: dict[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Run the standard pipeline's authorization checks over the admitted identity.
|
||||
|
||||
Mirrors the ``user_api_key_auth`` wrapper between the builder and its return: clear the
|
||||
|
|
@ -1246,7 +1255,7 @@ class MCPRequestHandler:
|
|||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=admitted,
|
||||
request=request,
|
||||
request_data=await _read_request_body(request=request),
|
||||
request_data=await _read_request_body(request=request) if request_data is None else request_data,
|
||||
route=route,
|
||||
)
|
||||
except (HTTPException, ProxyException):
|
||||
|
|
|
|||
|
|
@ -1041,6 +1041,7 @@ async def exchange_token_with_server(
|
|||
scope: str | None = None,
|
||||
client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
|
||||
connection_binding: ConnectionBinding | None = None,
|
||||
connection_claim: tuple[str, int] | None = None,
|
||||
):
|
||||
_raise_if_not_oauth2(mcp_server)
|
||||
if grant_type not in ("authorization_code", "refresh_token"):
|
||||
|
|
@ -1197,6 +1198,10 @@ async def exchange_token_with_server(
|
|||
)
|
||||
|
||||
async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
if connection_claim is not None:
|
||||
claimed: Final = await claim_connection_once(*connection_claim)
|
||||
if claimed is not None:
|
||||
return claimed
|
||||
try:
|
||||
response: Final = await async_client.post(
|
||||
token_url,
|
||||
|
|
@ -3092,11 +3097,8 @@ async def complete_connection(request: Request, flow: str = Form(...), decision:
|
|||
if opened is None or opened.jti != flow or decision not in ("approve", "deny"):
|
||||
return _oauth_error(400, "invalid_request", "Invalid or expired consent")
|
||||
server: Final = await validate_connection_binding(request, opened.binding)
|
||||
claim: Final = await claim_connection_once(f"consent:{opened.jti}", opened.exp)
|
||||
if claim is not None:
|
||||
return claim
|
||||
if decision == "deny":
|
||||
denied: Final = RedirectResponse(
|
||||
response: Final = (
|
||||
RedirectResponse(
|
||||
append_connection_query(
|
||||
opened.redirect_uri,
|
||||
(
|
||||
|
|
@ -3106,21 +3108,25 @@ async def complete_connection(request: Request, flow: str = Form(...), decision:
|
|||
),
|
||||
status_code=302,
|
||||
)
|
||||
cookie_path, _ = _cookie_path_and_secure(request)
|
||||
denied.delete_cookie(cookie_name, path=cookie_path)
|
||||
return denied
|
||||
response: Final = await authorize_with_server(
|
||||
request,
|
||||
server,
|
||||
opened.client_id,
|
||||
opened.redirect_uri,
|
||||
opened.state,
|
||||
opened.code_challenge,
|
||||
"S256",
|
||||
"code",
|
||||
opened.scope,
|
||||
connection=opened,
|
||||
if decision == "deny"
|
||||
else await authorize_with_server(
|
||||
request,
|
||||
server,
|
||||
opened.client_id,
|
||||
opened.redirect_uri,
|
||||
opened.state,
|
||||
opened.code_challenge,
|
||||
"S256",
|
||||
"code",
|
||||
opened.scope,
|
||||
connection=opened,
|
||||
)
|
||||
)
|
||||
if response.status_code >= 400:
|
||||
return response
|
||||
claim: Final = await claim_connection_once(f"consent:{opened.jti}", opened.exp)
|
||||
if claim is not None:
|
||||
return claim
|
||||
cookie_path, _ = _cookie_path_and_secure(request)
|
||||
response.delete_cookie(cookie_name, path=cookie_path)
|
||||
return response
|
||||
|
|
@ -3153,9 +3159,6 @@ async def exchange_connection_token(
|
|||
or not _pkce_verifier_matches(code_verifier, opened.authorization.code_challenge)
|
||||
):
|
||||
return _oauth_error(400, "invalid_grant", "Invalid authorization code or PKCE verifier")
|
||||
claim: Final = await claim_connection_once(f"code:{opened.jti}", opened.exp)
|
||||
if claim is not None:
|
||||
return claim
|
||||
return await exchange_token_with_server(
|
||||
request,
|
||||
server,
|
||||
|
|
@ -3167,6 +3170,7 @@ async def exchange_connection_token(
|
|||
code_verifier,
|
||||
scope=opened.authorization.scope,
|
||||
connection_binding=binding,
|
||||
connection_claim=(f"code:{opened.jti}", opened.exp),
|
||||
)
|
||||
if grant_type == "refresh_token":
|
||||
refreshed: Final = open_connection_credential(refresh_token or "", refresh=True)
|
||||
|
|
@ -3174,9 +3178,6 @@ async def exchange_connection_token(
|
|||
return _oauth_error(400, "invalid_grant", "Invalid refresh credential")
|
||||
if scope and not frozenset(scope.split()).issubset((refreshed.scope or "").split()):
|
||||
return _oauth_error(400, "invalid_scope", "Refresh cannot expand the granted scopes")
|
||||
claimed: Final = await claim_connection_once(f"refresh:{refreshed.jti}", refreshed.exp)
|
||||
if claimed is not None:
|
||||
return claimed
|
||||
return await exchange_token_with_server(
|
||||
request,
|
||||
server,
|
||||
|
|
@ -3189,5 +3190,6 @@ async def exchange_connection_token(
|
|||
refresh_token=refreshed.token.get_secret_value(),
|
||||
scope=scope or refreshed.scope,
|
||||
connection_binding=binding,
|
||||
connection_claim=(f"refresh:{refreshed.jti}", refreshed.exp),
|
||||
)
|
||||
return _oauth_error(400, "unsupported_grant_type", "Unsupported connection grant type")
|
||||
|
|
|
|||
|
|
@ -1632,7 +1632,9 @@ async def validate_connection_binding(request: Request, binding: ConnectionBindi
|
|||
if binding.resource != f"{get_request_base_url(request)}/mcp":
|
||||
raise HTTPException(status_code=400, detail="Invalid connection resource")
|
||||
key: Final = await MCPRequestHandler._reload_admitted_key(binding.key_hash) # pyright: ignore[reportPrivateUsage] # reuse key revocation and SCIM checks
|
||||
await MCPRequestHandler._enforce_admitted_live_policy(key.model_copy(), request, "/mcp") # pyright: ignore[reportPrivateUsage] # enforce the same MCP route and budget policy
|
||||
await MCPRequestHandler._enforce_admitted_live_policy( # pyright: ignore[reportPrivateUsage] # enforce the same MCP route and budget policy
|
||||
key.model_copy(), request, "/mcp", request_data={}
|
||||
)
|
||||
allowed: Final = await MCPRequestHandler.get_allowed_mcp_servers(key)
|
||||
server: Final = global_mcp_server_manager.get_mcp_server_by_id(
|
||||
binding.server_id, client_ip=IPAddressUtils.get_mcp_client_ip(request)
|
||||
|
|
|
|||
|
|
@ -953,7 +953,9 @@ async def test_get_tools_from_mcp_servers():
|
|||
client_ip=None,
|
||||
user_api_key_auth=None,
|
||||
oauth2_headers=None,
|
||||
connection_credential=None,
|
||||
):
|
||||
assert connection_credential is None
|
||||
if server.server_id == "server1_id":
|
||||
return [mock_tool_1]
|
||||
return [mock_tool_2]
|
||||
|
|
|
|||
|
|
@ -1817,7 +1817,8 @@ class TestMCPPublicRouteGuard:
|
|||
await MCPRequestHandler.process_mcp_request(scope)
|
||||
assert exc_info.value.status_code == 401
|
||||
|
||||
async def test_legitimate_well_known_path_still_bypasses_auth(self):
|
||||
@pytest.mark.parametrize("bearer", [None, "llm_caccess_stale", "llm_crefresh_stale"])
|
||||
async def test_legitimate_well_known_path_still_bypasses_auth(self, bearer):
|
||||
"""
|
||||
Real OAuth discovery routes registered under /.well-known/ must remain
|
||||
public so unauthenticated clients can fetch them per RFC 8414/9728.
|
||||
|
|
@ -1826,16 +1827,21 @@ class TestMCPPublicRouteGuard:
|
|||
"type": "http",
|
||||
"method": "GET",
|
||||
"path": "/.well-known/oauth-protected-resource",
|
||||
"headers": [],
|
||||
"headers": [(b"authorization", f"Bearer {bearer}".encode())] if bearer else [],
|
||||
}
|
||||
|
||||
# No mock needed — public path should not call user_api_key_auth at all
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
|
||||
) as mock_auth:
|
||||
auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope)
|
||||
auth_result, mcp_header, _, server_headers, oauth_headers, raw_headers = (
|
||||
await MCPRequestHandler.process_mcp_request(scope)
|
||||
)
|
||||
mock_auth.assert_not_called()
|
||||
assert isinstance(auth_result, UserAPIKeyAuth)
|
||||
assert "litellm.mcp.connection_grant" not in scope
|
||||
assert not mcp_header and not server_headers and not oauth_headers
|
||||
assert "authorization" not in raw_headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -9449,7 +9455,7 @@ class TestScopedSessionAdmission:
|
|||
(None, "connection-target", 401),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("grant_state", ["valid", "expired", "refresh", "oversized"])
|
||||
@pytest.mark.parametrize("grant_state", ["valid", "expired", "refresh", "oversized", "other-resource"])
|
||||
@pytest.mark.parametrize("server_mode", ["managed", "delegated"])
|
||||
async def test_connection_credential_requires_exact_key_and_server(
|
||||
monkeypatch, key, selector, expected, grant_state, server_mode
|
||||
|
|
@ -9485,7 +9491,9 @@ async def test_connection_credential_requires_exact_key_and_server(
|
|||
)
|
||||
monkeypatch.setattr(admission, "user_api_key_auth", AsyncMock(return_value=auth))
|
||||
binding = ConnectionBinding(
|
||||
key_hash=hash_token("sk-original"), server_id=server.server_id, resource="https://gateway.example/mcp"
|
||||
key_hash=hash_token("sk-original"),
|
||||
server_id=server.server_id,
|
||||
resource="https://other.example/mcp" if grant_state == "other-resource" else "https://gateway.example/mcp",
|
||||
)
|
||||
issued = json.loads(
|
||||
flow.mint_connection_tokens(
|
||||
|
|
@ -9528,11 +9536,18 @@ async def test_connection_credential_requires_exact_key_and_server(
|
|||
assert exc.value.status_code == expected_status
|
||||
if (
|
||||
server_mode == "managed"
|
||||
and key == "sk-original"
|
||||
and key is not None
|
||||
and selector == "connection-target"
|
||||
and grant_state != "valid"
|
||||
):
|
||||
assert "resource_metadata=" in exc.value.headers["www-authenticate"]
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
challenge = exc.value.headers["www-authenticate"]
|
||||
metadata_url = challenge.split('resource_metadata="', 1)[1].split('"', 1)[0]
|
||||
bootstrap = parse_qs(urlparse(metadata_url).query)["connection"][0]
|
||||
assert flow.open_connection_bootstrap(bootstrap) == ConnectionBinding(
|
||||
key_hash=hash_token(key), server_id=server.server_id, resource="https://gateway.example/mcp"
|
||||
)
|
||||
assert exc.value.headers["Cache-Control"] == "no-store"
|
||||
assert flow.CONNECTION_SCOPE_KEY not in scope
|
||||
return
|
||||
_, _, _, server_headers, oauth_headers, raw_headers = await MCPRequestHandler.process_mcp_request(scope)
|
||||
|
|
|
|||
|
|
@ -12574,6 +12574,79 @@ def test_keyed_connection_headerless_exchange_and_rotating_refresh(keyed_oauth_c
|
|||
harness.vault.assert_not_called()
|
||||
|
||||
|
||||
def test_keyed_connection_consent_can_retry_failed_preparation(keyed_oauth_client, monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
|
||||
|
||||
harness = keyed_oauth_client
|
||||
_, _, handle = _start_keyed_oauth(harness)
|
||||
prepare = AsyncMock(side_effect=[HTTPException(status_code=503, detail="discovery unavailable"), harness.server])
|
||||
monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", prepare)
|
||||
payload = {"flow": handle, "decision": "approve"}
|
||||
|
||||
failed = harness.client.post("/authorize/connection/complete", data=payload)
|
||||
assert failed.status_code == 503
|
||||
assert "location" not in failed.headers
|
||||
assert not failed.headers.get("set-cookie")
|
||||
retried = harness.client.post("/authorize/connection/complete", data=payload)
|
||||
assert retried.status_code == 303
|
||||
assert retried.headers["location"].startswith("https://provider.example/authorize?")
|
||||
harness.upstream.post.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("grant_type", ["authorization_code", "refresh_token"])
|
||||
def test_keyed_connection_exchange_can_retry_failed_preparation(keyed_oauth_client, monkeypatch, grant_type):
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints as endpoints
|
||||
|
||||
harness = keyed_oauth_client
|
||||
code_payload = _complete_keyed_oauth(harness)
|
||||
tokens = harness.client.post("/token", data=code_payload).json() if grant_type == "refresh_token" else None
|
||||
payload = (
|
||||
{"grant_type": grant_type, "client_id": code_payload["client_id"], "refresh_token": tokens["refresh_token"]}
|
||||
if tokens is not None else code_payload
|
||||
)
|
||||
harness.upstream.post.reset_mock()
|
||||
prepare = AsyncMock(side_effect=[HTTPException(status_code=503, detail="discovery unavailable"), harness.server])
|
||||
monkeypatch.setattr(endpoints, "_server_with_oauth_endpoints", prepare)
|
||||
|
||||
failed = harness.client.post("/token", data=payload)
|
||||
assert failed.status_code == 503
|
||||
harness.upstream.post.assert_not_awaited()
|
||||
retried = harness.client.post("/token", data=payload)
|
||||
assert retried.status_code == 200, retried.text
|
||||
assert retried.json()["access_token"].startswith("llm_caccess_")
|
||||
harness.upstream.post.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("grant_type", ["authorization_code", "refresh_token"])
|
||||
@pytest.mark.parametrize("failure", ["response", "timeout"])
|
||||
def test_keyed_connection_dispatched_provider_failure_cannot_redeem_twice(keyed_oauth_client, grant_type, failure):
|
||||
import httpx
|
||||
|
||||
harness = keyed_oauth_client
|
||||
code_payload = _complete_keyed_oauth(harness)
|
||||
tokens = harness.client.post("/token", data=code_payload).json() if grant_type == "refresh_token" else None
|
||||
payload = (
|
||||
{"grant_type": grant_type, "client_id": code_payload["client_id"], "refresh_token": tokens["refresh_token"]}
|
||||
if tokens is not None else code_payload
|
||||
)
|
||||
harness.upstream.post.reset_mock()
|
||||
harness.upstream.post.return_value = httpx.Response(
|
||||
503, json={"error": "temporarily_unavailable"}, request=httpx.Request("POST", "https://provider.example/token")
|
||||
)
|
||||
|
||||
if failure == "timeout":
|
||||
harness.upstream.post.side_effect = httpx.ReadTimeout("provider response unavailable")
|
||||
with pytest.raises(httpx.ReadTimeout):
|
||||
harness.client.post("/token", data=payload)
|
||||
else:
|
||||
failed = harness.client.post("/token", data=payload)
|
||||
assert failed.status_code >= 500
|
||||
replayed = harness.client.post("/token", data=payload)
|
||||
assert replayed.status_code == 400
|
||||
assert replayed.json()["error"] == "invalid_grant"
|
||||
harness.upstream.post.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("change", ["verifier", "redirect", "resource", "tamper"])
|
||||
def test_keyed_connection_rejects_bad_exchange_before_upstream(keyed_oauth_client, change):
|
||||
harness = keyed_oauth_client
|
||||
|
|
@ -12672,7 +12745,7 @@ def test_keyed_connection_client_cannot_enter_session_login_flow(keyed_oauth_cli
|
|||
|
||||
|
||||
@pytest.mark.parametrize("stage", ["authorize", "exchange", "refresh"])
|
||||
@pytest.mark.parametrize("policy", ["blocked", "expired", "denied_server", "denied_route", "allowed"])
|
||||
@pytest.mark.parametrize("policy", ["blocked", "expired", "denied_server", "denied_route", "over_budget", "allowed"])
|
||||
def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, monkeypatch, stage, policy):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import gateway_dcr_flow as flow
|
||||
|
|
@ -12702,7 +12775,12 @@ def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, m
|
|||
lookup = AsyncMock(return_value=key)
|
||||
monkeypatch.setattr(auth_checks, "get_key_object", lookup)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
|
||||
monkeypatch.setattr(admission, "_run_centralized_common_checks", AsyncMock())
|
||||
import litellm
|
||||
|
||||
common_checks = AsyncMock(
|
||||
side_effect=litellm.BudgetExceededError(current_cost=2, max_budget=1) if policy == "over_budget" else None
|
||||
)
|
||||
monkeypatch.setattr(admission, "_run_centralized_common_checks", common_checks)
|
||||
monkeypatch.setitem(global_mcp_server_manager.registry, harness.server.server_id, harness.server)
|
||||
if stage == "authorize":
|
||||
challenge = urlsafe_b64encode(hashlib.sha256(payload["code_verifier"].encode()).digest()).rstrip(b"=").decode()
|
||||
|
|
@ -12718,25 +12796,30 @@ def test_keyed_connection_reloads_live_key_before_provider(keyed_oauth_client, m
|
|||
},
|
||||
)
|
||||
elif stage == "exchange":
|
||||
response = harness.client.post("/token", data=payload)
|
||||
response = harness.client.post("/token", data={**payload, "model": "untrusted\nforged log entry"})
|
||||
else:
|
||||
assert token is not None and token.status_code == 200
|
||||
response = harness.client.post(
|
||||
"/token",
|
||||
data={
|
||||
"model": "untrusted\nforged log entry",
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": payload["client_id"],
|
||||
"refresh_token": token.json()["refresh_token"],
|
||||
"resource": harness.binding.resource,
|
||||
},
|
||||
)
|
||||
assert response.status_code == (200 if policy == "allowed" else 401 if policy in ("blocked", "expired") else 403), (
|
||||
assert response.status_code == (200 if policy == "allowed" else 401 if policy in ("blocked", "expired") else 422 if policy == "over_budget" else 403), (
|
||||
response.text
|
||||
)
|
||||
lookup.assert_awaited_once()
|
||||
assert lookup.call_args.kwargs["hashed_token"] == harness.binding.key_hash
|
||||
assert harness.upstream.post.call_count == int(policy == "allowed" and stage != "authorize")
|
||||
harness.vault.assert_not_called()
|
||||
if policy not in ("blocked", "expired", "denied_route"):
|
||||
common_checks.assert_awaited_once()
|
||||
assert common_checks.call_args.kwargs["request_data"] == {}
|
||||
assert common_checks.call_args.kwargs["route"] == "/mcp"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue