From 466f06df6dca01a4a9a8d76db49a90bf4e351b11 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 14 May 2026 00:33:36 +0530 Subject: [PATCH 1/9] fix(mcp): surface upstream 401 for token-forwarding MCP servers (#27847) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(mcp): surface upstream 401 for token-forwarding MCP servers For MCP servers configured with extra_headers: [Authorization], the gateway forwards the client token directly to the upstream. When that token is rejected (expired or invalid) the upstream returns 401, but the MCP SDK starts the SSE stream with 200 OK before calling handlers, so the 401 can't be returned mid-stream. Fix: add a pre-flight httpx probe in handle_streamable_http_mcp — before the SDK opens the session — so the gateway can still return HTTP 401 with WWW-Authenticate: Bearer authorization_uri= when the upstream rejects the token. The probe fails-open (returns 200) on network errors so a transient hiccup does not block valid requests. Co-authored-by: Cursor * fix(mcp): parallelize pre-flight auth probes and use HEAD to avoid side effects - Extract forwarded_auth outside the pass-through server loop (was called N times for the same scope value) - Gather all upstream auth probes concurrently with asyncio.gather instead of sequentially; eliminates N×5 s worst-case latency - Switch probe from POST+initialize JSON-RPC body to HEAD request; HEAD carries the Authorization header so the upstream rejects invalid tokens with 401 but never allocates a session or writes an audit entry Co-authored-by: Cursor * fix(mcp): use get_async_httpx_client in _probe_upstream_auth Replaces bare httpx.AsyncClient with the project-standard get_async_httpx_client(httpxSpecialProvider.MCP) to satisfy the ensure_async_clients_test code coverage check and avoid the +500 ms per-request overhead of creating a new client on every probe call. Co-authored-by: Cursor * refactor(mcp): extract pre-flight probe into _check_passthrough_upstream_auth Moves the parallel upstream auth probe logic out of handle_streamable_http_mcp into a dedicated helper to satisfy Ruff PLR0915 (Too many statements > 50). Co-authored-by: Cursor * fix(mcp): gate pre-flight probes on authorized server set to prevent bypass _check_passthrough_upstream_auth was resolving user-supplied server names directly before authorization ran, letting any permitted LiteLLM key trigger an upstream HEAD probe to a server it was not allowed to use. Changes: - Call _get_allowed_mcp_servers inside the helper so only servers the caller's key is authorized for are probed. - Move the call site to after toolset scoping so the auth context is fully resolved before the probe list is built. - Thread user_api_key_auth into the helper signature (replaces the raw mcp_servers name list). Co-authored-by: Cursor * Add async HTTP HEAD support Co-authored-by: Yassin Kortam * fix(mcp): use Scope type annotation in _get_forwarded_auth_from_scope Co-authored-by: Cursor * Fix MCP upstream auth probe method Co-authored-by: Yassin Kortam * Remove unused AsyncHTTPHandler head method Co-authored-by: Yassin Kortam * fix(mcp): exclude has_client_credentials servers from pre-flight auth probe _prepare_mcp_server_headers skips caller Authorization when the server uses OAuth client-credentials (M2M), but the pre-flight probe was still selecting those servers and forwarding the caller's raw token in the HEAD request. Exclude servers with has_client_credentials from the probe list to match the actual downstream header-preparation logic. Co-authored-by: Cursor * fix(mcp): propagate upstream 403 as 403, not 401 with WWW-Authenticate Per RFC 9110, 401 means "go get new credentials." Mapping an upstream 403 to a gateway 401 causes OAuth clients to restart the authorization flow, obtain a fresh token with identical scopes, hit 403 again, and loop indefinitely. 401 from upstream → gateway 401 + WWW-Authenticate (re-authorize) 403 from upstream → gateway 403 (no WWW-Authenticate hint) Co-authored-by: Cursor * fix(mcp): skip auth probe when Authorization may be the LiteLLM proxy key The pre-flight upstream probe must not forward the caller's Authorization header when it could itself be the LiteLLM proxy API key. Restrict the probe to requests that supply x-litellm-api-key explicitly — only then is the Authorization header unambiguously the upstream OAuth token the caller wants forwarded. * Fix MCP ASGI HTTPException propagation Co-authored-by: Yassin Kortam * fix(mcp): use public AsyncHTTPHandler.post() in auth probe Use AsyncHTTPHandler.post() and catch httpx.HTTPStatusError explicitly so the 401/403 we want to surface is not silently swallowed by the broad fail-open except Exception block. Avoids reaching into the handler's private client attribute, which would silently regress to fail-open if AsyncHTTPHandler is ever refactored. * Fix MCP auth probe tests Co-authored-by: Yassin Kortam * test(mcp): add coverage for httpx.HTTPStatusError path in auth probe AsyncHTTPHandler.post() calls raise_for_status() internally, so a real upstream 401/403 lands as httpx.HTTPStatusError. Add a test that exercises that specific exception path so a regression that swallows the error in the broad fail-open except Exception would be caught. --------- Co-authored-by: Cursor Co-authored-by: Yassin Kortam Co-authored-by: claude-bot --- .../proxy/_experimental/mcp_server/server.py | 165 +++++++++++++++++- litellm/proxy/proxy_server.py | 11 +- .../mcp_server/test_mcp_server.py | 134 ++++++++++++++ .../proxy/test_mcp_asgi_response.py | 36 ++++ 4 files changed, 344 insertions(+), 2 deletions(-) create mode 100644 tests/test_litellm/proxy/test_mcp_asgi_response.py diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 276a6e8a3bb..685232d0495 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -23,6 +23,7 @@ from typing import ( cast, ) +import httpx from fastapi import FastAPI, HTTPException from pydantic import AnyUrl, ConfigDict from starlette.requests import Request as StarletteRequest @@ -51,13 +52,17 @@ from litellm.proxy._experimental.mcp_server.utils import ( get_server_prefix, iter_known_server_prefixes, ) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, get_chain_id_from_headers, ) -from litellm.types.mcp import MCPAuth +from litellm.types.mcp import MCPAuth, MCPSpecVersion from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall from litellm.utils import Rules, client, function_setup @@ -2754,6 +2759,157 @@ if MCP_AVAILABLE: ) return user_api_key_auth.model_copy(update={"object_permission": updated_op}) + def _get_forwarded_auth_from_scope(scope: Scope) -> Optional[str]: + """Return the upstream-bound ``Authorization`` header value, or None. + + Only returns the ``Authorization`` header when ``x-litellm-api-key`` is + also present. In that case ``Authorization`` is unambiguously the + upstream token the caller wants forwarded to the MCP server. When + ``x-litellm-api-key`` is absent the ``Authorization`` header may itself + be the LiteLLM proxy API key (backward-compat path in + ``MCPRequestHandler.process_mcp_request``), and forwarding it upstream + would leak the proxy key to a third-party MCP server. + """ + authorization = None + has_litellm_key_header = False + for key, value in scope.get("headers", []): + key_lower = key.lower() + if key_lower == b"authorization": + authorization = value.decode("latin-1") + elif key_lower == b"x-litellm-api-key": + has_litellm_key_header = True + if not has_litellm_key_header: + return None + return authorization + + async def _probe_upstream_auth( + url: str, + auth_header: str, + timeout: float = 5.0, + ) -> tuple: + """JSON-RPC initialize-probe the upstream URL to check whether the token is accepted. + + Uses POST so StreamableHTTP MCP servers run the same auth path as a + real client request. Returns (status_code, www_authenticate). + Fails-open with (200, None) on network errors so a transient hiccup + does not block valid requests. + + Uses the public ``AsyncHTTPHandler.post()`` interface and catches + ``httpx.HTTPStatusError`` separately so the 401/403 we want to surface + is not swallowed by the broad fail-open ``except Exception`` below. + """ + client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.MCP, + params={"timeout": timeout}, + ) + probe_payload = { + "jsonrpc": "2.0", + "id": "litellm-mcp-auth-probe", + "method": "initialize", + "params": { + "protocolVersion": MCPSpecVersion.jun_2025.value, + "capabilities": {}, + "clientInfo": { + "name": "litellm-mcp-auth-probe", + "version": "1.0.0", + }, + }, + } + probe_headers = { + "Authorization": auth_header, + "Accept": "application/json, text/event-stream", + } + try: + resp = await client.post( + url=url, + headers=probe_headers, + json=probe_payload, + timeout=timeout, + ) + return resp.status_code, resp.headers.get("www-authenticate") + except httpx.HTTPStatusError as exc: + # AsyncHTTPHandler.post() calls raise_for_status(); a 401/403 from + # upstream lands here. Return its status so the caller can map it + # to the appropriate response. + return exc.response.status_code, exc.response.headers.get( + "www-authenticate" + ) + except Exception as exc: + verbose_logger.debug( + f"_probe_upstream_auth: probe to {url} failed ({exc}), allowing request through" + ) + return 200, None + + async def _check_passthrough_upstream_auth( + scope: Scope, + user_api_key_auth: Optional[UserAPIKeyAuth], + mcp_servers: Optional[List[str]], + client_ip: Optional[str], + ) -> None: + """Probe pass-through upstream servers in parallel before the MCP session starts. + + Only servers the caller's key is already authorized to reach are probed — + the list is derived from _get_allowed_mcp_servers so that a user cannot + trigger an upstream probe against a server their key is not permitted for. + + The MCP SDK commits HTTP 200 headers before invoking handlers, so a 401 + can only be returned before that point. This function raises HTTPException(401) + with a WWW-Authenticate header if any upstream rejects the client token. + Fails-open: network errors are logged and the request is allowed through. + """ + forwarded_auth = _get_forwarded_auth_from_scope(scope) + if not forwarded_auth: + return + + # Use the authorized server set, not the raw user-supplied names, so that + # a caller cannot force a probe to a server their key is not allowed to use. + allowed_servers = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + passthrough_servers = [ + srv + for srv in allowed_servers + if srv.extra_headers + and any(h.lower() == "authorization" for h in srv.extra_headers) + # Exclude M2M servers: _prepare_mcp_server_headers skips caller + # Authorization when has_client_credentials is set, so probing + # those with the caller's token would send the wrong credential. + and not srv.has_client_credentials + ] + if not passthrough_servers: + return + + probe_results = await asyncio.gather( + *[ + _probe_upstream_auth(srv.url or "", forwarded_auth) + for srv in passthrough_servers + ] + ) + request = StarletteRequest(scope) + base_url = get_request_base_url(request) + for srv, (probe_status, _) in zip(passthrough_servers, probe_results): + if probe_status == 401: + # Token is missing or expired — direct the client to re-authorize. + authorization_uri = ( + f"Bearer authorization_uri=" + f"{base_url}/.well-known/oauth-authorization-server/{srv.name}" + ) + raise HTTPException( + status_code=401, + detail="Unauthorized", + headers={"WWW-Authenticate": authorization_uri}, + ) + if probe_status == 403: + # Token is valid but the caller lacks permission — do not hint + # at re-authorization (RFC 9110: a fresh token with the same + # scopes would just hit 403 again and loop indefinitely). + raise HTTPException( + status_code=403, + detail="Forbidden", + ) + async def handle_streamable_http_mcp( scope: Scope, receive: Receive, send: Send ) -> None: @@ -2827,6 +2983,13 @@ if MCP_AVAILABLE: user_api_key_auth, active_toolset_id ) + # Pre-flight auth check for pass-through servers. Must run after + # toolset scoping so the probe list is derived from the fully-authorized + # server set, not the raw user-supplied names. + await _check_passthrough_upstream_auth( + scope, user_api_key_auth, mcp_servers, _client_ip + ) + # Inject masked debug headers when client sends x-litellm-mcp-debug: true _debug_headers = MCPDebug.maybe_build_debug_headers( raw_headers=raw_headers, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5851d2550bd..ecef96351ae 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15026,8 +15026,17 @@ async def _stream_mcp_asgi_response( # If the handler task dies (exception or cancellation) without sending the EOF # sentinel, body_iter() would block forever on body_queue.get(). The callback # below guarantees the queue gets unblocked regardless of how the task ends. + # When this happens before response headers, propagate the original exception + # instead of waiting for the header timeout. def _ensure_eof(task: asyncio.Task) -> None: - if task.cancelled() or task.exception() is not None: + if task.cancelled(): + body_queue.put_nowait(None) + return + + task_exception = task.exception() + if task_exception is not None: + if not headers_ready.done(): + headers_ready.set_exception(task_exception) body_queue.put_nowait(None) handler_task.add_done_callback(_ensure_eof) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index a7649502bde..e1eddfc9c7a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -3256,3 +3256,137 @@ async def test_call_tool_empty_extra_headers_returns_none(): ), "P2 API consistency issue: expected None for empty extra_headers, got: " + str( captured_extra_headers ) + + +# --------------------------------------------------------------------------- +# Pre-flight upstream auth check tests +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_probe_upstream_auth_returns_upstream_status(): + """_probe_upstream_auth forwards the status code from the upstream server.""" + from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth + + mock_response = MagicMock() + mock_response.status_code = 401 + mock_response.headers = {"www-authenticate": 'Bearer realm="test"'} + + mock_client = MagicMock() + mock_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.server.get_async_httpx_client", + return_value=mock_client, + ): + status, www_auth = await _probe_upstream_auth( + "http://upstream/mcp", "Bearer some-token" + ) + + assert status == 401 + assert www_auth == 'Bearer realm="test"' + mock_client.post.assert_awaited_once() + _, kwargs = mock_client.post.call_args + assert kwargs["headers"]["Authorization"] == "Bearer some-token" + assert kwargs["json"]["method"] == "initialize" + + +@pytest.mark.asyncio +async def test_probe_upstream_auth_surfaces_httpx_status_error(): + """Probe extracts status + WWW-Authenticate from httpx.HTTPStatusError. + + AsyncHTTPHandler.post() calls raise_for_status() internally, so when the + upstream returns 401/403 the call raises httpx.HTTPStatusError rather than + returning the response. The probe must catch that specifically (before the + fail-open `except Exception`) so the auth check is not silently defeated. + """ + import httpx + + from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth + + mock_response = MagicMock() + mock_response.status_code = 401 + mock_response.headers = {"www-authenticate": 'Bearer realm="test"'} + request = httpx.Request("POST", "http://upstream/mcp") + error = httpx.HTTPStatusError( + message="401 Unauthorized", request=request, response=mock_response + ) + + mock_client = MagicMock() + mock_client.post = AsyncMock(side_effect=error) + + with patch( + "litellm.proxy._experimental.mcp_server.server.get_async_httpx_client", + return_value=mock_client, + ): + status, www_auth = await _probe_upstream_auth( + "http://upstream/mcp", "Bearer some-token" + ) + + assert status == 401 + assert www_auth == 'Bearer realm="test"' + + +@pytest.mark.asyncio +async def test_probe_upstream_auth_fails_open_on_network_error(): + """_probe_upstream_auth returns (200, None) when the network call fails.""" + from litellm.proxy._experimental.mcp_server.server import _probe_upstream_auth + + mock_client = MagicMock() + mock_client.post = AsyncMock(side_effect=Exception("connection refused")) + + with patch( + "litellm.proxy._experimental.mcp_server.server.get_async_httpx_client", + return_value=mock_client, + ): + status, www_auth = await _probe_upstream_auth( + "http://upstream/mcp", "Bearer some-token" + ) + + assert status == 200 + assert www_auth is None + + +def test_get_forwarded_auth_from_scope_extracts_header(): + """Returns Authorization value when x-litellm-api-key is also present.""" + from litellm.proxy._experimental.mcp_server.server import ( + _get_forwarded_auth_from_scope, + ) + + scope = { + "headers": [ + (b"content-type", b"application/json"), + (b"x-litellm-api-key", b"sk-litellm-proxy-key"), + (b"authorization", b"Bearer my-token"), + ] + } + assert _get_forwarded_auth_from_scope(scope) == "Bearer my-token" + + +def test_get_forwarded_auth_from_scope_returns_none_when_missing(): + from litellm.proxy._experimental.mcp_server.server import ( + _get_forwarded_auth_from_scope, + ) + + assert _get_forwarded_auth_from_scope({"headers": []}) is None + + +def test_get_forwarded_auth_from_scope_skips_when_no_litellm_key_header(): + """Skip when ``x-litellm-api-key`` is absent. + + Without ``x-litellm-api-key``, the ``Authorization`` header may itself be + the LiteLLM proxy API key (backward-compat). Forwarding it upstream would + leak the proxy key, so the helper must return None and the probe must + not fire. + """ + from litellm.proxy._experimental.mcp_server.server import ( + _get_forwarded_auth_from_scope, + ) + + scope = { + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer ambiguous-token"), + ] + } + assert _get_forwarded_auth_from_scope(scope) is None diff --git a/tests/test_litellm/proxy/test_mcp_asgi_response.py b/tests/test_litellm/proxy/test_mcp_asgi_response.py new file mode 100644 index 00000000000..d030f65af4b --- /dev/null +++ b/tests/test_litellm/proxy/test_mcp_asgi_response.py @@ -0,0 +1,36 @@ +import asyncio + +import pytest +from fastapi import HTTPException + +from litellm.proxy.proxy_server import _stream_mcp_asgi_response + + +@pytest.mark.asyncio +async def test_stream_mcp_asgi_response_propagates_pre_header_http_exception(): + async def handle_fn(_scope, _receive, _send): + raise HTTPException( + status_code=401, + detail="Unauthorized", + headers={ + "WWW-Authenticate": "Bearer authorization_uri=https://example.test/auth" + }, + ) + + async def receive(): + return {"type": "http.request", "body": b"", "more_body": False} + + with pytest.raises(HTTPException) as exc_info: + await asyncio.wait_for( + _stream_mcp_asgi_response( + handle_fn, + {"type": "http", "method": "POST", "path": "/mcp", "headers": []}, + receive, + ), + timeout=1.0, + ) + + assert exc_info.value.status_code == 401 + assert exc_info.value.headers == { + "WWW-Authenticate": "Bearer authorization_uri=https://example.test/auth" + } From a74e269f7d180b9c31185ed83894ef4840249ba3 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 14 May 2026 00:35:53 +0530 Subject: [PATCH 2/9] fix(cost): align vertex_ai/gemini-embedding-2-preview with Vertex multimodal pricing (#27848) * fix(cost): align vertex_ai/gemini-embedding-2-preview with Vertex multimodal pricing Co-authored-by: Cursor * fix(cost): align vertex_ai/gemini-embedding-2 GA source URL with preview Per Greptile review on #27848: GA entry referenced ai.google.dev while the preview entry was updated to the canonical Vertex AI pricing page. Both share identical pricing values; sync the source URL for consistency. https://claude.ai/code/session_01W8jRwstnmduadGw8Z8egxe --------- Co-authored-by: Cursor Co-authored-by: Claude --- litellm/model_prices_and_context_window_backup.json | 9 ++++++--- model_prices_and_context_window.json | 9 ++++++--- tests/test_litellm/test_utils.py | 3 ++- 3 files changed, 14 insertions(+), 7 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 8df25d2c9b5..61e0dd07968 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -15551,14 +15551,17 @@ "uses_embed_content": true }, "vertex_ai/gemini-embedding-2-preview": { - "input_cost_per_token": 1.5e-07, + "input_cost_per_audio_per_second": 0.00016, + "input_cost_per_image": 0.00012, + "input_cost_per_token": 2e-07, + "input_cost_per_video_per_second": 0.00079, "litellm_provider": "vertex_ai", "max_input_tokens": 8192, "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal", + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, @@ -15573,7 +15576,7 @@ "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal", + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 70b065aa918..e62f8686e9d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -15556,14 +15556,17 @@ "uses_embed_content": true }, "vertex_ai/gemini-embedding-2-preview": { - "input_cost_per_token": 1.5e-07, + "input_cost_per_audio_per_second": 0.00016, + "input_cost_per_image": 0.00012, + "input_cost_per_token": 2e-07, + "input_cost_per_video_per_second": 0.00079, "litellm_provider": "vertex_ai", "max_input_tokens": 8192, "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal", + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, @@ -15578,7 +15581,7 @@ "mode": "embedding", "output_cost_per_token": 0, "output_vector_size": 3072, - "source": "https://ai.google.dev/gemini-api/docs/embeddings#multimodal", + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supports_multimodal": true, "uses_embed_content": true }, diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index d07af922ea6..bc60375f906 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -3002,7 +3002,7 @@ def test_model_info_for_openrouter_kimi_k2_5(): def test_gemini_embedding_2_ga_in_cost_map(): - """GA gemini-embedding-2 entries align with preview multimodal unit pricing.""" + """GA and Vertex preview gemini-embedding-2 entries align with multimodal unit pricing.""" import json from pathlib import Path @@ -3013,6 +3013,7 @@ def test_gemini_embedding_2_ga_in_cost_map(): for key, provider in ( ("gemini/gemini-embedding-2", "gemini"), ("vertex_ai/gemini-embedding-2", "vertex_ai"), + ("vertex_ai/gemini-embedding-2-preview", "vertex_ai"), ("gemini-embedding-2", "vertex_ai-embedding-models"), ): info = model_cost.get(key) From 18f77ff7bcd329d87114a73e71e55b41bd4147a5 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 14 May 2026 00:36:13 +0530 Subject: [PATCH 3/9] feat(mcp): add delegate_auth_to_upstream flag for PKCE passthrough (#27834) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(mcp): add delegate_auth_to_upstream flag for PKCE passthrough Adds an opt-in per-server flag that lets clients (e.g. VS Code) complete PKCE directly with an upstream OAuth2 MCP server, instead of LiteLLM double-gating with its own API-key/SSO check. Only honored when auth_type=oauth2 and the operator explicitly sets the flag; mixed-target or non-oauth2 requests fail closed. - Adds the field to Pydantic models, Prisma schema, and a migration - New MCPRequestHandler._target_servers_delegate_auth_to_upstream gate that runs only when no x-litellm-api-key is present, so authenticated users still get user_id resolution + stored-credential lookup - Anonymous callers now see delegate servers in get_allowed_mcp_servers (scoped to delegate servers only; the upstream still enforces auth) - mcp_management_endpoints: allow anonymous /authorize and /token for delegate servers so VS Code can complete PKCE without a LiteLLM session - UI toggle (shown only for oauth2) + payload/view wiring - Tests covering: oauth2 on/off, non-oauth2 with flag, mixed targets, no resolvable target, explicit key precedence, and 401 emission Co-authored-by: Cursor * Enforce oauth2 for delegated MCP auth bypass Co-authored-by: Yassin Kortam * fix(mcp): close secondary Authorization bypass for delegate servers The delegate-auth bypass gated only on the primary `x-litellm-api-key` header, so a LiteLLM key sent via `Authorization: Bearer sk-...` (the secondary header) was silently dropped — skipping spend tracking and rate limiting. Gate on the resolved litellm_api_key (which considers both headers) so the bypass fires only when neither is present. Also update the existing "Authorization header present" test to reflect that an upstream OAuth token now flows through the existing oauth2 fallback (LiteLLM auth attempt → fail → anonymous), not via the delegate branch. Co-authored-by: Cursor * Avoid duplicate MCP OAuth credential lookup Co-authored-by: Yassin Kortam * fix(mcp): block delegate bypass for M2M and internal-only servers Two security issues flagged in code review: 1. High – client_credentials (M2M) servers must not be delegatable: LiteLLM auto-fetches the upstream token using stored credentials, so allowing anonymous bypass would let any external caller invoke tools authenticated as LiteLLM's service account. Fix: check `server.has_client_credentials` in `_target_servers_delegate_auth_to_upstream`, the anonymous allow-list in `get_allowed_mcp_servers`, and `_mcp_oauth_user_api_key_auth`. 2. Medium – internal-only servers exposed to public internet: The anonymous delegate allow-list was not filtering by `available_on_public_internet`, so external callers with an upstream OAuth token could invoke tools on servers marked internal-only. Fix: add `available_on_public_internet` guard to the anonymous delegate server list in `get_allowed_mcp_servers`. Tests added for both cases. Co-authored-by: Cursor * Require public MCP delegate auth servers Co-authored-by: Yassin Kortam * fix(mcp): align delegate auth path parsing with downstream routing `_extract_target_server_names_from_path` used a naive segments-based split while `server.py::_get_mcp_servers_in_path` uses a regex that allows server names with one embedded slash and comma-separated lists. With the old parser, a request to `/mcp//` was parsed as targeting `` by the auth gate (bypassing LiteLLM auth) while the routing layer parsed it as `/` — when that name did not resolve, the request fell back to the anonymous allow-list, which can include `allow_all_keys` servers that normally require a LiteLLM key. Replace the parser with the same regex logic as `_get_mcp_servers_in_path` so auth gating sees the exact target name(s) downstream routing sees. Add regression tests covering parser parity and the specific extra-path-segment bypass attempt. https://claude.ai/code/session_01SjyPmwfmrq8fveFgw9iHW9 * fix(mcp): close header/path TOCTOU in MCP delegate auth gate `_target_servers_delegate_auth_to_upstream` and `_target_servers_use_oauth2` trusted the `x-mcp-servers` header when present, but `server.py::extract_mcp_auth_context` overrides that header with the path-derived list for `/mcp/...` routes. An attacker could set `x-mcp-servers: ` while pointing the URL path at a non-delegate server, flipping the auth gate without changing the target downstream routing actually uses. Extract a shared `_resolve_target_server_names` helper that mirrors the downstream override (path-derived names for `/mcp/...` routes, header value otherwise). Add regression tests covering the TOCTOU attempt and the helper's path-vs-header precedence. https://claude.ai/code/session_01SjyPmwfmrq8fveFgw9iHW9 * Fix delegated MCP OAuth test mock Co-authored-by: Yassin Kortam * fix(mcp): drop unreachable /{server}/mcp branch in auth path parser `_extract_target_server_names_from_path` also matched the ``/{server_name}/mcp`` form, but the downstream parser ``_get_mcp_servers_in_path`` only handles ``/mcp/...`` — and ``dynamic_mcp_route`` in ``proxy_server`` rewrites ``/{name}/mcp`` to ``/mcp/{name}`` on the scope before the MCP handler runs. Parsing the un-rewritten form on the auth side was therefore unreachable in production, and contradicted the docstring's claim of mirroring the downstream parser — exactly the kind of mismatch that risks a future header/path TOCTOU if any new entry point skips the rewrite. Drop the branch; the canonical ``/mcp/...`` path matches both parsers. Update the regression test to assert the new behavior. https://claude.ai/code/session_01SjyPmwfmrq8fveFgw9iHW9 * Fix MCP path auth target resolution Co-authored-by: Yassin Kortam * fix(mcp): require auth for refresh_token grants on delegate-auth servers `_mcp_oauth_user_api_key_auth` gates the unauthenticated PKCE flow for ``delegate_auth_to_upstream`` servers, but the bypass applied to BOTH ``/authorize`` and ``/token`` regardless of grant type. ``mcp_token`` accepts ``grant_type=refresh_token`` as well as ``authorization_code``, and ``exchange_token_with_server`` attaches the server's stored ``client_secret`` to whatever is forwarded upstream. An unauthenticated caller holding a refresh token issued to that OAuth client could mint fresh upstream access tokens through LiteLLM. Limit the anonymous bypass on ``/token`` to ``grant_type=authorization_code`` (the only grant PKCE actually protects via ``code_verifier``); fall through to normal LiteLLM auth for ``refresh_token`` and any other grant. ``/authorize`` continues to allow anonymous PKCE redirects. https://claude.ai/code/session_01SjyPmwfmrq8fveFgw9iHW9 * fix(ui): clear delegate_auth_to_upstream when switching off oauth2 The ``delegate_auth_to_upstream`` form field is rendered inside an ``isOAuth2 && (...)`` conditional, so the Form.Item unmounts when the user changes ``auth_type`` away from ``oauth2``. The follow-up ``form.setFieldValue("delegate_auth_to_upstream", false)`` runs after the field has already deregistered, so ``onFinish`` receives ``undefined`` and the fallback ``?? mcpServer.delegate_auth_to_upstream`` preserved the old ``true``. The flag then persisted in the database for a non-oauth2 server and silently re-activated if ``auth_type`` was later switched back to ``oauth2``. In the edit payload, force the flag to ``false`` whenever ``auth_type !== oauth2``; only trust the form value (and the existing DB fallback) when the server is actually oauth2. Backend defense-in-depth already ignores the flag for non-oauth2 servers, but the DB state should stay clean too. https://claude.ai/code/session_01SjyPmwfmrq8fveFgw9iHW9 * Fix MCP delegate auth reset on edit Co-authored-by: Yassin Kortam --------- Co-authored-by: Cursor Co-authored-by: Yassin Kortam Co-authored-by: Claude --- .../migration.sql | 2 + .../litellm_proxy_extras/schema.prisma | 1 + .../mcp_server/auth/user_api_key_auth_mcp.py | 158 +++- .../mcp_server/mcp_server_manager.py | 34 + .../proxy/_experimental/mcp_server/server.py | 24 +- litellm/proxy/_types.py | 3 + .../mcp_management_endpoints.py | 53 +- litellm/proxy/schema.prisma | 1 + .../types/mcp_server/mcp_server_manager.py | 6 + schema.prisma | 1 + .../auth/test_user_api_key_auth_mcp.py | 721 ++++++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 45 ++ .../mcp_server/test_mcp_stale_session.py | 195 +++-- .../test_mcp_management_endpoints.py | 102 +++ .../mcp_tools/MCPPermissionManagement.tsx | 41 +- .../mcp_tools/create_mcp_server.tsx | 2 + .../mcp_tools/mcp_server_edit.test.tsx | 41 + .../components/mcp_tools/mcp_server_edit.tsx | 10 + .../components/mcp_tools/mcp_server_view.tsx | 17 + .../src/components/mcp_tools/types.tsx | 1 + .../src/components/networking.tsx | 13 +- .../src/hooks/useMcpOAuthFlow.tsx | 1 + .../src/hooks/useUserMcpOAuthFlow.tsx | 1 + 23 files changed, 1402 insertions(+), 71 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260513120000_add_delegate_auth_to_upstream_to_mcp_servers/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260513120000_add_delegate_auth_to_upstream_to_mcp_servers/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260513120000_add_delegate_auth_to_upstream_to_mcp_servers/migration.sql new file mode 100644 index 00000000000..50a48743901 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260513120000_add_delegate_auth_to_upstream_to_mcp_servers/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "delegate_auth_to_upstream" BOOLEAN NOT NULL DEFAULT false; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 84ce99557e3..b53507abe6a 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -323,6 +323,7 @@ model LiteLLM_MCPServerTable { registration_url String? allow_all_keys Boolean @default(false) available_on_public_internet Boolean @default(true) + delegate_auth_to_upstream Boolean @default(false) is_byok Boolean @default(false) byok_description String[] @default([]) byok_api_key_help_url String? diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a05af66118c..c87e8c414cd 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -1,3 +1,4 @@ +import re from typing import Dict, List, Optional, Set, Tuple, cast from fastapi import HTTPException @@ -122,6 +123,24 @@ class MCPRequestHandler: # cannot be smuggled via query string, hostname, or a deeper URL segment. if request.url.path.startswith("/.well-known/"): validated_user_api_key_auth = UserAPIKeyAuth() + elif ( + not litellm_api_key + and MCPRequestHandler._target_servers_delegate_auth_to_upstream( # noqa: E501 + path=request.url.path, mcp_servers=mcp_servers + ) + ): + # Operator opted this oauth2 server into upstream-delegated auth + # (PKCE passthrough): skip LiteLLM API-key/SSO entirely so the + # client authenticates directly with the upstream MCP server. + # Fires ONLY when neither x-litellm-api-key nor Authorization is + # present. If any LiteLLM key is supplied (primary or secondary + # header), we fall through so user_id is resolved, spend/rate + # limiting apply, and any stored OAuth token can be retrieved + # and forwarded upstream. Gated by + # _target_servers_delegate_auth_to_upstream, which only returns + # True when EVERY target is auth_type=oauth2 AND has the + # delegate_auth_to_upstream flag set — fails closed otherwise. + validated_user_api_key_auth = UserAPIKeyAuth() elif has_explicit_litellm_key: # Explicit x-litellm-api-key provided - always validate normally validated_user_api_key_auth = await user_api_key_auth( @@ -181,23 +200,62 @@ class MCPRequestHandler: @staticmethod def _extract_target_server_names_from_path(path: str) -> List[str]: """ - Extract the target MCP server name from the standard MCP transport - URL patterns: ``/mcp/{server_name}[/...]`` and + Extract the target MCP server name(s) from the standard MCP transport + URL patterns: ``/mcp/{server_name_or_csv}[/...]`` and ``/{server_name}/mcp[/...]``. Returns ``[]`` for any other path so callers fail closed when the target cannot be resolved. + Mirrors the regex-based parser in ``server.py::_get_mcp_servers_in_path`` + so the names used for auth gating match the names used for downstream + filtering. Without this alignment, an attacker could craft + ``/mcp//`` so that auth treats the request + as targeting the delegate server (bypassing LiteLLM auth) while + downstream filtering sees a different (non-existent) target and falls + back to the caller's full allowed-server set. + REST/admin endpoints, OAuth2 server endpoints (``/{server_name}/authorize``, ``/token`` etc.), and ``.well-known`` discovery routes intentionally fall through — those flows do not need OAuth2 token passthrough. Clients aggregating multiple servers should - use ``x-mcp-servers``, which takes precedence over path parsing. + use ``x-mcp-servers`` on a path that does not encode a target. """ + # ``/{server_name}/mcp[/...]`` form — single server. The literal + # ``mcp`` must be the second segment (not the first, which would be + # the ``/mcp/...`` form handled below). This branch must stay in sync + # with ``server.py::_get_mcp_servers_in_path``, which also accepts the + # un-rewritten form (some entry points may skip the + # ``dynamic_mcp_route`` rewrite). segments = [s for s in path.split("/") if s] - if len(segments) >= 2 and segments[0] == "mcp": - return [segments[1]] - if len(segments) >= 2 and segments[1] == "mcp": + if len(segments) >= 2 and segments[1] == "mcp" and segments[0] != "mcp": return [segments[0]] - return [] + + # ``/mcp/...`` form — server name(s) may contain a slash (e.g. + # ``custom_solutions/user_123``) and may be a comma-separated list. + # Use the same parsing logic as ``_get_mcp_servers_in_path`` so the + # parsed names match downstream routing. + mcp_path_match = re.match(r"^/mcp/([^?#]+)(?:\?.*)?(?:#.*)?$", path) + if not mcp_path_match: + return [] + servers_and_path = mcp_path_match.group(1) + if not servers_and_path: + return [] + + if "," in servers_and_path: + # Comma-separated servers, possibly followed by a trailing path. + path_match = re.search(r"/([^/,]+(?:/[^/,]+)*)$", servers_and_path) + if path_match: + servers_part = servers_and_path[: -(len(path_match.group(1)) + 1)] + else: + servers_part = servers_and_path + return [s.strip() for s in servers_part.split(",") if s.strip()] + + # Single-server case — server name may contain at most one slash. + single_server_match = re.match( + r"^([^/]+(?:/[^/]+)?)(?:/.*)?$", servers_and_path + ) + if single_server_match: + return [single_server_match.group(1)] + return [servers_and_path] @staticmethod def _target_servers_use_oauth2(path: str, mcp_servers: Optional[List[str]]) -> bool: @@ -217,13 +275,13 @@ class MCPRequestHandler: ) from litellm.types.mcp import MCPAuth - # Use the x-mcp-servers header verbatim when present (including the - # explicitly-empty list, which means "no targets" → fail closed). - # Only fall back to path parsing when the header was absent entirely. - target_names = ( - mcp_servers - if mcp_servers is not None - else MCPRequestHandler._extract_target_server_names_from_path(path) + # Resolve the same target list downstream routing will use. For + # ``/mcp/...`` routes, ``extract_mcp_auth_context`` overrides the + # ``x-mcp-servers`` header with path-derived names, so we must mirror + # that here — otherwise a caller could set the header to a permissive + # server while the path targets a stricter one (header/path TOCTOU). + target_names = MCPRequestHandler._resolve_target_server_names( + path=path, mcp_servers_header=mcp_servers ) if not target_names: return False @@ -234,6 +292,78 @@ class MCPRequestHandler: return False return True + @staticmethod + def _target_servers_delegate_auth_to_upstream( + path: str, mcp_servers: Optional[List[str]] + ) -> bool: + """ + True only when EVERY MCP server the request targets is configured for + ``auth_type == oauth2`` AND has ``delegate_auth_to_upstream=True``. + Fails closed when any target does not opt in or cannot be resolved. + + Used by :meth:`process_mcp_request` to skip LiteLLM API-key/SSO auth + entirely (PKCE passthrough) so the client authenticates directly with + the upstream MCP server. Mixed-target requests (e.g. one delegated + + one non-delegated server) fall back to normal LiteLLM auth. + """ + # Inline imports avoid a circular dependency: mcp_server_manager imports + # from this module. + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + + # See _target_servers_use_oauth2: must mirror the downstream + # header-vs-path override or an attacker could set + # ``x-mcp-servers`` to a delegate-enabled server while the URL path + # targets a non-delegate server, skipping LiteLLM auth for it. + target_names = MCPRequestHandler._resolve_target_server_names( + path=path, mcp_servers_header=mcp_servers + ) + if not target_names: + return False + + for name in target_names: + server = global_mcp_server_manager.get_mcp_server_by_name(name) + if server is None or server.auth_type != MCPAuth.oauth2: + return False + # `is True` is intentional: opt-in must be an explicit boolean + # True. A MagicMock attribute (in tests) or any other truthy + # non-bool must not silently enable the bypass. + if getattr(server, "delegate_auth_to_upstream", False) is not True: + return False + if not getattr(server, "available_on_public_internet", True): + return False + # Never delegate for M2M (client_credentials) servers: LiteLLM + # fetches the upstream token automatically using stored credentials, + # so allowing anonymous bypass would let any external caller invoke + # tools authenticated as LiteLLM's service account. + if server.has_client_credentials: + return False + return True + + @staticmethod + def _resolve_target_server_names( + path: str, mcp_servers_header: Optional[List[str]] + ) -> List[str]: + """ + Resolve the target MCP server names exactly as downstream routing + does (``server.py::extract_mcp_auth_context``). + + For ``/mcp/...`` paths, downstream routing **overrides** any + ``x-mcp-servers`` header value with the path-derived names. Mirror + that here so an attacker cannot use a permissive header value to + flip an auth gate while the path targets a stricter server + (header/path TOCTOU). For non-``/mcp/...`` paths (where the path + does not encode targets), fall back to the header. + """ + path_targets = MCPRequestHandler._extract_target_server_names_from_path(path) + if path_targets: + return path_targets + # Path did not resolve to /mcp/... targets — trust the header + # (including an explicitly empty list, which means "no targets"). + return mcp_servers_header if mcp_servers_header is not None else [] + @staticmethod def _get_mcp_auth_header_from_headers(headers: Headers) -> Optional[str]: """ diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4901bc76d2f..31ed0918f3c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -402,6 +402,9 @@ class MCPServerManager: available_on_public_internet=bool( server_config.get("available_on_public_internet", True) ), + delegate_auth_to_upstream=bool( + server_config.get("delegate_auth_to_upstream", False) + ), # AWS SigV4 fields aws_access_key_id=server_config.get("aws_access_key_id", None), aws_secret_access_key=server_config.get("aws_secret_access_key", None), @@ -796,6 +799,9 @@ class MCPServerManager: available_on_public_internet=bool( getattr(mcp_server, "available_on_public_internet", True) ), + delegate_auth_to_upstream=bool( + getattr(mcp_server, "delegate_auth_to_upstream", False) + ), created_at=getattr(mcp_server, "created_at", None), updated_at=getattr(mcp_server, "updated_at", None), tool_name_to_display_name=_deserialize_json_dict( @@ -967,6 +973,34 @@ class MCPServerManager: if not in_toolset_scope: combined_servers.update(allow_all_server_ids) + # For anonymous callers (no user_id, no role), also surface any + # servers the operator has opted into upstream-delegated auth. + # These servers handle their own auth at the upstream level, so + # LiteLLM granting access here does not bypass any security gate. + is_anonymous = not ( + user_api_key_auth + and ( + getattr(user_api_key_auth, "user_id", None) + or getattr(user_api_key_auth, "user_role", None) + or getattr(user_api_key_auth, "api_key", None) + ) + ) + if is_anonymous: + delegate_server_ids = [ + server.server_id + for server in self.get_registry().values() + if getattr(server, "auth_type", None) == MCPAuth.oauth2 + and getattr(server, "delegate_auth_to_upstream", False) is True + # M2M servers must not be exposed anonymously: an + # unauthenticated caller would get LiteLLM to proxy tool + # calls using its stored client_credentials. + and not server.has_client_credentials + # Internal-only servers must not be reachable from public + # internet callers who happen to carry an upstream token. + and getattr(server, "available_on_public_internet", True) + ] + combined_servers.update(delegate_server_ids) + if len(combined_servers) == 0: verbose_logger.debug( "No allowed MCP Servers found for user api key auth." diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 685232d0495..0a74a92f9ce 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1337,8 +1337,24 @@ if MCP_AVAILABLE: raw_headers=raw_headers, ) - # If no OAuth2 token came from request headers, fall back to pre-fetched creds - if extra_headers is None and server.auth_type == MCPAuth.oauth2: + # Prefer server-stored per-user OAuth when configured, so a stale + # Authorization header from the MCP client cannot override Redis/DB + # (same issue as call_tool in mcp_server_manager: VS Code caches tokens). + if ( + server.auth_type == MCPAuth.oauth2 + and getattr(server, "needs_user_oauth_token", False) + and user_api_key_auth is not None + ): + db_headers = await _get_user_oauth_extra_headers_from_db( + server, + user_api_key_auth, + prefetched_creds=_prefetched_oauth_creds, + ) + if db_headers: + extra_headers = db_headers + + # If still no OAuth2 token, fall back to pre-fetched creds (non-stale-client path) + elif extra_headers is None and server.auth_type == MCPAuth.oauth2: extra_headers = await _get_user_oauth_extra_headers_from_db( server, user_api_key_auth, @@ -2541,6 +2557,10 @@ if MCP_AVAILABLE: import re mcp_servers_from_path: Optional[List[str]] = None + segments = [s for s in path.split("/") if s] + if len(segments) >= 2 and segments[1] == "mcp" and segments[0] != "mcp": + return [segments[0]] + # Match /mcp/ # Where servers can be comma-separated list of server names # Server names can contain slashes (e.g., "custom_solutions/user_123") diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f513381c868..9a9d27b9b84 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1274,6 +1274,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None allow_all_keys: bool = False available_on_public_internet: bool = True + delegate_auth_to_upstream: bool = False is_byok: bool = False byok_description: List[str] = Field(default_factory=list) byok_api_key_help_url: Optional[str] = None @@ -1356,6 +1357,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): registration_url: Optional[str] = None allow_all_keys: bool = False available_on_public_internet: bool = True + delegate_auth_to_upstream: bool = False is_byok: bool = False byok_description: List[str] = Field(default_factory=list) byok_api_key_help_url: Optional[str] = None @@ -1427,6 +1429,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): registration_url: Optional[str] = None allow_all_keys: bool = False available_on_public_internet: bool = True + delegate_auth_to_upstream: bool = False is_byok: bool = False byok_description: List[str] = Field(default_factory=list) byok_api_key_help_url: Optional[str] = None diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index f2e64fdde83..587b80d4726 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -163,7 +163,7 @@ if MCP_AVAILABLE: ) from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.management_helpers.utils import management_endpoint_wrapper - from litellm.types.mcp import MCPCredentials + from litellm.types.mcp import MCPAuth, MCPCredentials from litellm.types.mcp_server.mcp_server_manager import MCPServer @dataclass @@ -1551,6 +1551,57 @@ if MCP_AVAILABLE: except _jwt.InvalidTokenError: pass + # For delegate_auth_to_upstream servers the entire PKCE handshake + # (both /authorize browser redirect and /token authorization_code + # exchange) must work without a LiteLLM session. /authorize is opened + # in a VS Code webview that may have no cookie; /token is a programmatic + # POST from VS Code. PKCE security (code_verifier) guarantees the + # authorization_code exchange cannot be replayed, so anonymous access + # is safe for that grant only. + # + # Importantly, NOT safe for refresh_token grants: ``mcp_token`` will + # forward the request to the upstream issuer with LiteLLM's stored + # ``client_secret`` attached, so any caller holding a refresh token + # issued to this client could mint fresh upstream access tokens through + # us. Require normal LiteLLM auth for those. + if not api_key: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 + global_mcp_server_manager, + ) + + server_id = request.path_params.get("server_id", "") + if server_id: + _s = global_mcp_server_manager.get_mcp_server_by_id(server_id) + if not _s: + _s = global_mcp_server_manager.get_mcp_server_by_name(server_id) + if ( + _s + and getattr(_s, "auth_type", None) == MCPAuth.oauth2 + and getattr(_s, "delegate_auth_to_upstream", False) is True + and getattr(_s, "available_on_public_internet", True) + # M2M servers fetch tokens with stored credentials; never + # expose their /authorize or /token endpoints anonymously. + and not _s.has_client_credentials + ): + # For /token, require PKCE authorization_code; refresh_token + # grants must NOT bypass auth (see comment above). + path_lower = (request.url.path or "").rstrip("/").lower() + if path_lower.endswith("/token"): + body_data = await _read_request_body(request=request) + grant_type = (body_data or {}).get("grant_type", "") + if grant_type != "authorization_code": + # Fall through to normal LiteLLM auth (will 401 if + # no key supplied). + pass + else: + return UserAPIKeyAuth() + else: + # /authorize and other PKCE-flow GETs are safe to + # bypass: PKCE binds the upstream issuer's ``code`` + # to the original ``code_challenge`` so no anonymous + # token can be minted via the redirect alone. + return UserAPIKeyAuth() + request_data = await _read_request_body(request=request) request_data = populate_request_with_path_params( request_data=request_data, request=request diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 84ce99557e3..b53507abe6a 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -323,6 +323,7 @@ model LiteLLM_MCPServerTable { registration_url String? allow_all_keys Boolean @default(false) available_on_public_internet Boolean @default(true) + delegate_auth_to_upstream Boolean @default(false) is_byok Boolean @default(false) byok_description String[] @default([]) byok_api_key_help_url String? diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 268d064eacc..776c7fa67a6 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -68,6 +68,12 @@ class MCPServer(BaseModel): access_groups: Optional[List[str]] = None allow_all_keys: bool = False available_on_public_internet: bool = True + # When True AND auth_type == oauth2, MCP requests targeting this server + # bypass LiteLLM API-key/SSO auth (and the pre-emptive 401) so the client + # completes PKCE directly with the upstream MCP server. Honored only for + # auth_type=oauth2; ignored for any other auth_type. See + # MCPRequestHandler._target_servers_delegate_auth_to_upstream. + delegate_auth_to_upstream: bool = False is_byok: bool = False byok_description: List[str] = [] byok_api_key_help_url: Optional[str] = None diff --git a/schema.prisma b/schema.prisma index 84ce99557e3..b53507abe6a 100644 --- a/schema.prisma +++ b/schema.prisma @@ -323,6 +323,7 @@ model LiteLLM_MCPServerTable { registration_url String? allow_all_keys Boolean @default(false) available_on_public_internet Boolean @default(true) + delegate_auth_to_upstream Boolean @default(false) is_byok Boolean @default(false) byok_description String[] @default([]) byok_api_key_help_url String? diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 6e0dadcd4d8..5bb16a4cd48 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -1075,6 +1075,727 @@ class TestMCPOAuth2FallbackTargetGating: await MCPRequestHandler.process_mcp_request(scope) +@pytest.mark.asyncio +class TestMCPDelegateAuthToUpstream: + """ + Tests for the ``delegate_auth_to_upstream`` per-server flag. + + When set on an ``auth_type=oauth2`` MCP server, LiteLLM must skip its own + API-key/SSO check entirely so the client completes PKCE directly with the + upstream MCP server. The gate must fail closed for any non-oauth2 server, + any mixed-target request, and any request where the target cannot be + resolved. + """ + + @staticmethod + def _make_server(auth_type, delegate_auth_to_upstream=False): + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + return MCPServer( + server_id="test-server-id", + name="test-server", + transport="http", + auth_type=auth_type, + delegate_auth_to_upstream=delegate_auth_to_upstream, + ) + + async def test_delegate_skips_litellm_auth_with_no_authorization(self): + """ + oauth2 + delegate_auth_to_upstream=True, no Authorization header at + all → anonymous UserAPIKeyAuth and ``user_api_key_auth`` is never + called. + """ + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/delegated_oauth_server", + "headers": [], + } + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + ) + (auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope) + assert isinstance(auth_result, UserAPIKeyAuth) + mock_auth.assert_not_called() + + async def test_delegate_with_upstream_token_in_authorization_falls_back_to_anonymous( + self, + ): + """ + oauth2 + delegate_auth_to_upstream=True with an upstream OAuth token in + ``Authorization`` (not a LiteLLM key): LiteLLM auth is attempted first + (and fails), then the existing oauth2 fallback returns anonymous so the + bearer is forwarded upstream untouched. The delegate branch itself does + not fire when Authorization is present — that is what protects spend + tracking for callers using Authorization-style LiteLLM keys. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/delegated_oauth_server", + "headers": [(b"authorization", b"Bearer upstream-pkce-token")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + ) + ( + auth_result, + _, + _, + _, + oauth2_headers, + _, + ) = await MCPRequestHandler.process_mcp_request(scope) + assert isinstance(auth_result, UserAPIKeyAuth) + assert oauth2_headers.get("Authorization") == "Bearer upstream-pkce-token" + + async def test_delegate_off_still_requires_litellm_auth(self): + """ + oauth2 server but delegate flag is OFF → existing behaviour: a missing + / invalid LiteLLM key still 401s (no anonymous fast-path). + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/non_delegated_oauth_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=False, + ) + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + + async def test_delegate_ignored_for_non_oauth2_server(self): + """ + Defense in depth: even if an operator turns on delegate_auth_to_upstream + for a non-oauth2 server (api_key, bearer_token, etc.), the gate must + not fire — only oauth2 servers may delegate. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/api_key_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.api_key, + delegate_auth_to_upstream=True, + ) + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + + async def test_delegate_mixed_targets_fail_closed(self): + """ + x-mcp-servers can list multiple targets. If ANY of them does not opt in + to delegate_auth_to_upstream, the bypass must NOT fire — otherwise an + attacker could mix one delegated server in to skip auth on the others. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"x-mcp-servers", b"delegated_oauth,plain_oauth"), + ], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + def mock_lookup(name, client_ip=None): + if name == "delegated_oauth": + return TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + return TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=False, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.side_effect = mock_lookup + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + + async def test_delegate_no_resolvable_target_fail_closed(self): + """ + If the target server cannot be resolved at all (e.g. admin/REST path + that isn't ``/mcp/{name}`` or ``/{name}/mcp``), we cannot prove the + gate's preconditions, so we must fail closed and run normal auth. + """ + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "GET", + "path": "/admin/whatever", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = None + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + + async def test_explicit_litellm_key_takes_precedence_over_delegate(self): + """ + When ``x-litellm-api-key`` is present, normal auth runs even for a + delegate server, so ``user_id`` is resolved and any stored upstream + OAuth credentials can be looked up and forwarded. The bypass only + fires when no LiteLLM key is supplied. + """ + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/delegated_oauth_server", + "headers": [(b"x-litellm-api-key", b"Bearer sk-1234")], + } + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + return_value=UserAPIKeyAuth(user_id="real-user"), + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + ) + (auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope) + assert isinstance(auth_result, UserAPIKeyAuth) + assert auth_result.user_id == "real-user" + mock_auth.assert_called_once() + + async def test_litellm_key_via_authorization_header_not_bypassed(self): + """ + Regression: a LiteLLM key sent via the secondary ``Authorization`` header + (e.g. ``Authorization: Bearer sk-...``) must still trigger normal auth + and not be silently swallowed by the delegate bypass — otherwise spend + tracking and rate limiting are skipped for those callers. + """ + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/delegated_oauth_server", + "headers": [(b"authorization", b"Bearer sk-1234")], + } + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + return_value=UserAPIKeyAuth(user_id="real-user"), + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + ) + (auth_result, *_rest) = await MCPRequestHandler.process_mcp_request(scope) + assert isinstance(auth_result, UserAPIKeyAuth) + assert auth_result.user_id == "real-user" + mock_auth.assert_called_once() + + async def test_delegate_ignored_for_client_credentials_server(self): + """ + oauth2 + delegate_auth_to_upstream=True but oauth2_flow=client_credentials + → bypass must NOT fire; normal LiteLLM auth must be attempted. + + M2M servers fetch the upstream token automatically using stored + credentials, so allowing anonymous bypass would let any external + caller invoke tools as LiteLLM's service account. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/m2m_server", + "headers": [], + } + + m2m_server = MCPServer( + server_id="m2m-server-id", + name="m2m_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow="client_credentials", + ) + + async def mock_auth_raises(*_args, **_kwargs): + raise HTTPException(status_code=401, detail="No key provided") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_auth_raises, + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = m2m_server + # No delegate bypass → normal auth is attempted → 401 raised + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + mock_auth.assert_called_once() + + async def test_delegate_ignored_for_non_public_server(self): + """ + Internal-only delegate servers must not bypass LiteLLM auth for + anonymous public callers. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/internal_server", + "headers": [], + } + + internal_server = MCPServer( + server_id="internal-server-id", + name="internal_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + available_on_public_internet=False, + ) + + async def mock_auth_raises(*_args, **_kwargs): + raise HTTPException(status_code=401, detail="No key provided") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_auth_raises, + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = internal_server + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + mock_auth.assert_called_once() + + async def test_get_allowed_servers_excludes_client_credentials_delegate(self): + """ + get_allowed_mcp_servers must not surface M2M (client_credentials) delegate + servers to anonymous callers even if delegate_auth_to_upstream=True. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + pkce_server = MCPServer( + server_id="pkce-server", + name="pkce_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + available_on_public_internet=True, + ) + m2m_server = MCPServer( + server_id="m2m-server", + name="m2m_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow="client_credentials", + available_on_public_internet=True, + ) + manager.registry = { + pkce_server.server_id: pkce_server, + m2m_server.server_id: m2m_server, + } + + with patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + result = await manager.get_allowed_mcp_servers(None) + + assert "pkce-server" in result + assert "m2m-server" not in result + + async def test_get_allowed_servers_excludes_non_public_delegate(self): + """ + Internal-only (available_on_public_internet=False) delegate servers + must not appear in the anonymous allow-list. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + public_server = MCPServer( + server_id="public-server", + name="public_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + available_on_public_internet=True, + ) + internal_server = MCPServer( + server_id="internal-server", + name="internal_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + available_on_public_internet=False, + ) + manager.registry = { + public_server.server_id: public_server, + internal_server.server_id: internal_server, + } + + with patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + result = await manager.get_allowed_mcp_servers(None) + + assert "public-server" in result + assert "internal-server" not in result + + def test_extract_target_server_names_matches_routing_parser(self): + """ + Regression: _extract_target_server_names_from_path must match the + downstream regex parser in server.py::_get_mcp_servers_in_path. + + Previously, a request to ``/mcp//garbage`` was parsed as + targeting ```` by the auth gate (bypassing LiteLLM auth) + while the routing layer parsed it as ``/garbage`` — when + that name did not resolve, the request fell back to the anonymous + allow-list which can include ``allow_all_keys`` servers that normally + require a LiteLLM key. + """ + from litellm.proxy._experimental.mcp_server.server import ( + _get_mcp_servers_in_path, + ) + + cases = [ + # Single server, single segment. + ("/mcp/foo", ["foo"]), + # Server name with one embedded slash (two segments). + ("/mcp/foo/bar", ["foo/bar"]), + # Server name with embedded slash + extra path → name stays at two segments. + ("/mcp/foo/bar/tools", ["foo/bar"]), + # Comma-separated servers, no trailing path. + ("/mcp/foo,bar", ["foo", "bar"]), + # Comma-separated servers with trailing path. + ("/mcp/foo,bar/tools", ["foo", "bar"]), + # Alternative form ``//mcp`` is also parsed (both auth + # parser and routing parser handle it for defense-in-depth — some + # entry points may not be rewritten by ``dynamic_mcp_route``). + ("/foo/mcp", ["foo"]), + ("/foo/mcp/tools", ["foo"]), + # Non-MCP paths → empty (fail closed). + ("/.well-known/oauth-authorization-server", []), + ("/v1/keys", []), + ("/", []), + ] + for path_input, expected in cases: + assert ( + MCPRequestHandler._extract_target_server_names_from_path(path_input) + == expected + ), f"path={path_input!r} → expected {expected!r}" + assert ( + _get_mcp_servers_in_path(path_input) or [] + ) == expected, f"path={path_input!r} → routing expected {expected!r}" + + async def test_delegate_does_not_bypass_on_extra_path_segment(self): + """ + Regression: ``/mcp//`` must NOT bypass auth. + + The bypass key check is now performed against the same parsed target + as downstream routing — ``/`` — which will not + resolve to a delegate-enabled server, so normal LiteLLM auth runs. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/delegated_server/extra", + "headers": [], + } + + delegate_server = TestMCPDelegateAuthToUpstream._make_server( + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + + def lookup_by_name(name): + # Only the *exact* delegated name resolves. Anything else (e.g. + # ``delegated_server/extra``) returns None so the bypass fails. + if name == "delegated_server": + return delegate_server + return None + + async def mock_auth_raises(*_args, **_kwargs): + raise HTTPException(status_code=401, detail="No key provided") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_auth_raises, + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.side_effect = lookup_by_name + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + # Auth was attempted (not bypassed) because the parsed target + # name does not match any registered delegate server. + mock_auth.assert_called_once() + + async def test_delegate_ignores_x_mcp_servers_header_for_mcp_paths(self): + """ + Regression (header/path TOCTOU): For ``/mcp/...`` routes, downstream + routing overrides ``x-mcp-servers`` with the path-derived names. + The auth bypass must do the same — otherwise an attacker could send + ``x-mcp-servers: `` while the URL path targets a + non-delegate server, flipping the auth gate on a server that should + require a LiteLLM key. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/non_delegate_server", + "headers": [(b"x-mcp-servers", b"delegated_server")], + } + + delegate_server = MCPServer( + server_id="delegate-id", + name="delegated_server", + transport="http", + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + available_on_public_internet=True, + ) + non_delegate = MCPServer( + server_id="non-delegate-id", + name="non_delegate_server", + transport="http", + auth_type=MCPAuth.api_key, + ) + + def lookup_by_name(name): + return { + "delegated_server": delegate_server, + "non_delegate_server": non_delegate, + }.get(name) + + async def mock_auth_raises(*_args, **_kwargs): + raise HTTPException(status_code=401, detail="No key provided") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_auth_raises, + ) as mock_auth, + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.side_effect = lookup_by_name + # Bypass MUST NOT fire — path-derived target is the non-delegate + # server. Normal auth runs and 401s. + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + assert exc_info.value.status_code == 401 + mock_auth.assert_called_once() + + async def test_resolve_target_server_names_prefers_path_over_header(self): + """ + ``_resolve_target_server_names`` must: + + - For ``/mcp/`` paths, return the path-derived list and ignore + the header (mirrors downstream routing). + - For non-MCP paths, fall back to the header (including the explicit + empty-list case, which fails closed). + """ + # Path matches /mcp/... — header is ignored. + assert MCPRequestHandler._resolve_target_server_names( + path="/mcp/foo", mcp_servers_header=["evil"] + ) == ["foo"] + assert MCPRequestHandler._resolve_target_server_names( + path="/mcp/foo,bar", mcp_servers_header=["evil"] + ) == ["foo", "bar"] + assert MCPRequestHandler._resolve_target_server_names( + path="/foo/mcp", mcp_servers_header=["evil"] + ) == ["foo"] + # Path does not match — header is trusted. + assert MCPRequestHandler._resolve_target_server_names( + path="/.well-known/oauth-authorization-server", + mcp_servers_header=["foo"], + ) == ["foo"] + # Explicit empty list on a non-MCP path → empty (fail closed). + assert ( + MCPRequestHandler._resolve_target_server_names( + path="/.well-known/oauth-authorization-server", + mcp_servers_header=[], + ) + == [] + ) + # No header on a non-MCP path → empty. + assert ( + MCPRequestHandler._resolve_target_server_names( + path="/.well-known/oauth-authorization-server", + mcp_servers_header=None, + ) + == [] + ) + + class TestMCPCustomHeaderName: """Test suite for custom MCP authentication header name functionality""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index b53420f0000..794864f658b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2456,6 +2456,51 @@ class TestMCPServerManager: assert "test_server_1" in result assert "test_server_2" in result + @pytest.mark.asyncio + async def test_get_allowed_mcp_servers_anonymous_delegate_requires_oauth2(self): + """Anonymous delegated auth listing should only include oauth2 servers.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + manager = MCPServerManager() + oauth_delegate_server = MCPServer( + server_id="oauth-delegate", + name="oauth_delegate", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + ) + api_key_delegate_server = MCPServer( + server_id="api-key-delegate", + name="api_key_delegate", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + delegate_auth_to_upstream=True, + ) + oauth_non_delegate_server = MCPServer( + server_id="oauth-non-delegate", + name="oauth_non_delegate", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=False, + ) + manager.registry = { + oauth_delegate_server.server_id: oauth_delegate_server, + api_key_delegate_server.server_id: api_key_delegate_server, + oauth_non_delegate_server.server_id: oauth_non_delegate_server, + } + + with patch.object( + MCPRequestHandler, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[], + ): + result = await manager.get_allowed_mcp_servers(None) + + assert set(result) == {"oauth-delegate"} + def test_get_mcp_server_from_tool_name_uses_server_name_not_name(self): """ Test that _get_mcp_server_from_tool_name uses server.server_name instead of server.name diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index a0bfbff4222..7b3bb81e04b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -497,31 +497,39 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401(): oauth_server.auth_type = MCPAuth.oauth2 oauth_server.needs_user_oauth_token = True - with patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new_callable=AsyncMock, - return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", - new_callable=AsyncMock, - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", - new_callable=AsyncMock, - return_value=None, - ) as mock_get_stored_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", - return_value=oauth_server, - ), patch.object( - session_manager, - "handle_request", - new_callable=AsyncMock, - ) as mock_handle_request: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value=None, + ) as mock_get_stored_token, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=oauth_server, + ), + patch.object( + session_manager, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): with pytest.raises(HTTPException) as exc_info: await handle_streamable_http_mcp(scope, receive, send) @@ -562,32 +570,119 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401(): oauth_server.auth_type = MCPAuth.oauth2 oauth_server.needs_user_oauth_token = True - with patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new_callable=AsyncMock, - return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", - new_callable=AsyncMock, - return_value=False, - ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", - new_callable=AsyncMock, - return_value={"Authorization": "Bearer cached-token"}, - ) as mock_get_stored_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", - return_value=oauth_server, - ), patch.object( - session_manager, - "handle_request", - new_callable=AsyncMock, - ) as mock_handle_request: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, ["repro_oauth_server"], None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + new_callable=AsyncMock, + return_value={"Authorization": "Bearer cached-token"}, + ) as mock_get_stored_token, + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=oauth_server, + ), + patch.object( + session_manager, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): await handle_streamable_http_mcp(scope, receive, send) assert mock_get_stored_token.await_count == 1 assert mock_handle_request.await_count == 1 + + +@pytest.mark.asyncio +async def test_handle_streamable_http_mcp_emits_401_for_delegated_server_without_token(): + """ + OAuth2 server with ``delegate_auth_to_upstream=True`` and no Authorization + header must still emit a pre-emptive 401 with WWW-Authenticate so the + client kicks off PKCE. The 401 points at LiteLLM's discovery shim, which + in turn delegates to the upstream OAuth issuer. + """ + from fastapi import HTTPException + + try: + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager, + ) + except ImportError: + pytest.skip("MCP server not available") + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"host", b"litellm.example.com"), + ], + } + receive = AsyncMock() + send = AsyncMock() + user_auth = MagicMock() + user_auth.user_id = None + delegated_server = MagicMock() + delegated_server.auth_type = MCPAuth.oauth2 + delegated_server.delegate_auth_to_upstream = True + delegated_server.needs_user_oauth_token = True + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + user_auth, + None, + ["delegated_oauth_server"], + None, + None, + None, + ), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_stale_mcp_session", + new_callable=AsyncMock, + return_value=False, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + return_value=delegated_server, + ), + patch.object( + session_manager, + "handle_request", + new_callable=AsyncMock, + ) as mock_handle_request, + ): + with pytest.raises(HTTPException) as exc_info: + await handle_streamable_http_mcp(scope, receive, send) + + assert exc_info.value.status_code == 401 + assert "www-authenticate" in exc_info.value.headers + assert mock_handle_request.await_count == 0 diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 30ad84e18b8..47e058786a5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1645,6 +1645,108 @@ class TestTemporaryMCPSessionEndpoints: _, call_kwargs = auth_builder_mock.call_args assert call_kwargs["api_key"] == "Bearer sk-header-key" + @pytest.mark.asyncio + async def test_mcp_oauth_user_api_key_auth_requires_oauth2_for_delegate_bypass( + self, + ): + """Non-oauth2 servers must not get anonymous access from the delegate flag.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _mcp_oauth_user_api_key_auth, + ) + + expected_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + mock_request = MagicMock() + mock_request.headers = {} + mock_request.cookies = {} + mock_request.path_params = {"server_id": "server-1"} + non_oauth_server = MagicMock() + non_oauth_server.auth_type = MCPAuth.api_key + non_oauth_server.delegate_auth_to_upstream = True + mock_manager = MagicMock() + mock_manager.get_mcp_server_by_id.return_value = non_oauth_server + mock_manager.get_mcp_server_by_name.return_value = None + fake_proxy_server = types.SimpleNamespace(master_key=None) + + with ( + patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + mock_manager, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + AsyncMock(return_value=expected_auth), + ) as auth_builder_mock, + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + AsyncMock(return_value={}), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.populate_request_with_path_params", + side_effect=lambda request_data, request: request_data, + ), + ): + result = await _mcp_oauth_user_api_key_auth(mock_request) + + assert result is expected_auth + auth_builder_mock.assert_awaited_once() + _, call_kwargs = auth_builder_mock.call_args + assert call_kwargs["api_key"] == "" + + @pytest.mark.asyncio + async def test_mcp_oauth_user_api_key_auth_requires_public_server_for_delegate_bypass( + self, + ): + """Internal-only delegate servers must still require LiteLLM auth.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _mcp_oauth_user_api_key_auth, + ) + + expected_auth = generate_mock_user_api_key_auth( + user_role=LitellmUserRoles.PROXY_ADMIN + ) + mock_request = MagicMock() + mock_request.headers = {} + mock_request.cookies = {} + mock_request.path_params = {"server_id": "server-1"} + internal_server = MagicMock() + internal_server.auth_type = MCPAuth.oauth2 + internal_server.delegate_auth_to_upstream = True + internal_server.available_on_public_internet = False + internal_server.has_client_credentials = False + mock_manager = MagicMock() + mock_manager.get_mcp_server_by_id.return_value = internal_server + mock_manager.get_mcp_server_by_name.return_value = None + fake_proxy_server = types.SimpleNamespace(master_key=None) + + with ( + patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + mock_manager, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._user_api_key_auth_builder", + AsyncMock(return_value=expected_auth), + ) as auth_builder_mock, + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + AsyncMock(return_value={}), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.populate_request_with_path_params", + side_effect=lambda request_data, request: request_data, + ), + ): + result = await _mcp_oauth_user_api_key_auth(mock_request) + + assert result is expected_auth + auth_builder_mock.assert_awaited_once() + _, call_kwargs = auth_builder_mock.call_args + assert call_kwargs["api_key"] == "" + @pytest.mark.asyncio async def test_mcp_authorize_proxies_to_discoverable_endpoint(self): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx index dcb7298c830..8f60e50d2b2 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx @@ -1,7 +1,7 @@ import React, { useEffect } from "react"; import { Form, Select, Tooltip, Collapse, Input, Space, Button, Switch } from "antd"; import { InfoCircleOutlined, MinusCircleOutlined, PlusOutlined } from "@ant-design/icons"; -import { MCPServer } from "./types"; +import { MCPServer, AUTH_TYPE } from "./types"; const { Panel } = Collapse; interface MCPPermissionManagementProps { @@ -23,6 +23,8 @@ const MCPPermissionManagement: React.FC = ({ getAccessGroupOptions, }) => { const form = Form.useFormInstance(); + const watchedAuthType = Form.useWatch("auth_type", form); + const isOAuth2 = watchedAuthType === AUTH_TYPE.OAUTH2; // Set initial values when mcpServer changes useEffect(() => { @@ -40,12 +42,25 @@ const MCPPermissionManagement: React.FC = ({ if (typeof mcpServer.available_on_public_internet === "boolean") { form.setFieldValue("available_on_public_internet", mcpServer.available_on_public_internet); } + if (typeof mcpServer.delegate_auth_to_upstream === "boolean") { + form.setFieldValue("delegate_auth_to_upstream", mcpServer.delegate_auth_to_upstream); + } } else { form.setFieldValue("allow_all_keys", false); form.setFieldValue("available_on_public_internet", true); + form.setFieldValue("delegate_auth_to_upstream", false); } }, [mcpServer, form]); + // delegate_auth_to_upstream is only honored server-side when auth_type=oauth2. + // Force it back to false whenever the user switches away from oauth2 so a + // stale toggle value doesn't get persisted with another auth type. + useEffect(() => { + if (!isOAuth2) { + form.setFieldValue("delegate_auth_to_upstream", false); + } + }, [isOAuth2, form]); + return ( = ({ + {isOAuth2 && ( +
+
+ + Delegate auth to upstream (PKCE passthrough) + + + + +

+ Bypass LiteLLM auth so clients authenticate directly with the upstream OAuth MCP server. +

+
+ + + +
+ )} + diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index 17bcd59c43e..f8b0141b25d 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -284,6 +284,7 @@ const CreateMCPServer: React.FC = ({ credentials: credentialValues, allow_all_keys: allowAllKeysRaw, available_on_public_internet: availableOnPublicInternetRaw, + delegate_auth_to_upstream: delegateAuthToUpstreamRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -388,6 +389,7 @@ const CreateMCPServer: React.FC = ({ tool_name_to_description: Object.keys(toolNameToDescription).length > 0 ? toolNameToDescription : null, allow_all_keys: Boolean(allowAllKeysRaw), available_on_public_internet: Boolean(availableOnPublicInternetRaw), + delegate_auth_to_upstream: Boolean(delegateAuthToUpstreamRaw), static_headers: staticHeaders, ...(tokenValidation !== null && { token_validation: tokenValidation }), }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx index 760504d5797..1f2864f6759 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -177,6 +177,47 @@ describe("MCPServerEdit (stdio)", () => { }); }); +describe("MCPServerEdit (delegate auth)", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should clear delegate auth flag when saving a non-oauth2 server", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + auth_type: "none", + delegate_auth_to_upstream: false, + }); + + render( + , + ); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.auth_type).toBe("none"); + expect(payload.delegate_auth_to_upstream).toBe(false); + }); +}); + describe("MCPServerEdit (interactive OAuth)", () => { beforeEach(() => { vi.clearAllMocks(); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index 997e76982d4..9278d41c3e3 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -384,6 +384,7 @@ const MCPServerEdit: React.FC = ({ args: rawArgs, allow_all_keys: allowAllKeysRaw, available_on_public_internet: availableOnPublicInternetRaw, + delegate_auth_to_upstream: delegateAuthToUpstreamRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -552,6 +553,15 @@ const MCPServerEdit: React.FC = ({ static_headers: staticHeaders, allow_all_keys: Boolean(allowAllKeysRaw ?? mcpServer.allow_all_keys), available_on_public_internet: Boolean(availableOnPublicInternetRaw ?? mcpServer.available_on_public_internet), + // ``delegate_auth_to_upstream`` is only honored server-side for + // ``auth_type=oauth2``. The Form.Item is conditionally rendered so the + // value drops out of the form on auth_type change; force false for any + // non-oauth2 server to avoid persisting a stale ``true`` that would + // silently re-activate if auth_type is later switched back to oauth2. + delegate_auth_to_upstream: + restValues.auth_type === AUTH_TYPE.OAUTH2 + ? Boolean(delegateAuthToUpstreamRaw ?? mcpServer.delegate_auth_to_upstream) + : false, // Include token_validation when it is set (non-null) or when clearing an existing value ...(tokenValidation !== null || mcpServer.token_validation ? { token_validation: tokenValidation } diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index f9a3d57e952..1f8f7f68d33 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -272,6 +272,23 @@ export const MCPServerView: React.FC = ({ )} + {handleAuth(mcpServer.auth_type) === "oauth2" && ( +
+ Delegate Auth to Upstream +
+ {mcpServer.delegate_auth_to_upstream ? ( + + + Enabled (PKCE passthrough) + + ) : ( + + Disabled + + )} +
+
+ )}
Access Groups
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 814c5d74f46..7cfe08d9ee5 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -202,6 +202,7 @@ export interface MCPServer { tool_name_to_description?: Record; allow_all_keys?: boolean; available_on_public_internet?: boolean; + delegate_auth_to_upstream?: boolean; /** Stdio-only fields (present when transport === 'stdio') */ command?: string | null; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 7f28dbe6da9..756348f4937 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -8800,6 +8800,7 @@ interface ExchangeMcpOAuthTokenParams { clientSecret?: string; codeVerifier: string; redirectUri: string; + accessToken?: string | null; } export const exchangeMcpOAuthToken = async ({ @@ -8809,6 +8810,7 @@ export const exchangeMcpOAuthToken = async ({ clientSecret, codeVerifier, redirectUri, + accessToken, }: ExchangeMcpOAuthTokenParams) => { const base = getProxyBaseUrl(); const normalizedServerId = encodeURIComponent(serverId.trim()); @@ -8826,11 +8828,16 @@ export const exchangeMcpOAuthToken = async ({ body.set("code_verifier", codeVerifier); body.set("redirect_uri", redirectUri); + const headers: Record = { + "Content-Type": "application/x-www-form-urlencoded", + }; + if (accessToken) { + headers["Authorization"] = `Bearer ${accessToken}`; + } + const response = await fetch(url, { method: "POST", - headers: { - "Content-Type": "application/x-www-form-urlencoded", - }, + headers, body: body.toString(), }); diff --git a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx index 9157fedbe21..7edeade4cbd 100644 --- a/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useMcpOAuthFlow.tsx @@ -279,6 +279,7 @@ export const useMcpOAuthFlow = ({ clientSecret: flowState.clientSecret, codeVerifier: flowState.codeVerifier, redirectUri: flowState.redirectUri, + accessToken, }); onTokenReceived(token); diff --git a/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx b/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx index aa7ce84de5e..cf0a81dcadf 100644 --- a/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx +++ b/ui/litellm-dashboard/src/hooks/useUserMcpOAuthFlow.tsx @@ -224,6 +224,7 @@ export const useUserMcpOAuthFlow = ({ clientSecret: flowState.clientSecret, codeVerifier: flowState.codeVerifier, redirectUri: flowState.redirectUri, + accessToken, }); // Persist the token for this user via the backend. From 7e61dbb1df1ade78ad5ee9f84a495eed0a5cc270 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 14 May 2026 00:47:06 +0530 Subject: [PATCH 4/9] fix(responses): preserve cache_control in Responses API -> Chat Completion transformation (#27727) * fix(responses): preserve cache_control in Responses API -> Chat Completion transformation cache_control injected by AnthropicCacheControlHook was silently dropped when _transform_responses_api_content_to_chat_completion_content rebuilt content blocks with only {type, text}. Now copies cache_control through so Anthropic prompt caching works correctly when using client.responses.create with cache_control_injection_points. Co-authored-by: Cursor * fix(responses): preserve cache_control for input_image and input_file blocks Extends the cache_control fix to image and file content blocks, which were also silently dropping cache_control during the Responses API -> Chat Completion transformation. Adds tests for all three content block types. Co-authored-by: Cursor --------- Co-authored-by: Cursor Co-authored-by: Claude Babysitter --- .../transformation.py | 30 ++++--- .../test_litellm_completion_responses.py | 83 +++++++++++++++++++ 2 files changed, 100 insertions(+), 13 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 48b12a5fba9..e2ba8353591 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -1238,6 +1238,8 @@ class LiteLLMCompletionResponsesConfig: file_dict["file_data"] = item["file_data"] new_item: Dict[str, Any] = {"type": "file", "file": file_dict} + if "cache_control" in item: + new_item["cache_control"] = item["cache_control"] return new_item @staticmethod @@ -1282,26 +1284,28 @@ class LiteLLMCompletionResponsesConfig: ) ) elif item.get("type") == "input_image": - content_list.append( - dict( - LiteLLMCompletionResponsesConfig._transform_input_image_item_to_image_item( - item - ) + image_block = dict( + LiteLLMCompletionResponsesConfig._transform_input_image_item_to_image_item( + item ) ) + if "cache_control" in item: + image_block["cache_control"] = item["cache_control"] + content_list.append(image_block) else: # Skip text blocks with None text to avoid downstream errors text_value = item.get("text") if text_value is None: continue - content_list.append( - { - "type": LiteLLMCompletionResponsesConfig._get_chat_completion_request_content_type( - item.get("type") or "text" - ), - "text": text_value, - } - ) + content_block: Dict[str, Any] = { + "type": LiteLLMCompletionResponsesConfig._get_chat_completion_request_content_type( + item.get("type") or "text" + ), + "text": text_value, + } + if "cache_control" in item: + content_block["cache_control"] = item["cache_control"] + content_list.append(content_block) return content_list else: raise ValueError(f"Invalid content type: {type(content)}") diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 5d8ff8022e5..503a610e016 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -2170,3 +2170,86 @@ class TestEnsureOutputItemContentPartAdded: events = iterator._pending_response_events assert len(events) == 2 + + +class TestCacheControlPreservation: + def test_cache_control_preserved_in_content_transformation(self): + """cache_control injected by AnthropicCacheControlHook must survive + the Responses API -> Chat Completion content transformation.""" + content = [ + { + "type": "text", + "text": "hello", + "cache_control": {"type": "ephemeral"}, + } + ] + result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content( + content + ) + assert isinstance(result, list) + assert len(result) == 1 + assert result[0]["cache_control"] == {"type": "ephemeral"} + + def test_content_without_cache_control_unaffected(self): + """Content blocks that don't have cache_control should be unaffected.""" + content = [{"type": "text", "text": "hello"}] + result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content( + content + ) + assert isinstance(result, list) + assert len(result) == 1 + assert "cache_control" not in result[0] + + def test_cache_control_preserved_in_input_item_transformation(self): + """cache_control survives the full input-item -> messages transformation.""" + input_item = { + "role": "user", + "content": [ + { + "type": "text", + "text": "long context", + "cache_control": {"type": "ephemeral"}, + } + ], + } + messages = LiteLLMCompletionResponsesConfig._transform_responses_api_input_item_to_chat_completion_message( + input_item + ) + assert len(messages) == 1 + msg_content = ( + messages[0].get("content") + if isinstance(messages[0], dict) + else getattr(messages[0], "content", None) + ) + assert isinstance(msg_content, list) + assert msg_content[0]["cache_control"] == {"type": "ephemeral"} + + def test_cache_control_preserved_for_input_file_block(self): + content = [ + { + "type": "input_file", + "file_id": "file-abc123", + "cache_control": {"type": "ephemeral"}, + } + ] + result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content( + content + ) + assert isinstance(result, list) + assert len(result) == 1 + assert result[0]["cache_control"] == {"type": "ephemeral"} + + def test_cache_control_preserved_for_input_image_block(self): + content = [ + { + "type": "input_image", + "image_url": "https://example.com/img.png", + "cache_control": {"type": "ephemeral"}, + } + ] + result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content( + content + ) + assert isinstance(result, list) + assert len(result) == 1 + assert result[0]["cache_control"] == {"type": "ephemeral"} From d6fa08307b05c0675b983a284c9c87cf037a4044 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Wed, 13 May 2026 13:02:38 -0700 Subject: [PATCH 5/9] fix(proxy): expose db status on public /health/readiness MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit External readiness probes consumed the legacy detailed payload's `db` field to drive alerting and pod-rotation decisions. Stripping the body to `{"status": "healthy"}` broke those probes silently — the HTTP code still flipped to 503, but probes checking `body.db == "connected"` treated the response as healthy. Add `db` back to the unauthenticated payload. Keep the rest of the diagnostic fields (litellm_version, callbacks, cache, log_level) gated behind /health/readiness/details so the recon-leak gate from #26912 holds. Values match the legacy contract: "connected", "disconnected", "Not connected". --- .../health_endpoints/_health_endpoints.py | 22 +++++++++++++------ .../health_endpoints/test_health_endpoints.py | 15 ++++++++----- 2 files changed, 25 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 3032c9d38cc..ff3df11c448 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -1548,15 +1548,21 @@ def _allow_public_health_readiness_details() -> bool: return general_settings.get("allow_public_health_readiness_details") is True -async def _set_public_readiness_status(response: Response) -> None: +async def _resolve_public_readiness_db(response: Response) -> str: + """ + Return the db status string for the public probe and flip the response to + 503 when a configured DB is unreachable. Mirrors the legacy values: + "Not connected" (no DB configured), "connected", "disconnected". + """ from litellm.proxy.proxy_server import prisma_client if prisma_client is None: - return + return "Not connected" db_health_status = await _db_health_readiness_check() if db_health_status["status"] != "connected": response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE + return db_health_status["status"] @router.get( @@ -1565,15 +1571,17 @@ async def _set_public_readiness_status(response: Response) -> None: ) async def health_readiness(response: Response): """ - Public readiness probe. Keep this low-detail for unauthenticated load - balancers by default. Admins can opt into the legacy detailed public - payload with general_settings.allow_public_health_readiness_details. + Public readiness probe. Returns a low-detail payload safe to expose to + unauthenticated load balancers — `status` plus `db` so orchestrators and + external probes can distinguish "healthy" from "DB unreachable" without a + credential. Admins can opt into the legacy detailed payload with + general_settings.allow_public_health_readiness_details. """ if _allow_public_health_readiness_details(): return await _get_health_readiness_details(response=response) - await _set_public_readiness_status(response=response) - return {"status": "healthy"} + db_status = await _resolve_public_readiness_db(response=response) + return {"status": "healthy", "db": db_status} @router.get( diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index bcd7fcb37b3..80a4804956c 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -844,9 +844,14 @@ def test_health_readiness(proxy_client): duration_ms < 500 ), f"Health check took {duration_ms:.2f}ms, expected < 500ms for readiness endpoint" - # Assert response contains only low-detail public probe fields + # Assert response contains only low-detail public probe fields. `db` is + # included so unauthenticated probes can distinguish "DB unreachable" + # from a fully-healthy worker; its value depends on whether the test env + # exposes DATABASE_URL. response_data = response.json() - assert response_data == {"status": "healthy"} + assert set(response_data.keys()) == {"status", "db"} + assert response_data["status"] == "healthy" + assert response_data["db"] in {"connected", "disconnected", "Not connected"} print(f"Response time: {duration_ms:.2f}ms") @@ -1750,7 +1755,7 @@ async def test_health_readiness_returns_503_when_db_disconnected(): result = await health_readiness(response=response) assert response.status_code == 503 - assert result == {"status": "healthy"} + assert result == {"status": "healthy", "db": "disconnected"} @pytest.mark.asyncio @@ -1773,7 +1778,7 @@ async def test_health_readiness_returns_200_when_db_connected(): result = await health_readiness(response=response) assert response.status_code == 200 - assert result == {"status": "healthy"} + assert result == {"status": "healthy", "db": "connected"} @pytest.mark.asyncio @@ -1792,7 +1797,7 @@ async def test_health_readiness_returns_200_when_no_db_configured(): result = await health_readiness(response=response) assert response.status_code == 200 - assert result == {"status": "healthy"} + assert result == {"status": "healthy", "db": "Not connected"} def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): From ca82761e4843e110d4192d586731f812d2eb3d60 Mon Sep 17 00:00:00 2001 From: oss-agent-shin Date: Wed, 13 May 2026 13:28:22 -0700 Subject: [PATCH 6/9] docs(budget_manager): add docstring to BudgetManager.reset_cost (#27867) Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> --- litellm/budget_manager.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/litellm/budget_manager.py b/litellm/budget_manager.py index b25967579e0..bbebb6042cb 100644 --- a/litellm/budget_manager.py +++ b/litellm/budget_manager.py @@ -178,6 +178,18 @@ class BudgetManager: return list(self.user_dict.keys()) def reset_cost(self, user): + """ + Reset the tracked spend for a user back to zero. + + Clears both the aggregate ``current_cost`` and the per-model + ``model_cost`` breakdown stored for the given user. + + Args: + user: The user identifier whose cost should be reset. + + Returns: + dict: ``{"user": }`` reflecting the reset state. + """ self.user_dict[user]["current_cost"] = 0 self.user_dict[user]["model_cost"] = {} return {"user": self.user_dict[user]} From 714664fe21f49e22fd7b012ef680cd0b42bd646a Mon Sep 17 00:00:00 2001 From: oss-agent-shin Date: Wed, 13 May 2026 13:54:00 -0700 Subject: [PATCH 7/9] docs: add class docstring to _LoopWrapper (#27870) Document the purpose of the daemon thread that backs the sync branch of the timeout decorator. Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> --- litellm/timeout.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/litellm/timeout.py b/litellm/timeout.py index f9bf036cea2..0d03a3e45e8 100644 --- a/litellm/timeout.py +++ b/litellm/timeout.py @@ -90,6 +90,13 @@ def timeout(timeout_duration: float = 0.0, exception_to_raise=Timeout): class _LoopWrapper(Thread): + """Daemon thread that owns a dedicated asyncio event loop. + + Used by the sync branch of :func:`timeout` to run a coroutine on a + background event loop so the calling thread can wait on it with a + timeout via :func:`asyncio.run_coroutine_threadsafe`. + """ + def __init__(self): super().__init__(daemon=True) self.loop = asyncio.new_event_loop() From 06c1c3497f8c390c4e663fcc184fb1ccfc3ad1b0 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 13 May 2026 14:06:58 -0700 Subject: [PATCH 8/9] =?UTF-8?q?fix:=20Fix=20Redis=20Sentinel=20client=20ha?= =?UTF-8?q?ndling=20to=20solve=20authentication=20error=E2=80=A6=20(#26302?= =?UTF-8?q?)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: Fix Redis Sentinel client handling to solve authentication error with password protected sentinel (#25625) * fix Redis Sentinel authentication handling * test: cover Redis Sentinel auth routing * refactor: align Redis Sentinel kwargs threading * fix: avoid duplicate Redis Sentinel socket timeouts * Address review comments * refactor(_redis): return set from _get_redis_kwargs for O(1) lookup Align _get_redis_kwargs() with the cluster helper by returning a set instead of a list, so the sentinel connection-kwargs filter uses O(1) membership tests. Addresses Greptile review feedback on PR #26302. * fix(_redis): restore Azure-specific kwargs in cluster kwargs set The set-literal refactor of _get_redis_cluster_kwargs dropped four LiteLLM-custom Azure keys (azure_redis_ad_token, azure_client_id, azure_tenant_id, azure_client_secret) that the prior list form had explicitly appended. Because they are not in RedisCluster's argspec, they were silently stripped, breaking Azure IAM auth on cluster clients. Re-add them to the explicit include set. --------- Co-authored-by: Kristin Cowalcijk Co-authored-by: Sameer Kankute Co-authored-by: krrish-berri-2 Co-authored-by: claude --- litellm/_redis.py | 68 +++++++++++------- tests/test_litellm/test_redis.py | 114 +++++++++++++++++++++++++++++++ 2 files changed, 156 insertions(+), 26 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index 1c11ea829ba..65284162663 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -41,7 +41,7 @@ def _get_redis_kwargs(): "retry", } - include_args = [ + include_args = { "url", "redis_connect_func", "gcp_service_account", @@ -50,9 +50,9 @@ def _get_redis_kwargs(): "azure_client_id", "azure_tenant_id", "azure_client_secret", - ] + } - available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args + available_args = {x for x in arg_spec.args if x not in exclude_args} | include_args return available_args @@ -84,23 +84,23 @@ def _get_redis_cluster_kwargs(client=None): # Only allow primitive arguments exclude_args = {"self", "connection_pool", "retry", "host", "port", "startup_nodes"} - available_args = [x for x in arg_spec.args if x not in exclude_args] - available_args.append("password") - available_args.append("username") - available_args.append("ssl") - available_args.append("ssl_cert_reqs") - available_args.append("ssl_check_hostname") - available_args.append("ssl_ca_certs") - available_args.append( - "redis_connect_func" - ) # Needed for sync clusters and IAM detection - available_args.append("gcp_service_account") - available_args.append("gcp_ssl_ca_certs") - available_args.append("azure_redis_ad_token") - available_args.append("azure_client_id") - available_args.append("azure_tenant_id") - available_args.append("azure_client_secret") - available_args.append("max_connections") + available_args = {x for x in arg_spec.args if x not in exclude_args} + available_args |= { + "password", + "username", + "ssl", + "ssl_cert_reqs", + "ssl_check_hostname", + "ssl_ca_certs", + "redis_connect_func", # Needed for sync clusters and IAM detection + "gcp_service_account", + "gcp_ssl_ca_certs", + "azure_redis_ad_token", + "azure_client_id", + "azure_tenant_id", + "azure_client_secret", + "max_connections", + } return available_args @@ -479,10 +479,24 @@ def init_redis_cluster(redis_kwargs) -> redis.RedisCluster: return redis.RedisCluster(startup_nodes=new_startup_nodes, **cluster_kwargs) # type: ignore +def _get_redis_sentinel_connection_kwargs(redis_kwargs: dict) -> dict: + connection_kwargs = {} + args = _get_redis_kwargs() + for arg in redis_kwargs: + if arg in args: + connection_kwargs[arg] = redis_kwargs[arg] + + return connection_kwargs + + def _init_redis_sentinel(redis_kwargs) -> redis.Redis: sentinel_nodes = redis_kwargs.get("sentinel_nodes") sentinel_password = redis_kwargs.get("sentinel_password") service_name = redis_kwargs.get("service_name") + connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) + connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) + sentinel_kwargs = dict(connection_kwargs) + sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: raise ValueError( @@ -494,19 +508,22 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: # Set up the Sentinel client sentinel = redis.Sentinel( sentinel_nodes, - socket_timeout=REDIS_SOCKET_TIMEOUT, - password=sentinel_password, + sentinel_kwargs=sentinel_kwargs, ) # Return the master instance for the given service - return sentinel.master_for(service_name) + return sentinel.master_for(service_name, **connection_kwargs) def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: sentinel_nodes = redis_kwargs.get("sentinel_nodes") sentinel_password = redis_kwargs.get("sentinel_password") service_name = redis_kwargs.get("service_name") + connection_kwargs = _get_redis_sentinel_connection_kwargs(redis_kwargs) + connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) + sentinel_kwargs = dict(connection_kwargs) + sentinel_kwargs["password"] = sentinel_password if not sentinel_nodes or not service_name: raise ValueError( @@ -518,13 +535,12 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: # Set up the Sentinel client sentinel = async_redis.Sentinel( sentinel_nodes, - socket_timeout=REDIS_SOCKET_TIMEOUT, - password=sentinel_password, + sentinel_kwargs=sentinel_kwargs, ) # Return the master instance for the given service - return sentinel.master_for(service_name) + return sentinel.master_for(service_name, **connection_kwargs) def get_redis_client(**env_overrides): diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 2483469db23..282c9d72d51 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -426,6 +426,120 @@ def test_sync_client_prefers_cluster_over_url_via_env_var( assert len(call_kwargs["startup_nodes"]) == 1 +@patch("litellm._redis.redis.Sentinel") +def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_cls): + """Sentinel auth must be passed to the sentinel, not the Redis master client.""" + mock_sentinel = MagicMock() + mock_sentinel_cls.return_value = mock_sentinel + + get_redis_client( + sentinel_nodes=[("sentinel-1", 26379)], + sentinel_password="sentinel-secret", + service_name="mymaster", + password="redis-secret", + username="redis-user", + ssl=True, + ssl_cert_reqs="required", + ssl_check_hostname=True, + ssl_ca_certs="/tmp/test-ca.pem", + max_connections=17, + socket_timeout=5, + ) + + mock_sentinel_cls.assert_called_once() + sentinel_call_kwargs = mock_sentinel_cls.call_args[1] + assert "password" not in sentinel_call_kwargs + assert "username" not in sentinel_call_kwargs + assert "ssl" not in sentinel_call_kwargs + assert "ssl_cert_reqs" not in sentinel_call_kwargs + assert "ssl_check_hostname" not in sentinel_call_kwargs + assert "ssl_ca_certs" not in sentinel_call_kwargs + assert "max_connections" not in sentinel_call_kwargs + assert "socket_timeout" not in sentinel_call_kwargs + assert sentinel_call_kwargs["sentinel_kwargs"] == { + "password": "sentinel-secret", + "username": "redis-user", + "ssl": True, + "ssl_cert_reqs": "required", + "ssl_check_hostname": True, + "ssl_ca_certs": "/tmp/test-ca.pem", + "max_connections": 17, + "socket_timeout": 5, + } + assert "service_name" not in sentinel_call_kwargs["sentinel_kwargs"] + assert "sentinel_nodes" not in sentinel_call_kwargs["sentinel_kwargs"] + assert "sentinel_password" not in sentinel_call_kwargs["sentinel_kwargs"] + mock_sentinel.master_for.assert_called_once_with( + "mymaster", + password="redis-secret", + username="redis-user", + ssl=True, + ssl_cert_reqs="required", + ssl_check_hostname=True, + ssl_ca_certs="/tmp/test-ca.pem", + max_connections=17, + socket_timeout=5, + ) + + +@patch("litellm._redis.async_redis.Sentinel") +def test_async_sentinel_uses_sentinel_password_and_master_password( + mock_sentinel_cls, +): + """Async sentinel auth must mirror the sync sentinel password routing.""" + mock_sentinel = MagicMock() + mock_sentinel_cls.return_value = mock_sentinel + + get_redis_async_client( + sentinel_nodes=[("sentinel-1", 26379)], + sentinel_password="sentinel-secret", + service_name="mymaster", + password="redis-secret", + username="redis-user", + ssl=True, + ssl_cert_reqs="required", + ssl_check_hostname=True, + ssl_ca_certs="/tmp/test-ca.pem", + max_connections=17, + socket_timeout=5, + ) + + mock_sentinel_cls.assert_called_once() + sentinel_call_kwargs = mock_sentinel_cls.call_args[1] + assert "password" not in sentinel_call_kwargs + assert "username" not in sentinel_call_kwargs + assert "ssl" not in sentinel_call_kwargs + assert "ssl_cert_reqs" not in sentinel_call_kwargs + assert "ssl_check_hostname" not in sentinel_call_kwargs + assert "ssl_ca_certs" not in sentinel_call_kwargs + assert "max_connections" not in sentinel_call_kwargs + assert "socket_timeout" not in sentinel_call_kwargs + assert sentinel_call_kwargs["sentinel_kwargs"] == { + "password": "sentinel-secret", + "username": "redis-user", + "ssl": True, + "ssl_cert_reqs": "required", + "ssl_check_hostname": True, + "ssl_ca_certs": "/tmp/test-ca.pem", + "max_connections": 17, + "socket_timeout": 5, + } + assert "service_name" not in sentinel_call_kwargs["sentinel_kwargs"] + assert "sentinel_nodes" not in sentinel_call_kwargs["sentinel_kwargs"] + assert "sentinel_password" not in sentinel_call_kwargs["sentinel_kwargs"] + mock_sentinel.master_for.assert_called_once_with( + "mymaster", + password="redis-secret", + username="redis-user", + ssl=True, + ssl_cert_reqs="required", + ssl_check_hostname=True, + ssl_ca_certs="/tmp/test-ca.pem", + max_connections=17, + socket_timeout=5, + ) + + @patch("litellm._redis.init_redis_cluster") def test_sync_client_preserves_password_for_cluster_when_url_also_set( mock_init_cluster, monkeypatch From bbeb094d00ee890850996e012c5c7490e5b04fee Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 14 May 2026 02:39:12 +0530 Subject: [PATCH 9/9] Litellm agent oss staging 05 11 2026 (#27733) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(ollama): Include provider in model list for ollama (#26135) * Include provider in model names for ollama * Fix unit tests * fix(ollama): process both thinking and content in same streaming chunk (#26098) * fix(health_check): skip max_tokens for image_generation mode (#26417) * fix(health_check): skip max_tokens for image_generation mode `_update_litellm_params_for_health_check` injected `max_tokens` for every deployment. OpenAI `/v1/images/generations` strictly rejects unknown fields, so health checks for dall-e-* and gpt-image-1 always failed with `400 "Unknown parameter: 'max_tokens'"` even though the actual image endpoint calls succeed. Skip the `max_tokens` injection when `model_info.mode == "image_generation"`. `messages` still gets injected (downstream `_filter_model_params` already strips it for non-chat handlers). * Switch to allow-list with per-deployment override Per @krrishdholakia review: deny-listing image_generation only re-introduces the same bug for every other non-chat mode (embedding, audio_*, rerank, video_generation, ocr, search, moderation, ...). Replace the single image_generation skip with `_MAX_TOKEN_SUPPORT_MODES = {chat, completion, responses}`. Missing `mode` is treated as chat for backward compatibility. New modes are safe by default. Add `model_info.health_check_supports_max_tokens` as an operator escape hatch — True forces injection on a non-listed deployment (operator wants to bound probe tokens), False suppresses it on a chat-style deployment behind a strict-schema provider. Tests: parametrize over 3 chat-style + 10 non-chat modes, plus override on/off and the no-mode legacy path. * fix(http_handler): handle RequestNotRead in MaskedHTTPStatusError for multipart uploads (#26718) Squash-merged by litellm-agent from dawidkulpa's PR. * fix(ollama): guard against double 'ollama/' prefix in live model listing Greptile flagged that Ollama servers can return names that already start with 'ollama/'. Check the prefix before prepending so we don't produce 'ollama/ollama/...'. Adds a regression test. * Fix Ollama empty reasoning stream chunks Co-authored-by: Yassin Kortam --------- Co-authored-by: James Myatt Co-authored-by: VHash <225398745+vhash0@users.noreply.github.com> Co-authored-by: hayden Co-authored-by: dawidkulpa <84176950+dawidkulpa@users.noreply.github.com> Co-authored-by: Cursor Co-authored-by: Claude Co-authored-by: Yassin Kortam --- litellm/llms/custom_httpx/http_handler.py | 7 +- litellm/llms/ollama/chat/transformation.py | 4 +- litellm/llms/ollama/common_utils.py | 2 +- litellm/proxy/health_check.py | 39 +++++- .../test_credential_leak_prevention.py | 25 ++++ .../ollama/test_ollama_chat_transformation.py | 65 +++++++++ .../llms/ollama/test_ollama_model_info.py | 30 ++++- .../proxy/test_health_check_max_tokens.py | 126 ++++++++++++++++++ 8 files changed, 289 insertions(+), 9 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index af18c666679..e11d8532dbf 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -485,11 +485,16 @@ class MaskedHTTPStatusError(httpx.HTTPStatusError): if k.lower() not in ("content-encoding", "content-length") } + try: + request_content = original_error.request.content + except httpx.RequestNotRead: + request_content = b"" + masked_request = httpx.Request( method=original_error.request.method, url=masked_url, headers=original_error.request.headers, - content=original_error.request.content, + content=request_content, ) super().__init__( diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index 48534799c97..e36150a4954 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -507,10 +507,10 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator): # PROCESS REASONING CONTENT reasoning_content: Optional[str] = None content: Optional[str] = None - if chunk["message"].get("thinking") is not None: + if chunk["message"].get("thinking"): reasoning_content = chunk["message"].get("thinking") self.started_reasoning_content = True - elif chunk["message"].get("content") is not None: + if chunk["message"].get("content"): if ( self.started_reasoning_content and not self.finished_reasoning_content diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index 8aedd9b3500..8ca8b7d383a 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -108,7 +108,7 @@ class OllamaModelInfo(BaseLLMModelInfo): continue nm = entry.get("name") or entry.get("model") if isinstance(nm, str): - names.add(nm) + names.add(nm if nm.startswith("ollama/") else f"ollama/{nm}") except Exception as e: verbose_logger.warning(f"Error retrieving ollama tag endpoint: {e}") # If tags endpoint fails, fall back to static list diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 585ba883947..4a28143e617 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -36,6 +36,31 @@ ADMIN_ONLY_HEALTH_DISPLAY_PARAMS = ("api_base", "api_version") MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"] +# Modes whose health-check probe is a chat-style completion call and +# therefore accept `max_tokens`. Other modes (embedding, image_generation, +# audio_*, rerank, video_generation, ocr, search, moderation, ...) hit +# endpoints that reject unknown fields with 400 "Unknown parameter: +# 'max_tokens'". Allow-list so new modes are safe by default. +# Per-deployment override: `model_info.health_check_supports_max_tokens`. +_MAX_TOKEN_SUPPORT_MODES: frozenset = frozenset({"chat", "completion", "responses"}) + + +def _should_inject_health_check_max_tokens(model_info: dict) -> bool: + """ + Whether the health-check probe should include `max_tokens`. + + Order: + 1. `model_info.health_check_supports_max_tokens` (operator override). + 2. `_MAX_TOKEN_SUPPORT_MODES`. Missing `mode` is treated as `chat` + for backward compatibility. + """ + explicit = model_info.get("health_check_supports_max_tokens") + if explicit is not None: + return bool(explicit) + mode = model_info.get("mode") or "chat" + return mode in _MAX_TOKEN_SUPPORT_MODES + + # Health-check modes that forward `reasoning_effort` to the provider (chat-style calls). _HEALTH_CHECK_MODES_SUPPORTING_REASONING_EFFORT = frozenset( (None, "chat", "completion") @@ -389,14 +414,22 @@ def _update_litellm_params_for_health_check( Update the litellm params for health check. - gets a short `messages` param for health check + - adds a bounded `max_tokens` when the deployment is a chat-style mode + (`chat`, `completion`, `responses`) or the operator explicitly opts in + via `model_info.health_check_supports_max_tokens`. Non-chat endpoints + (image, embedding, audio_*, rerank, video, ocr, search, moderation, ...) + reject unknown fields with 400 "Unknown parameter: 'max_tokens'". - updates the `model` param with the `health_check_model` if it exists Doc: https://docs.litellm.ai/docs/proxy/health#wildcard-routes - updates the `voice` param with the `health_check_voice` for `audio_speech` mode if it exists Doc: https://docs.litellm.ai/docs/proxy/health#text-to-speech-models - for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID """ litellm_params["messages"] = _get_random_llm_message() - _resolved_max_tokens = _resolve_health_check_max_tokens(model_info, litellm_params) - if _resolved_max_tokens is not None: - litellm_params["max_tokens"] = _resolved_max_tokens + if _should_inject_health_check_max_tokens(model_info): + _resolved_max_tokens = _resolve_health_check_max_tokens( + model_info, litellm_params + ) + if _resolved_max_tokens is not None: + litellm_params["max_tokens"] = _resolved_max_tokens # Per-model reasoning effort for health checks only (e.g. reasoning_effort=none). if model_info.get("mode", None) in _HEALTH_CHECK_MODES_SUPPORTING_REASONING_EFFORT: diff --git a/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py b/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py index 72b4da7b38b..0a3bf403bf8 100644 --- a/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py +++ b/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py @@ -104,6 +104,31 @@ class TestMaskedHTTPStatusError: # The attached request must be the masked one, not the original. assert "KEY_X" not in str(req.url) + def test_handles_streaming_request_content(self): + """MaskedHTTPStatusError must not crash when request body is streamed.""" + streaming_request = httpx.Request( + "POST", + "https://api.openai.com/v1/images/edits?key=SECRET_KEY", + stream=httpx.ByteStream(b"multipart-data"), + ) + response = httpx.Response( + 400, + request=streaming_request, + content=b'{"error": "bad request"}', + ) + orig = httpx.HTTPStatusError( + message="400 Bad Request", + request=streaming_request, + response=response, + ) + + masked = MaskedHTTPStatusError(orig) + + assert masked.status_code == 400 + assert masked.response.status_code == 400 + assert masked.response.request is not None + assert "SECRET_KEY" not in str(masked.request.url) + def test_strips_content_encoding_to_avoid_double_decode(self): """If the upstream response declared Content-Encoding (e.g. gzip), the rebuilt Response must not carry that header over — otherwise httpx diff --git a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py index 05b96b88228..906c51d8064 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py @@ -694,6 +694,71 @@ class TestOllamaReasoningContentStreaming: # reasoning_content is not set when there's no thinking in the chunk assert getattr(result2.choices[0].delta, "reasoning_content", None) is None + def test_thinking_and_content_in_same_chunk(self): + """ + Test that a chunk containing both thinking and content preserves both fields. + """ + iterator = OllamaChatCompletionResponseIterator( + streaming_response=iter([]), + sync_stream=True, + ) + + chunk = { + "model": "deepseek-r1", + "message": { + "role": "assistant", + "thinking": "Let me reason first.", + "content": "Final answer.", + }, + "done": False, + } + + result = iterator.chunk_parser(chunk) + + assert result.choices[0].delta.reasoning_content == "Let me reason first." + assert result.choices[0].delta.content == "Final answer." + + def test_streaming_chunks_ignore_inactive_empty_reasoning_fields(self): + """ + Test that Ollama chunks with inactive empty fields stay in the active delta. + """ + iterator = OllamaChatCompletionResponseIterator( + streaming_response=iter([]), + sync_stream=True, + ) + + chunk = { + "model": "deepseek-r1", + "message": { + "role": "assistant", + "thinking": "Let me reason first.", + "content": "", + }, + "done": False, + } + + result = iterator.chunk_parser(chunk) + + assert result.choices[0].delta.reasoning_content == "Let me reason first." + assert result.choices[0].delta.content is None + assert iterator.finished_reasoning_content is False + + content_chunk = { + "model": "deepseek-r1", + "message": { + "role": "assistant", + "thinking": "", + "content": "Final answer.", + }, + "done": False, + } + + result = iterator.chunk_parser(content_chunk) + + assert getattr(result.choices[0].delta, "reasoning_content", None) is None + assert result.choices[0].delta.content == "Final answer." + assert iterator.finished_reasoning_content is True + def test_think_tags_in_content(self): """ Test that tags embedded in content are properly parsed. diff --git a/tests/test_litellm/llms/ollama/test_ollama_model_info.py b/tests/test_litellm/llms/ollama/test_ollama_model_info.py index 95fc80b7fd6..448a26bafe1 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_model_info.py +++ b/tests/test_litellm/llms/ollama/test_ollama_model_info.py @@ -73,7 +73,7 @@ class TestOllamaModelInfo: info = OllamaModelInfo() models = info.get_models() # Only 'alpha' and 'zeta' should be returned, sorted alphabetically - assert models == ["alpha", "zeta"] + assert models == ["ollama/alpha", "ollama/zeta"] # Ensure correct endpoint was called assert calls and calls[0].endswith("/api/tags") assert call_headers and call_headers[0] == {} @@ -122,7 +122,7 @@ class TestOllamaModelInfo: monkeypatch.setattr(httpx, "get", mock_get) info = OllamaModelInfo() models = info.get_models() - assert models == ["m1", "m2"] + assert models == ["ollama/m1", "ollama/m2"] def test_get_models_fallback_on_error(self, monkeypatch): """ @@ -139,6 +139,32 @@ class TestOllamaModelInfo: # Default static ollama_models is ['llama2'], so expect ['ollama/llama2'] assert models == ["ollama/llama2"] + def test_get_models_no_double_prefix(self, monkeypatch): + """ + Names that already carry the 'ollama/' prefix (or are returned by an + Ollama server that's been configured to emit them) should not be + prefixed a second time. + """ + sample = { + "models": [ + {"name": "ollama/already-prefixed"}, + {"name": "fresh"}, + {"name": "hf.co/Qwen/Qwen3-14B:latest"}, + ] + } + + def mock_get(url, headers): + return DummyResponse(sample, status_code=200) + + monkeypatch.setattr(httpx, "get", mock_get) + info = OllamaModelInfo() + models = info.get_models() + assert models == [ + "ollama/already-prefixed", + "ollama/fresh", + "ollama/hf.co/Qwen/Qwen3-14B:latest", + ] + class TestOllamaGetModelInfo: """Tests for OllamaConfig.get_model_info() api_base threading and graceful fallback.""" diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py index 72d77862b5d..5cb7cdacc60 100644 --- a/tests/test_litellm/proxy/test_health_check_max_tokens.py +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -227,6 +227,132 @@ def test_wildcard_ignores_reasoning_split_model_info(monkeypatch): assert _resolve_health_check_max_tokens(model_info, litellm_params) is None +# --------------------------------------------------------------------------- +# image_generation must not receive max_tokens. +# +# _update_litellm_params_for_health_check injected `max_tokens` for every +# deployment. For `mode: image_generation` that leaked into OpenAI +# `/v1/images/generations`, which strictly rejects unknown fields with +# `400 "Unknown parameter: 'max_tokens'"`, marking dall-e-* and +# gpt-image-1 as permanently unhealthy even though their actual image +# calls succeed. `messages` still gets injected (downstream +# `_filter_model_params` already strips it for non-chat handlers). +# --------------------------------------------------------------------------- + + +def test_image_generation_mode_skips_max_tokens(): + """image_generation must not receive max_tokens.""" + model_info = {"mode": "image_generation"} + litellm_params = {"model": "openai/dall-e-3", "api_key": "sk-test"} + + updated = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert "max_tokens" not in updated + # connection-level params must still pass through unchanged + assert updated["api_key"] == "sk-test" + + +def test_health_check_max_tokens_value_is_ignored_for_non_chat_modes(): + """A configured `health_check_max_tokens` *value* (the int that controls + how many tokens to inject) is still skipped when the mode is outside the + allow-list — the inject decision runs before value resolution, so the + value never reaches `_resolve_health_check_max_tokens`. Note this is + distinct from `health_check_supports_max_tokens` (the bool that toggles + injection on/off per deployment).""" + model_info = {"mode": "image_generation", "health_check_max_tokens": 50} + litellm_params = {"model": "openai/dall-e-3"} + + updated = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert "max_tokens" not in updated + + +def test_chat_mode_still_injects_max_tokens(): + """Regression guard: the chat-style probe payload is unchanged.""" + model_info = {"mode": "chat"} + litellm_params = {"model": "gpt-4"} + + updated = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated["max_tokens"] == 5 + + +def test_no_mode_still_injects_max_tokens(): + """Regression guard: model_info without `mode` keeps the legacy path.""" + model_info: dict = {} + litellm_params = {"model": "gpt-4"} + + updated = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated["max_tokens"] == 5 + + +# --------------------------------------------------------------------------- +# Allow-list behavior: only chat-style modes (chat / completion / responses) +# receive max_tokens. Every other mode is skipped by default. +# +# Per-deployment override via `health_check_supports_max_tokens` lets the +# operator force injection on (e.g. a non-listed but max_tokens-capable +# endpoint where they want to bound probe token usage) or off (e.g. a +# chat-style provider with a strict schema that rejects unknown fields). +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("mode", ["chat", "completion", "responses"]) +def test_chat_style_modes_inject_max_tokens(mode): + updated = _update_litellm_params_for_health_check( + {"mode": mode}, {"model": f"openai/dummy-{mode}"} + ) + + assert updated["max_tokens"] == 5 + + +@pytest.mark.parametrize( + "mode", + [ + "embedding", + "image_generation", + "image_edit", + "audio_speech", + "audio_transcription", + "rerank", + "video_generation", + "ocr", + "search", + "moderation", + ], +) +def test_non_chat_modes_skip_max_tokens(mode): + updated = _update_litellm_params_for_health_check( + {"mode": mode}, {"model": f"openai/dummy-{mode}"} + ) + + assert "max_tokens" not in updated + + +def test_explicit_override_true_forces_injection_outside_allowlist(): + """Operator opts a non-listed deployment in to bound probe token usage.""" + model_info = { + "mode": "image_generation", + "health_check_supports_max_tokens": True, + } + litellm_params = {"model": "openai/some-future-image-model"} + + updated = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated["max_tokens"] == 5 + + +def test_explicit_override_false_suppresses_injection_inside_allowlist(): + """Operator opts a chat-style deployment out (strict-schema provider).""" + model_info = {"mode": "chat", "health_check_supports_max_tokens": False} + litellm_params = {"model": "openai/strict-schema-chat"} + + updated = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert "max_tokens" not in updated + + def test_update_litellm_params_health_check_reasoning_effort(): """model_info.health_check_reasoning_effort sets reasoning_effort for chat-style health checks.""" model_info = {"health_check_reasoning_effort": "low"}