From 58f1814cd7c1086c79bf30a11c7531ed33af363c Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 10 Jul 2026 12:11:04 -0700 Subject: [PATCH 001/123] fix(mcp): surface rejected delegate-auth upstream tokens as connect-time 401 For MCP servers with auth_type=oauth2 + delegate_auth_to_upstream=true, a client-supplied upstream token that the upstream rejects was masked: the upstream 401 raised during tools/list is absorbed by the list handler, so on a single-server route a rejected token became HTTP 200 with an empty tool list. Clients showed "0 tools" instead of re-authenticating, and monitoring never saw an unauthorized signal. Extend the connect-time preflight _check_passthrough_upstream_auth to probe delegate-auth servers with the caller's bare Authorization bearer, reusing the existing _probe_upstream_auth and the RFC 6750 challenge builder, so a rejected token fails the connect with 401 + WWW-Authenticate error="invalid_token" and a compliant client re-runs the upstream OAuth flow. The bare Authorization header is a valid upstream token only when admission took the delegate bypass, so the delegate target is resolved through get_mcp_server_by_name (the same resolver admission uses) rather than the wider allowed-server prefix/access-group matching. A name that reaches a delegate server only via server_id or an access group is admitted as a real LiteLLM key, so probing it would leak that key upstream; requiring the admission-resolver match closes that gap. The probe is gated to single-server routes (matching the OBO preflight), keyed to the caller's authorized set by server_id, and the challenge echoes the requested name so aliased routes get the same resource_metadata URL as the tokenless preemptive challenge. Tokenless requests keep flowing to the preemptive discovery challenge unchanged. Resolves LIT-4194 --- .../proxy/_experimental/mcp_server/server.py | 121 +++-- .../mcp_server/test_mcp_server.py | 451 ++++++++++++++++++ 2 files changed, 542 insertions(+), 30 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 3550237dd65..2bec607f77b 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3684,8 +3684,15 @@ if MCP_AVAILABLE: headers={"www-authenticate": upstream_www_authenticate}, ) + def _get_authorization_header_from_scope(scope: Scope) -> Optional[str]: + """First ``Authorization`` header value in the ASGI scope, or None.""" + for key, value in scope.get("headers", []): + if key.lower() == b"authorization": + return value.decode("latin-1") + return None + def _scope_has_authorization_header(scope: Scope) -> bool: - return any(key.lower() == b"authorization" for key, _ in scope.get("headers", [])) + return _get_authorization_header_from_scope(scope) is not None def _get_forwarded_auth_from_scope(scope: Scope) -> Optional[str]: """Return the upstream-bound ``Authorization`` header value, or None. @@ -3698,17 +3705,24 @@ if MCP_AVAILABLE: ``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 + has_litellm_key_header = any(key.lower() == b"x-litellm-api-key" for key, _ in scope.get("headers", [])) if not has_litellm_key_header: return None - return authorization + return _get_authorization_header_from_scope(scope) + + def _is_delegate_upstream_probe_target(server: MCPServer) -> bool: + """Whether ``server`` is an interactive delegate-auth server whose client-supplied + token should be preflighted upstream. + + Mirrors the anonymous-delegate gate in ``get_allowed_mcp_servers``: the flow is + resolved via ``effective_oauth2_flow`` so an unstamped M2M-shape row fails closed + (its stored client credentials drive egress; the caller's bearer is irrelevant). + """ + return ( + server.auth_type == MCPAuth.oauth2 + and server.delegate_auth_to_upstream is True + and MCPServerManager.effective_oauth2_flow(server) != "client_credentials" + ) async def _probe_upstream_auth( url: str, @@ -3770,7 +3784,7 @@ if MCP_AVAILABLE: mcp_servers: Optional[List[str]], client_ip: Optional[str], ) -> None: - """Probe pass-through upstream servers in parallel before the MCP session starts. + """Probe pass-through and delegate-auth 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 @@ -3778,11 +3792,42 @@ if MCP_AVAILABLE: 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. + with a WWW-Authenticate header if any upstream rejects the client token, or 403 + if the upstream accepts it but forbids the caller. Fails-open: network errors are logged and the request is allowed through. + + Delegate-auth servers (``auth_type=oauth2`` + ``delegate_auth_to_upstream``) + are probed with the caller's bare ``Authorization`` bearer. That bearer is only + an upstream token (never a LiteLLM key) when admission took the delegate bypass, + so the delegate target is resolved through ``get_mcp_server_by_name`` -- the same + resolver admission used -- rather than the wider allowed-server prefix/access-group + matching. A name that only reaches a delegate server via server_id or an access + group would have been admitted as a real LiteLLM key, so probing it would leak that + key upstream; requiring the admission-resolver match closes that gap. Without the + probe a rejected token is absorbed by the tools/list handler and masked as an empty + tool list. Gated to single-server routes so one rejected token cannot 401 a + multi-server aggregate connect, matching the OBO preflight gating; the challenge + echoes the requested name so aliased routes get the same resource_metadata URL as + the tokenless preemptive challenge. """ forwarded_auth = _get_forwarded_auth_from_scope(scope) - if not forwarded_auth: + requested_single_target = mcp_servers[0] if mcp_servers is not None and len(mcp_servers) == 1 else None + # The bare Authorization header (no x-litellm-api-key) is a valid upstream token + # only when admission classified it as one, i.e. the single requested name resolves + # to a delegate server under admission's own resolver. Resolve it the same way here + # so a server_id- or access-group-named delegate (which admission would have treated + # as a LiteLLM key) is never probed with that key. + delegate_server = ( + global_mcp_server_manager.get_mcp_server_by_name(requested_single_target, client_ip=client_ip) + if requested_single_target + else None + ) + delegate_auth = ( + _get_authorization_header_from_scope(scope) + if delegate_server is not None and _is_delegate_upstream_probe_target(delegate_server) + else None + ) + if not forwarded_auth and not delegate_auth: return # Use the authorized server set, not the raw user-supplied names, so that @@ -3792,33 +3837,49 @@ if MCP_AVAILABLE: mcp_servers=mcp_servers, client_ip=client_ip, ) - passthrough_servers = [ - srv - for srv in allowed_servers - # Restrict to genuine OAuth pass-through servers (auth_type none + - # Authorization in extra_headers). Gateway-managed OAuth2 servers - # must not receive the ``resource_metadata=`` challenge emitted - # below — they require ``authorization_uri=`` pointing at the - # gateway AS metadata. ``is_oauth_passthrough`` already requires - # ``auth_type in (None, MCPAuth.none)``, which is mutually - # exclusive with ``has_client_credentials`` (oauth2 + M2M flow), - # so M2M servers are implicitly excluded here. - if srv.is_oauth_passthrough - ] - if not passthrough_servers: + passthrough_targets: Tuple[Tuple[MCPServer, str, str], ...] = ( + tuple( + (srv, forwarded_auth, srv.name) + for srv in allowed_servers + # Restrict to genuine OAuth pass-through servers (auth_type none + + # Authorization in extra_headers). Gateway-managed OAuth2 servers + # must not receive the ``resource_metadata=`` challenge emitted + # below — they require ``authorization_uri=`` pointing at the + # gateway AS metadata. ``is_oauth_passthrough`` already requires + # ``auth_type in (None, MCPAuth.none)``, which is mutually + # exclusive with ``has_client_credentials`` (oauth2 + M2M flow), + # so M2M servers are implicitly excluded here. + if srv.is_oauth_passthrough + ) + if forwarded_auth + else () + ) + # Probe the admission-resolved delegate server only when the caller is actually + # authorized for it (present in the IP-filtered allowed set), keyed by server_id. + delegate_targets: Tuple[Tuple[MCPServer, str, str], ...] = ( + tuple( + (srv, delegate_auth, requested_single_target) + for srv in allowed_servers + if delegate_server is not None and srv.server_id == delegate_server.server_id + ) + if delegate_auth and requested_single_target + else () + ) + probe_targets = passthrough_targets + delegate_targets + if not probe_targets: return probe_results = await asyncio.gather( - *[_probe_upstream_auth(srv.url or "", forwarded_auth) for srv in passthrough_servers] + *[_probe_upstream_auth(srv.url or "", auth_header) for srv, auth_header, _ in probe_targets] ) - for srv, (probe_status, _) in zip(passthrough_servers, probe_results): + for (srv, _, challenge_server_name), (probe_status, _) in zip(probe_targets, probe_results): if probe_status == 401: # Token is missing or expired: keep pass-through clients on the # protected-resource discovery flow so they re-authorize against # the upstream IdP metadata proxied by LiteLLM. www_authenticate = _get_passthrough_www_authenticate( scope=scope, - server_name=srv.name, + server_name=challenge_server_name, invalid_token=True, ) raise HTTPException( 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 7d25e0ba493..89c2f17f93d 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 @@ -5379,6 +5379,457 @@ def test_get_forwarded_auth_from_scope_skips_when_no_litellm_key_header(): assert _get_forwarded_auth_from_scope(scope) is None +def _delegate_auth_mcp_server(server_id: str = "delegate-1") -> MCPServer: + return MCPServer( + server_id=server_id, + name="delegate_test", + url="http://upstream:9401/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow="authorization_code", + ) + + +def _delegate_scope(headers: list) -> dict: + return { + "type": "http", + "method": "POST", + "path": "/mcp/delegate_test", + "scheme": "http", + "server": ("localhost", 4000), + "headers": headers, + } + + +def _patch_delegate_resolver(server: MCPServer, *resolvable_names: str): + """Patch the admission-parity resolver the delegate probe gates on. Returns + ``server`` only for names admission's ``get_mcp_server_by_name`` would match + (alias / server_name / name); every other name (server_id, access group) yields + None, exactly as the real resolver does.""" + + def _resolve(name, client_ip=None): + return server if name in resolvable_names else None + + return patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + side_effect=_resolve, + ) + + +@pytest.mark.asyncio +async def test_delegate_bad_token_gets_connect_time_401(): + """Regression (LIT-4194): a rejected upstream token on a delegate-auth server + must fail the connect with 401 + ``error="invalid_token"``, not be absorbed + into HTTP 200 + an empty tool list by the tools/list handler. + + Delegate-mode clients send only ``Authorization`` (no ``x-litellm-api-key``), + so ``_get_forwarded_auth_from_scope`` returns None and, before the fix, the + preflight returned early without probing. + """ + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + server = _delegate_auth_mcp_server() + scope = _delegate_scope([(b"authorization", b"Bearer bogus-token")]) + + with _patch_delegate_resolver(server, "delegate_test"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, 'Bearer realm="upstream", error="invalid_token"')), + ) as probe: + with pytest.raises(HTTPException) as exc_info: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["delegate_test"], + client_ip=None, + ) + + assert exc_info.value.status_code == 401 + challenge = exc_info.value.headers["www-authenticate"] + assert 'error="invalid_token"' in challenge + assert 'resource_metadata="http://localhost:4000/.well-known/oauth-protected-resource/mcp/delegate_test"' in challenge + probe.assert_awaited_once() + probe_url, probe_auth = probe.call_args.args + assert probe_url == "http://upstream:9401/mcp" + assert probe_auth == "Bearer bogus-token" + + +@pytest.mark.asyncio +async def test_delegate_valid_token_passes_preflight(): + """An upstream-accepted token must not be blocked by the delegate preflight.""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + server = _delegate_auth_mcp_server() + scope = _delegate_scope([(b"authorization", b"Bearer good-token")]) + + with _patch_delegate_resolver(server, "delegate_test"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(200, None)), + ) as probe: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["delegate_test"], + client_ip=None, + ) + + probe.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_delegate_valid_token_forbidden_returns_403(): + """An upstream that accepts the token but forbids the caller (403) must surface + as a bare 403 with no ``WWW-Authenticate`` re-auth hint (a fresh token with the + same scopes would loop), not as an invalid_token challenge.""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + server = _delegate_auth_mcp_server() + scope = _delegate_scope([(b"authorization", b"Bearer scoped-out-token")]) + + with _patch_delegate_resolver(server, "delegate_test"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(403, None)), + ): + with pytest.raises(HTTPException) as exc_info: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["delegate_test"], + client_ip=None, + ) + + assert exc_info.value.status_code == 403 + assert not (exc_info.value.headers or {}) + + +@pytest.mark.asyncio +async def test_delegate_tokenless_request_not_probed(): + """Tokenless delegate requests are the preemptive challenge's job; the + preflight must not probe upstream with an empty credential.""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + server = _delegate_auth_mcp_server() + scope = _delegate_scope([(b"content-type", b"application/json")]) + + with _patch_delegate_resolver(server, "delegate_test"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["delegate_test"], + client_ip=None, + ) + + probe.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_delegate_preflight_skipped_on_multi_server_routes(): + """The delegate probe is gated to single-server routes so one rejected token + cannot 401 a multi-server aggregate connect (matching the OBO preflight).""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + servers = [_delegate_auth_mcp_server("delegate-1"), _delegate_auth_mcp_server("delegate-2")] + scope = _delegate_scope([(b"authorization", b"Bearer bogus-token")]) + + with _patch_delegate_resolver(servers[0], "delegate_test", "other_server"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=servers), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["delegate_test", "other_server"], + client_ip=None, + ) + + probe.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bare_authorization_never_probes_passthrough_servers(): + """A bare ``Authorization`` header may be a LiteLLM key (backward-compat), so + only delegate servers (where admission classified it as an upstream token) + may be probed with it; ``is_oauth_passthrough`` servers still require the + unambiguous ``x-litellm-api-key`` + ``Authorization`` pair.""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + passthrough_server = MCPServer( + server_id="pt-1", + name="pt_server", + url="http://upstream:9402/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + oauth_passthrough=True, + extra_headers=["Authorization"], + ) + scope = _delegate_scope([(b"authorization", b"Bearer ambiguous-token")]) + + with _patch_delegate_resolver(passthrough_server, "pt_server"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[passthrough_server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["pt_server"], + client_ip=None, + ) + + probe.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_delegate_not_probed_when_named_only_via_server_id(): + """Security regression (LIT-4194): a delegate server reachable by the requested + name only through its server_id (or an access group) is admitted as a real + LiteLLM key by ``process_mcp_request`` (its ``get_mcp_server_by_name`` misses), + so the bare ``Authorization`` header is that LiteLLM key. The probe must resolve + the target through the SAME resolver and therefore skip it, never forwarding the + key upstream, even though the widened allowed-server set still contains it.""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + server = _delegate_auth_mcp_server(server_id="delegate-secret-id") + # Admission's resolver matches alias/server_name/name only, never server_id: the + # requested server_id resolves to None here, mirroring the real divergence. + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/delegate-secret-id", + "scheme": "http", + "server": ("localhost", 4000), + "headers": [(b"authorization", b"Bearer sk-litellm-proxy-key")], + } + + with _patch_delegate_resolver(server, "delegate_test"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(user_id="u1", api_key="hashed-sk"), + mcp_servers=["delegate-secret-id"], + client_ip=None, + ) + + probe.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_delegate_preflight_with_unpatched_probe(): + """Integration across the preflight and the unpatched ``_probe_upstream_auth``, + mocked only at the httpx-client boundary (tests/test_litellm is mocked-only; the + real-network proof lives in the PR's live-proxy evidence). The mock honors the + ``AsyncHTTPHandler.post`` contract by raising ``httpx.HTTPStatusError`` on the + upstream 401, so the production ``except httpx.HTTPStatusError`` branch is the one + exercised. A rejected token surfaces as the connect-time 401 challenge; an + accepted token passes untouched, and the caller's bearer reaches the delegate URL.""" + import httpx + + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + accepted = MagicMock() + accepted.status_code = 200 + accepted.headers = {} + rejected = MagicMock() + rejected.status_code = 401 + rejected.headers = {"www-authenticate": 'Bearer realm="stub-upstream", error="invalid_token"'} + + async def respond_by_token(url=None, headers=None, json=None, timeout=None, **kwargs): + if headers.get("Authorization") == "Bearer good-token": + return accepted + raise httpx.HTTPStatusError( + "401 Unauthorized", + request=httpx.Request("POST", url), + response=rejected, + ) + + mock_client = MagicMock() + mock_client.post = AsyncMock(side_effect=respond_by_token) + + server = _delegate_auth_mcp_server() + + with _patch_delegate_resolver(server, "delegate_test"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server.get_async_httpx_client", + return_value=mock_client, + ): + with pytest.raises(HTTPException) as exc_info: + await _check_passthrough_upstream_auth( + scope=_delegate_scope([(b"authorization", b"Bearer bogus-token")]), + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["delegate_test"], + client_ip=None, + ) + + await _check_passthrough_upstream_auth( + scope=_delegate_scope([(b"authorization", b"Bearer good-token")]), + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["delegate_test"], + client_ip=None, + ) + + assert exc_info.value.status_code == 401 + challenge = exc_info.value.headers["www-authenticate"] + assert 'error="invalid_token"' in challenge + assert 'resource_metadata="http://localhost:4000/.well-known/oauth-protected-resource/mcp/delegate_test"' in challenge + probed_urls = [call.kwargs["url"] for call in mock_client.post.await_args_list] + assert probed_urls == ["http://upstream:9401/mcp", "http://upstream:9401/mcp"] + + +@pytest.mark.asyncio +async def test_delegate_challenge_echoes_requested_alias(): + """An alias-routed delegate request must be probed, and the challenge must echo + the requested alias (not the canonical server name) so the resource_metadata + URL matches what the tokenless preemptive challenge emits for the same route.""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + server = _delegate_auth_mcp_server().model_copy(update={"alias": "dt-alias"}) + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/dt-alias", + "scheme": "http", + "server": ("localhost", 4000), + "headers": [(b"authorization", b"Bearer bogus-token")], + } + + with _patch_delegate_resolver(server, "dt-alias"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, 'Bearer error="invalid_token"')), + ): + with pytest.raises(HTTPException) as exc_info: + await _check_passthrough_upstream_auth( + scope=scope, + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["dt-alias"], + client_ip=None, + ) + + challenge = exc_info.value.headers["www-authenticate"] + assert 'error="invalid_token"' in challenge + assert 'resource_metadata="http://localhost:4000/.well-known/oauth-protected-resource/mcp/dt-alias"' in challenge + + +@pytest.mark.asyncio +async def test_delegate_probe_not_fanned_out_to_access_group_members(): + """A single access-group name passes the one-target route gate but must not fan + the delegate probe out to group-expanded member servers; the group name resolves + to no server under admission's resolver, so no probe fires.""" + from litellm.proxy._experimental.mcp_server.server import ( + _check_passthrough_upstream_auth, + ) + from litellm.proxy._types import UserAPIKeyAuth + + group_member = _delegate_auth_mcp_server() + + with _patch_delegate_resolver(group_member, "delegate_test"), patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[group_member]), + ), patch( + "litellm.proxy._experimental.mcp_server.server._probe_upstream_auth", + new=AsyncMock(return_value=(401, None)), + ) as probe: + await _check_passthrough_upstream_auth( + scope=_delegate_scope([(b"authorization", b"Bearer bogus-token")]), + user_api_key_auth=UserAPIKeyAuth(), + mcp_servers=["prod_tools_group"], + client_ip=None, + ) + + probe.assert_not_awaited() + + +def test_is_delegate_upstream_probe_target_fails_closed_on_m2m_shape(): + """An unstamped M2M-shape row (null ``oauth2_flow`` + client credentials) + resolves to ``client_credentials`` and must not be probed with the caller's + bearer; its stored client credentials drive egress instead.""" + from litellm.proxy._experimental.mcp_server.server import ( + _is_delegate_upstream_probe_target, + ) + + assert _is_delegate_upstream_probe_target(_delegate_auth_mcp_server()) is True + + m2m_shape = MCPServer( + server_id="delegate-m2m", + name="delegate_m2m", + url="http://upstream:9401/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + delegate_auth_to_upstream=True, + oauth2_flow=None, + token_url="http://idp:9000/token", + client_id="client", + client_secret="secret", + ) + assert _is_delegate_upstream_probe_target(m2m_shape) is False + + non_delegate = MCPServer( + server_id="oauth2-plain", + name="oauth2_plain", + url="http://upstream:9401/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + assert _is_delegate_upstream_probe_target(non_delegate) is False + + @pytest.mark.asyncio async def test_create_mcp_client_sampling_disabled_by_default(): """Sampling callback must be None when allow_sampling is not set (default False).""" From 75dd70a67829c89a7d9a446c2781f8278df22d30 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Tue, 30 Jun 2026 17:08:04 -0400 Subject: [PATCH 002/123] fix(model_cost): add supports_reasoning: false to Gemini image generation models vertex_ai/gemini-2.5-flash-image, vertex_ai/gemini-3-pro-image-preview, vertex_ai/gemini-3.1-flash-image-preview, gemini/gemini-3-pro-image-preview, and gemini/gemini-3.1-flash-image-preview were missing supports_reasoning entries; _supports_factory then fell through to the vertex_ai provider-level config which returns true, causing requests with reasoning_effort to be sent to an API that rejects them. --- ...odel_prices_and_context_window_backup.json | 5 +++++ tests/test_litellm/test_utils.py | 19 +++++++++++++++++++ 2 files changed, 24 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2a269d693ec..1f1bec02f9c 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -18727,6 +18727,7 @@ ], "supports_function_calling": false, "supports_prompt_caching": true, + "supports_reasoning": false, "supports_response_schema": true, "supports_system_messages": true, "supports_vision": true, @@ -18768,6 +18769,7 @@ ], "supports_function_calling": false, "supports_prompt_caching": true, + "supports_reasoning": false, "supports_response_schema": true, "supports_system_messages": true, "supports_vision": true, @@ -36726,6 +36728,7 @@ "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, + "supports_reasoning": false, "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -36748,6 +36751,7 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, + "supports_reasoning": false, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, "vertex_ai/gemini-3.1-flash-image-preview": { @@ -36761,6 +36765,7 @@ "output_cost_per_image": 0.0672, "output_cost_per_image_token": 6e-05, "output_cost_per_token": 3e-06, + "supports_reasoning": false, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models" }, "vertex_ai/gemini-3.1-flash-lite-preview": { diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index d739f9c116a..e05a69f01c8 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -4709,3 +4709,22 @@ class TestValidateEnvironmentTencent: assert "TENCENT_API_KEY" in result["missing_keys"] + +@pytest.mark.parametrize( + "model", + [ + "vertex_ai/gemini-2.5-flash-image", + "vertex_ai/gemini-3-pro-image-preview", + "vertex_ai/gemini-3.1-flash-image-preview", + "gemini/gemini-2.5-flash-image", + "gemini/gemini-3-pro-image-preview", + "gemini/gemini-3.1-flash-image-preview", + ], +) +def test_gemini_image_models_do_not_support_reasoning( + model: str, local_model_cost_map: None +) -> None: + assert litellm.supports_reasoning(model) is False, ( + f"{model} incorrectly classified as reasoning-capable. " + "Add 'supports_reasoning: false' to its model_cost entry." + ) From aa717bc4d0dd56317d1d51a927c49f3b598ca7c6 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Tue, 30 Jun 2026 17:17:10 -0400 Subject: [PATCH 003/123] fix(model_cost): apply supports_reasoning: false to root pricing JSON The backup file is used by tests; the root model_prices_and_context_window.json is what gets published to the pricing URL and loaded by the proxy at runtime. Without this, the proxy would continue resolving supports_reasoning via the provider-level fallback and returning true for Gemini image generation models. Also covers vertex_ai/gemini-3-pro-image and vertex_ai/gemini-3.1-flash-image (non-preview variants) and gemini/gemini-3.1-flash-image which exist only in the root JSON. --- model_prices_and_context_window.json | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index e79ddbe35d2..cdcd048e2c6 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -18847,6 +18847,7 @@ ], "supports_function_calling": false, "supports_prompt_caching": true, + "supports_reasoning": false, "supports_response_schema": true, "supports_system_messages": true, "supports_vision": true, @@ -18888,6 +18889,7 @@ ], "supports_function_calling": false, "supports_prompt_caching": true, + "supports_reasoning": false, "supports_response_schema": true, "supports_system_messages": true, "supports_vision": true, @@ -18929,6 +18931,7 @@ ], "supports_function_calling": false, "supports_prompt_caching": true, + "supports_reasoning": false, "supports_response_schema": true, "supports_system_messages": true, "supports_vision": true, @@ -36900,6 +36903,7 @@ "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, + "supports_reasoning": false, "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, @@ -36922,6 +36926,7 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, + "supports_reasoning": false, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, "vertex_ai/gemini-3-pro-image-preview": { @@ -36937,6 +36942,7 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, + "supports_reasoning": false, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, "vertex_ai/gemini-3.1-flash-image": { @@ -36950,6 +36956,7 @@ "output_cost_per_image": 0.0672, "output_cost_per_image_token": 6e-05, "output_cost_per_token": 3e-06, + "supports_reasoning": false, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models" }, "vertex_ai/gemini-3.1-flash-image-preview": { @@ -36963,6 +36970,7 @@ "output_cost_per_image": 0.0672, "output_cost_per_image_token": 6e-05, "output_cost_per_token": 3e-06, + "supports_reasoning": false, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models" }, "vertex_ai/gemini-3.1-flash-lite-preview": { From fd862bb2b8bed8cf9cfd6890b759a078f5492150 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Tue, 30 Jun 2026 20:23:55 -0400 Subject: [PATCH 004/123] fix(model_cost): add supports_reasoning: false to gemini/gemini-3-pro-image --- ...odel_prices_and_context_window_backup.json | 44 +++++++++++++++++++ model_prices_and_context_window.json | 3 +- tests/test_litellm/test_utils.py | 1 + 3 files changed, 47 insertions(+), 1 deletion(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 1f1bec02f9c..ebb4b69d3f0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -18696,6 +18696,50 @@ }, "supports_image_size": false }, + "gemini/gemini-3-pro-image": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "gemini", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "rpm": 1000, + "tpm": 4000000, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3-pro-image", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_reasoning": false, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query", + "supports_service_tier": true + }, "gemini/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index cdcd048e2c6..d8d90e5b052 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -18814,7 +18814,8 @@ "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, - "web_search_billing_unit": "per_query" + "web_search_billing_unit": "per_query", + "supports_reasoning": false }, "gemini/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index e05a69f01c8..fddf1f4a0be 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -4717,6 +4717,7 @@ class TestValidateEnvironmentTencent: "vertex_ai/gemini-3-pro-image-preview", "vertex_ai/gemini-3.1-flash-image-preview", "gemini/gemini-2.5-flash-image", + "gemini/gemini-3-pro-image", "gemini/gemini-3-pro-image-preview", "gemini/gemini-3.1-flash-image-preview", ], From 43b69d84b317143c72ec9bb28126e7956612ab9b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 10 Jul 2026 22:06:25 +0000 Subject: [PATCH 005/123] fix(model_cost): align backup gemini/gemini-3-pro-image entry with root pricing JSON --- litellm/model_prices_and_context_window_backup.json | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index ebb4b69d3f0..9b1e60dff1a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -18737,8 +18737,7 @@ "search_context_size_medium": 0.014, "search_context_size_high": 0.014 }, - "web_search_billing_unit": "per_query", - "supports_service_tier": true + "web_search_billing_unit": "per_query" }, "gemini/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, From e5421bfe1ec29744dbfa86c6b2d77ca94d17ea01 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 10 Jul 2026 23:16:09 +0000 Subject: [PATCH 006/123] fix(model_cost): add missing backup entries for gemini image models gemini/gemini-3.1-flash-image, vertex_ai/gemini-3-pro-image, and vertex_ai/gemini-3.1-flash-image existed in the root pricing JSON but not in litellm/model_prices_and_context_window_backup.json, leaving deployments with LITELLM_LOCAL_MODEL_COST_MAP=True unprotected. Copies the root entries into the backup verbatim and extends the regression test to cover all ten gemini image models, asserting each exists in the local cost map so a missing backup entry fails the test instead of passing vacuously --- ...odel_prices_and_context_window_backup.json | 72 +++++++++++++++++++ tests/test_litellm/test_utils.py | 7 ++ 2 files changed, 79 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9b1e60dff1a..2e21be86168 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -18782,6 +18782,48 @@ }, "web_search_billing_unit": "per_query" }, + "gemini/gemini-3.1-flash-image": { + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "litellm_provider": "gemini", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.045, + "output_cost_per_image_token": 6e-05, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "rpm": 1000, + "tpm": 4000000, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-image", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_reasoning": false, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-3.1-flash-image-preview": { "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, @@ -36781,6 +36823,22 @@ "tpm": 8000000, "supports_image_size": false }, + "vertex_ai/gemini-3-pro-image": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "supports_reasoning": false, + "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + }, "vertex_ai/gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, @@ -36797,6 +36855,20 @@ "supports_reasoning": false, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, + "vertex_ai/gemini-3.1-flash-image": { + "input_cost_per_image": 0.00056, + "input_cost_per_token": 5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.0672, + "output_cost_per_image_token": 6e-05, + "output_cost_per_token": 3e-06, + "supports_reasoning": false, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models" + }, "vertex_ai/gemini-3.1-flash-image-preview": { "input_cost_per_image": 0.00056, "input_cost_per_token": 5e-07, diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index fddf1f4a0be..576f2f20242 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -4714,17 +4714,24 @@ class TestValidateEnvironmentTencent: "model", [ "vertex_ai/gemini-2.5-flash-image", + "vertex_ai/gemini-3-pro-image", "vertex_ai/gemini-3-pro-image-preview", + "vertex_ai/gemini-3.1-flash-image", "vertex_ai/gemini-3.1-flash-image-preview", "gemini/gemini-2.5-flash-image", "gemini/gemini-3-pro-image", "gemini/gemini-3-pro-image-preview", + "gemini/gemini-3.1-flash-image", "gemini/gemini-3.1-flash-image-preview", ], ) def test_gemini_image_models_do_not_support_reasoning( model: str, local_model_cost_map: None ) -> None: + assert model in litellm.model_cost, ( + f"{model} is missing from the local model cost map. " + "Add its entry to litellm/model_prices_and_context_window_backup.json." + ) assert litellm.supports_reasoning(model) is False, ( f"{model} incorrectly classified as reasoning-capable. " "Add 'supports_reasoning: false' to its model_cost entry." From 831dbbc4dfe5c7203ce6d4bc68352e12e8faf455 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 12:18:48 -0700 Subject: [PATCH 007/123] fix(anthropic): translate raw adaptive thinking for chat completions on pre-4.6 models Clients that pass thinking={"type": "adaptive"} directly (not via the reasoning_effort alias) on the /chat/completions interface had it forwarded unmodified to pre-4.6 Anthropic models, which reject the shape. Mirrors the translation already applied on the native /v1/messages passthrough (#32867): translate to legacy thinking={type: enabled, budget_tokens}, capped below max_tokens, dropping thinking when max_tokens can't fit even the minimum budget. Hoists the shared budget-capping helper onto AnthropicConfig so both paths use one implementation. --- litellm/llms/anthropic/chat/transformation.py | 53 +++++++++++++- .../messages/transformation.py | 18 +---- .../test_anthropic_chat_transformation.py | 73 +++++++++++++++++++ 3 files changed, 126 insertions(+), 18 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 6033e54fb77..722bf504845 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -227,6 +227,10 @@ DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING = ( "Sonnet 4.6+, and Mythos Preview." ) +DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING = ( + "Dropping adaptive `thinking` for model=%s: max_tokens is too small to fit the minimum thinking budget." +) + DROP_UNSUPPORTED_SPEED_WARNING = ( "Dropping unsupported `speed` for model=%s (drop_params=True). Fast mode is only supported on select Opus models." ) @@ -1220,6 +1224,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): llm_provider=llm_provider, ) + @staticmethod + def _cap_thinking_budget_to_max_tokens( + thinking: AnthropicThinkingParam, max_tokens: Optional[int] + ) -> Optional[AnthropicThinkingParam]: + """Cap a legacy ``thinking.budget_tokens`` below ``max_tokens`` (Anthropic + requires ``max_tokens > budget_tokens``). Returns the (possibly capped) + thinking dict, or ``None`` when ``max_tokens`` is too small to fit even the + minimum thinking budget and thinking should be dropped.""" + budget = thinking.get("budget_tokens") + if max_tokens is None or not isinstance(budget, int): + return thinking + if max_tokens <= ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: + return None + if budget < max_tokens: + return thinking + return AnthropicThinkingParam(type=thinking.get("type", "enabled"), budget_tokens=max_tokens - 1) + def _extract_json_schema_from_response_format(self, value: Optional[dict]) -> Optional[dict]: if value is None: return None @@ -1463,7 +1484,37 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ): optional_params["metadata"] = {"user_id": value} elif param == "thinking": - optional_params["thinking"] = value + if ( + isinstance(value, dict) + and value.get("type") == "adaptive" + and not AnthropicConfig._is_adaptive_thinking_model(model) + ): + # Callers (e.g. Claude Code) send adaptive thinking + # unconditionally; translate it down to the legacy + # `thinking={type: enabled, budget_tokens}` interface a + # pre-4.6 model actually supports instead of forwarding a + # shape the model will reject. + max_tokens = non_default_params.get("max_completion_tokens") or non_default_params.get("max_tokens") + legacy_thinking = AnthropicConfig._map_reasoning_effort( + reasoning_effort="medium", + model=model, + llm_provider=self.custom_llm_provider or "anthropic", + ) + capped_thinking = ( + AnthropicConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens) + if legacy_thinking is not None + else None + ) + if capped_thinking is not None: + optional_params["thinking"] = capped_thinking + else: + litellm.verbose_logger.warning( + DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING, + model, + ) + optional_params.pop("thinking", None) + else: + optional_params["thinking"] = value elif param == "reasoning_effort": # Accept both string ("low") and dict ({"effort": "low", # "summary": "concise"}). The Responses->Chat parser keeps the diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index 39713e0f003..cb220764965 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -3,7 +3,6 @@ from typing import Any, AsyncIterator, Dict, List, Optional, Tuple import httpx from litellm.constants import ( - ANTHROPIC_MIN_THINKING_BUDGET_TOKENS, DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET, @@ -353,7 +352,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): except _BadRequestError as e: raise AnthropicError(message=str(e.message), status_code=400) capped_thinking = ( - AnthropicMessagesConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens) + AnthropicConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens) if legacy_thinking is not None else None ) @@ -371,21 +370,6 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): else: optional_params.pop("output_config", None) - @staticmethod - def _cap_thinking_budget_to_max_tokens(thinking: Dict, max_tokens: Optional[int]) -> Optional[Dict]: - """Cap a legacy ``thinking.budget_tokens`` below ``max_tokens`` (Anthropic - requires ``max_tokens > budget_tokens``). Returns the (possibly capped) - thinking dict, or ``None`` when ``max_tokens`` is too small to fit even the - minimum thinking budget and thinking should be dropped.""" - budget = thinking.get("budget_tokens") - if max_tokens is None or not isinstance(budget, int): - return thinking - if max_tokens <= ANTHROPIC_MIN_THINKING_BUDGET_TOKENS: - return None - if budget < max_tokens: - return thinking - return {**thinking, "budget_tokens": max_tokens - 1} - def transform_anthropic_messages_request( self, model: str, diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index 7fb38544c52..43ef7fcd971 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -10,6 +10,7 @@ from unittest.mock import MagicMock, patch import litellm from litellm.constants import ( + ANTHROPIC_MIN_THINKING_BUDGET_TOKENS, DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MAX_THINKING_BUDGET, @@ -2443,6 +2444,78 @@ def test_reasoning_effort_maps_to_adaptive_thinking_for_claude_4_6_models(): assert result["output_config"]["effort"] == effort_map[effort] +def test_raw_adaptive_thinking_translates_to_legacy_for_pre_46_model(): + """Clients like Claude Code send ``thinking={"type": "adaptive"}`` directly + (not via ``reasoning_effort``) on every request, regardless of which model + the request routes to. For a pre-4.6 model that doesn't understand + adaptive thinking, this must be translated to the legacy + ``thinking={type: enabled, budget_tokens}`` interface instead of being + forwarded raw, which Anthropic would reject.""" + config = AnthropicConfig() + + result = config.map_openai_params( + non_default_params={"thinking": {"type": "adaptive"}, "max_tokens": 8192}, + optional_params={}, + model="claude-haiku-4-5-20251001", + drop_params=False, + ) + + assert result["thinking"]["type"] == "enabled" + assert result["thinking"]["budget_tokens"] == DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET + + +def test_raw_adaptive_thinking_budget_capped_below_max_tokens(): + """Anthropic requires ``max_tokens > thinking.budget_tokens``. When the + default medium budget wouldn't fit, it must be capped below max_tokens + rather than forwarded as an invalid combination.""" + config = AnthropicConfig() + + max_tokens = DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET - 100 + result = config.map_openai_params( + non_default_params={"thinking": {"type": "adaptive"}, "max_tokens": max_tokens}, + optional_params={}, + model="claude-haiku-4-5-20251001", + drop_params=False, + ) + + assert result["thinking"]["type"] == "enabled" + assert result["thinking"]["budget_tokens"] == max_tokens - 1 + + +def test_raw_adaptive_thinking_dropped_when_max_tokens_too_small(): + """When max_tokens can't fit even the minimum thinking budget, thinking + must be dropped entirely so the request still succeeds, matching how the + native /v1/messages passthrough already handles this.""" + config = AnthropicConfig() + + result = config.map_openai_params( + non_default_params={ + "thinking": {"type": "adaptive"}, + "max_tokens": ANTHROPIC_MIN_THINKING_BUDGET_TOKENS, + }, + optional_params={}, + model="claude-haiku-4-5-20251001", + drop_params=False, + ) + + assert "thinking" not in result + + +def test_raw_adaptive_thinking_untouched_for_46_plus_model(): + """Adaptive-thinking models understand ``thinking={"type": "adaptive"}`` + natively, so it must pass through unmodified.""" + config = AnthropicConfig() + + result = config.map_openai_params( + non_default_params={"thinking": {"type": "adaptive"}, "max_tokens": 8192}, + optional_params={}, + model="claude-sonnet-4-6-20260219", + drop_params=False, + ) + + assert result["thinking"] == {"type": "adaptive"} + + @pytest.fixture def local_model_cost_map(monkeypatch): original_model_cost = litellm.model_cost From 741220c80fea828f3548a06cce644616b3953c26 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 14:23:28 -0700 Subject: [PATCH 008/123] fix(bedrock-converse): translate adaptive thinking for pre-4.6 models Follow-up to #32867 (native /v1/messages) and the /chat/completions commit earlier on this branch, extending the same adaptive-thinking translation to the Bedrock Converse path. Clients like Claude Code send thinking={type: "adaptive"} on every request. When routed via Bedrock Converse to pre-4.6 models (claude-haiku-4-5, claude-sonnet-4-5), this was forwarded as-is and rejected by the model. Mirrors the translation already applied on the /chat/completions and /v1/messages paths: map to legacy thinking={type: enabled, budget_tokens}, capped below max_tokens. Also fixes the missing custom_llm_provider arg in the chat completions path's call to AnthropicConfig._map_reasoning_effort. --- litellm/llms/anthropic/chat/transformation.py | 1 + .../bedrock/chat/converse_transformation.py | 24 ++++++++- .../chat/test_converse_transformation.py | 49 +++++++++++++++++++ 3 files changed, 73 insertions(+), 1 deletion(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 722bf504845..83054f61ec9 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1498,6 +1498,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): legacy_thinking = AnthropicConfig._map_reasoning_effort( reasoning_effort="medium", model=model, + custom_llm_provider=self.custom_llm_provider or "anthropic", llm_provider=self.custom_llm_provider or "anthropic", ) capped_thinking = ( diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index be904fb27be..c38b3593465 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -33,6 +33,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( make_valid_bedrock_tool_name, ) from litellm.llms.anthropic.chat.transformation import ( + DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING, DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING, REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT, AnthropicConfig, @@ -899,7 +900,28 @@ class AmazonConverseConfig(BaseConfig): "tool_choice": {"disable_parallel_tool_use": disable_parallel} } if param == "thinking": - optional_params["thinking"] = value + if ( + isinstance(value, dict) + and value.get("type") == "adaptive" + and not AnthropicConfig._is_adaptive_thinking_model(model, "bedrock") + ): + max_tokens = non_default_params.get("max_completion_tokens") or non_default_params.get("max_tokens") + legacy_thinking = AnthropicConfig._map_reasoning_effort( + reasoning_effort="medium", + model=model, + custom_llm_provider="bedrock", + ) + capped = ( + AnthropicConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens) + if legacy_thinking is not None + else None + ) + if capped is not None: + optional_params["thinking"] = capped + else: + litellm.verbose_logger.warning(DROP_UNSUPPORTED_ADAPTIVE_THINKING_WARNING, model) + else: + optional_params["thinking"] = value elif param == "reasoning_effort" and isinstance(value, str): self._handle_reasoning_effort_parameter( model=model, reasoning_effort=value, optional_params=optional_params diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index 9f2a4168dec..fe06933e249 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -5767,3 +5767,52 @@ def test_message_level_cache_control_drops_ttl_for_unsupported_model(ttl_target) cache_points = _collect_cache_points(result) assert len(cache_points) == 1 assert "ttl" not in cache_points[0] + + +@pytest.mark.parametrize( + "model", + [ + "bedrock/converse/us.anthropic.claude-haiku-4-5", + "bedrock/converse/us.anthropic.claude-sonnet-4-5", + ], +) +def test_adaptive_thinking_translated_to_legacy_on_pre_46_converse(model): + """Raw thinking={type: adaptive} from callers like Claude Code must be + translated to legacy thinking={type: enabled, budget_tokens} for pre-4.6 + models on Bedrock Converse rather than forwarded as-is and rejected.""" + config = AmazonConverseConfig() + + optional_params = config.map_openai_params( + non_default_params={"thinking": {"type": "adaptive"}, "max_tokens": 8192}, + optional_params={}, + model=model, + drop_params=False, + ) + + thinking = optional_params.get("thinking") + assert thinking is not None + assert thinking["type"] == "enabled" + assert isinstance(thinking.get("budget_tokens"), int) + assert thinking["budget_tokens"] < 8192 + + +@pytest.mark.parametrize( + "model", + [ + "bedrock/converse/us.anthropic.claude-opus-4-7", + "bedrock/converse/us.anthropic.claude-sonnet-4-6", + ], +) +def test_adaptive_thinking_passes_through_on_46_plus_converse(model): + """thinking={type: adaptive} must be forwarded unchanged for 4.6+ models + that natively support adaptive thinking.""" + config = AmazonConverseConfig() + + optional_params = config.map_openai_params( + non_default_params={"thinking": {"type": "adaptive"}, "max_tokens": 8192}, + optional_params={}, + model=model, + drop_params=False, + ) + + assert optional_params.get("thinking") == {"type": "adaptive"} From b58ccb092ea6d2cc4ca988dec795debd919db22c Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 14:48:14 -0700 Subject: [PATCH 009/123] fix(anthropic): pass resolved provider to adaptive-thinking check The rebase onto staging changed _is_adaptive_thinking_model to require custom_llm_provider (no default), so the one-arg call in the raw adaptive thinking branch raised TypeError at runtime for any /chat/completions caller sending thinking={type: adaptive}. Use self._resolved_provider, matching the reasoning_effort branch just below. Caught by Greptile. --- litellm/llms/anthropic/chat/transformation.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 83054f61ec9..1f5b76f3d0a 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1487,7 +1487,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if ( isinstance(value, dict) and value.get("type") == "adaptive" - and not AnthropicConfig._is_adaptive_thinking_model(model) + and not AnthropicConfig._is_adaptive_thinking_model(model, self._resolved_provider) ): # Callers (e.g. Claude Code) send adaptive thinking # unconditionally; translate it down to the legacy @@ -1498,8 +1498,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): legacy_thinking = AnthropicConfig._map_reasoning_effort( reasoning_effort="medium", model=model, - custom_llm_provider=self.custom_llm_provider or "anthropic", - llm_provider=self.custom_llm_provider or "anthropic", + custom_llm_provider=self._resolved_provider, + llm_provider=self._resolved_provider, ) capped_thinking = ( AnthropicConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens) From a5b0f32a84be1c122715d9963916eaba97870eb0 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 14:50:32 -0700 Subject: [PATCH 010/123] test(bedrock-converse): cover adaptive-thinking drop when max_tokens too small Adds the regression test for the warning-drop branch in the Converse adaptive-thinking translation, mirroring the chat completions path's test_raw_adaptive_thinking_dropped_when_max_tokens_too_small. --- .../chat/test_converse_transformation.py | 21 +++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index fe06933e249..fc12ead36a1 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -5816,3 +5816,24 @@ def test_adaptive_thinking_passes_through_on_46_plus_converse(model): ) assert optional_params.get("thinking") == {"type": "adaptive"} + + +def test_adaptive_thinking_dropped_when_max_tokens_too_small_converse(): + """When max_tokens can't fit even the minimum thinking budget, the raw + adaptive block must be dropped entirely rather than translated, so the + Bedrock Converse request still succeeds.""" + from litellm.constants import ANTHROPIC_MIN_THINKING_BUDGET_TOKENS + + config = AmazonConverseConfig() + + optional_params = config.map_openai_params( + non_default_params={ + "thinking": {"type": "adaptive"}, + "max_tokens": ANTHROPIC_MIN_THINKING_BUDGET_TOKENS, + }, + optional_params={}, + model="bedrock/converse/us.anthropic.claude-sonnet-4-5", + drop_params=False, + ) + + assert "thinking" not in optional_params From 2b2e8cf2bf3cca3431733c444203c8e8b7b8795e Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 15:44:50 -0700 Subject: [PATCH 011/123] feat(ui): working Test Connection for the complexity auto router The consolidated auto-router tab dropped the Test Connection button because the shared prepareModelAddRequest helper returns an empty array for an auto router (it has no model_mappings), so the caller crashed destructuring result[0].litellmParamsObj. That is the crash in #31590 and the open PR #31794. #31794 only silenced the crash by pointing the test at auto_router/complexity_router, which is not a provider model, so the /health/test_connection health check (a real litellm.ahealth_check completion) would still error. Bring the button back and make it meaningful: an auto router dispatches to saved model groups, so Test Connection now probes those directly. It builds a deduped target list from the configured tiers (tiers sharing a model group collapse to one probe) plus the embedding model when semantic keyword matching is on, then runs a live /health/test_connection against each and shows per-target pass/fail. This never touches prepareModelAddRequest, so the original destructure crash cannot recur. Scope is the recommended complexity router only; the to-be-deprecated semantic router is untouched. No backend changes. Supersedes #31794. Resolves #31590. --- .../add_model/add_auto_router_tab.tsx | 70 +++++++++- .../auto_router_connection_test.test.tsx | 80 +++++++++++ .../add_model/auto_router_connection_test.tsx | 132 ++++++++++++++++++ .../build_auto_router_test_targets.test.ts | 68 +++++++++ .../build_auto_router_test_targets.ts | 42 ++++++ 5 files changed, 387 insertions(+), 5 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx create mode 100644 ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx create mode 100644 ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.test.ts create mode 100644 ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 79c07040210..55c62224a48 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -1,5 +1,5 @@ import React, { useEffect, useState } from "react"; -import { Card, Form, Button, Tooltip, Typography, Select as AntdSelect, Radio, Badge, Space } from "antd"; +import { Card, Form, Button, Tooltip, Typography, Select as AntdSelect, Radio, Badge, Space, Modal } from "antd"; import type { FormInstance } from "antd"; import { ThunderboltOutlined, BranchesOutlined } from "@ant-design/icons"; import { Text, TextInput } from "@tremor/react"; @@ -12,6 +12,8 @@ import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./Complexit import { KeywordTierRule } from "./KeywordTierRules"; import { DEFAULT_MATCH_THRESHOLD } from "./SemanticKeywordMatching"; import { buildComplexityRouterConfig, getSemanticConfigError } from "./build_complexity_router_config"; +import { buildAutoRouterTestTargets, AutoRouterTestTarget } from "./build_auto_router_test_targets"; +import AutoRouterConnectionTest from "./auto_router_connection_test"; import NotificationManager from "../molecules/notifications_manager"; interface AddAutoRouterTabProps { @@ -45,6 +47,11 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc // Semantic router config (existing) const [routerConfig, setRouterConfig] = useState(null); + const [isTestModalVisible, setIsTestModalVisible] = useState(false); + const [isTestingConnection, setIsTestingConnection] = useState(false); + const [connectionTestId, setConnectionTestId] = useState(0); + const [testTargets, setTestTargets] = useState([]); + useEffect(() => { const fetchModelAccessGroups = async () => { const response = await modelAvailableCall(accessToken, "", "", false, null, true, true); @@ -194,6 +201,24 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc } }; + const handleTestConnection = () => { + const targets = buildAutoRouterTestTargets({ + tiers: complexityRouterConfig.tiers, + semanticMatchingEnabled, + embeddingModel, + }); + + if (targets.length === 0) { + NotificationManager.fromBackend("Please select at least one model for a complexity tier"); + return; + } + + setTestTargets(targets); + setConnectionTestId((id) => id + 1); + setIsTestingConnection(true); + setIsTestModalVisible(true); + }; + return ( <> Add Auto Router @@ -355,10 +380,15 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc Need Help?
- {/* TODO: add back a Test Connection or JSON preview action here. Test Connection was removed - because prepareModelAddRequest can't build a valid pre-save payload for an auto router - (tiers are model-group references, not litellm_params); a JSON preview of the - complexity_router_config would be a good alternative. */} + {routerType === "recommended" && ( + + )}
+ + { + setIsTestModalVisible(false); + setIsTestingConnection(false); + }} + footer={[ + , + ]} + width={700} + > + {isTestModalVisible && ( + setIsTestingConnection(false)} + /> + )} + ); }; diff --git a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx new file mode 100644 index 00000000000..9b872d5edee --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx @@ -0,0 +1,80 @@ +import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils"; +import { vi } from "vitest"; +import AutoRouterConnectionTest from "./auto_router_connection_test"; +import { AutoRouterTestTarget } from "./build_auto_router_test_targets"; + +vi.mock("../networking", async () => { + const actual = await vi.importActual("../networking"); + return { + ...actual, + testConnectionRequest: vi.fn(), + }; +}); + +const getMock = async () => vi.mocked((await import("../networking")).testConnectionRequest); + +const targets: AutoRouterTestTarget[] = [ + { labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }, + { labels: ["MEDIUM", "COMPLEX"], modelGroup: "claude-sonnet-4", mode: "chat" }, + { labels: ["Embedding"], modelGroup: "voyage-3-5", mode: "embedding" }, +]; + +describe("AutoRouterConnectionTest", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("probes each target once with the right model and mode (chat for tiers, embedding for the embedding model)", async () => { + const mock = await getMock(); + mock.mockResolvedValue({ status: "success" }); + + renderWithProviders(); + + await waitFor(() => expect(mock).toHaveBeenCalledTimes(3)); + + expect(mock).toHaveBeenCalledWith("sk-test", { model: "gpt-4o-mini" }, {}, "chat"); + expect(mock).toHaveBeenCalledWith("sk-test", { model: "claude-sonnet-4" }, {}, "chat"); + expect(mock).toHaveBeenCalledWith("sk-test", { model: "voyage-3-5" }, {}, "embedding"); + }); + + it("shows a success indicator per target when the health check passes", async () => { + const mock = await getMock(); + mock.mockResolvedValue({ status: "success" }); + + renderWithProviders(); + + await waitFor(() => expect(screen.getAllByTestId("test-status-success")).toHaveLength(3)); + expect(screen.queryByTestId("test-status-error")).toBeNull(); + expect(screen.getByText("MEDIUM, COMPLEX")).toBeInTheDocument(); + }); + + it("renders the provider error message for a failing target while others pass", async () => { + const mock = await getMock(); + mock.mockImplementation((_token, litellmParams) => + litellmParams.model === "claude-sonnet-4" + ? Promise.resolve({ status: "error", result: { error: "litellm.AuthenticationError: invalid api key" } }) + : Promise.resolve({ status: "success" }), + ); + + renderWithProviders(); + + await waitFor(() => expect(screen.getByTestId("test-error-message")).toBeInTheDocument()); + expect(screen.getByTestId("test-error-message")).toHaveTextContent("invalid api key"); + expect(screen.getByTestId("test-error-message")).not.toHaveTextContent("litellm.AuthenticationError"); + expect(screen.getAllByTestId("test-status-success")).toHaveLength(2); + }); + + it("surfaces a thrown network error as a failing row", async () => { + const mock = await getMock(); + mock.mockRejectedValue(new Error("Network request failed")); + + renderWithProviders( + , + ); + + await waitFor(() => expect(screen.getByTestId("test-error-message")).toHaveTextContent("Network request failed")); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx new file mode 100644 index 00000000000..77588006e35 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx @@ -0,0 +1,132 @@ +import React from "react"; +import { Typography } from "antd"; +import { CheckCircleTwoTone, CloseCircleTwoTone, LoadingOutlined } from "@ant-design/icons"; +import { testConnectionRequest } from "../networking"; +import { AutoRouterTestTarget } from "./build_auto_router_test_targets"; + +const { Text } = Typography; + +interface AutoRouterConnectionTestProps { + accessToken: string; + targets: AutoRouterTestTarget[]; + onTestComplete?: () => void; +} + +type TargetResult = { status: "pending" } | { status: "success" } | { status: "error"; error: string }; + +interface NormalizedResponse { + ok: boolean; + error?: string; +} + +const normalizeTestConnectionResponse = (response: unknown): NormalizedResponse => { + if (typeof response !== "object" || response === null) { + return { ok: false, error: "Unexpected response from connection test" }; + } + const record = response as Record; + if (record.status === "success") { + return { ok: true }; + } + const result = + typeof record.result === "object" && record.result !== null ? (record.result as Record) : {}; + const resultError = typeof result.error === "string" ? result.error : undefined; + const recordMessage = typeof record.message === "string" ? record.message : undefined; + return { ok: false, error: resultError ?? recordMessage ?? "Unknown error" }; +}; + +const cleanErrorMessage = (error: string): string => { + const mainError = error.split("stack trace:")[0].trim(); + return mainError.replace(/^litellm\.(.*?)Error: /, ""); +}; + +const runTarget = async (accessToken: string, target: AutoRouterTestTarget): Promise => { + try { + const response = await testConnectionRequest(accessToken, { model: target.modelGroup }, {}, target.mode); + const normalized = normalizeTestConnectionResponse(response); + return normalized.ok + ? { status: "success" } + : { status: "error", error: cleanErrorMessage(normalized.error ?? "Unknown error") }; + } catch (error) { + return { status: "error", error: cleanErrorMessage(error instanceof Error ? error.message : String(error)) }; + } +}; + +const AutoRouterConnectionTest: React.FC = ({ + accessToken, + targets, + onTestComplete, +}) => { + const [results, setResults] = React.useState(() => targets.map(() => ({ status: "pending" }))); + + React.useEffect(() => { + let cancelled = false; + const run = async () => { + const settled = await Promise.all(targets.map((target) => runTarget(accessToken, target))); + if (cancelled) return; + setResults(settled); + if (onTestComplete) onTestComplete(); + }; + run(); + return () => { + cancelled = true; + }; + // eslint-disable-next-line react-hooks/exhaustive-deps -- probes run once per mount; the parent remounts via `key` to start a fresh test, and re-running on prop identity changes would refire paid health checks + }, []); + + if (targets.length === 0) { + return No complexity tiers are configured yet, so there is nothing to test.; + } + + return ( +
+ + Each configured tier routes to a saved model group. Test Connection runs a live health check against each one. + + {targets.map((target, index) => { + const result = results[index] ?? { status: "pending" }; + return ( +
+
+ {result.status === "pending" && } + {result.status === "success" && ( + + )} + {result.status === "error" && ( + + )} +
+
+ {target.labels.join(", ")}{" "} + + {"->"} {target.modelGroup} + {target.mode === "embedding" ? " (embedding)" : ""} + + {result.status === "error" && ( + + {result.error} + + )} +
+
+ ); + })} +
+ ); +}; + +export default AutoRouterConnectionTest; diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.test.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.test.ts new file mode 100644 index 00000000000..01b6470f17f --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.test.ts @@ -0,0 +1,68 @@ +import { buildAutoRouterTestTargets } from "./build_auto_router_test_targets"; + +const tiers = { + SIMPLE: "gpt-4o-mini", + MEDIUM: "claude-sonnet-4", + COMPLEX: "claude-sonnet-4", + REASONING: "o3", +}; + +describe("buildAutoRouterTestTargets", () => { + it("dedups tiers that share a model group into one chat target carrying both labels", () => { + const targets = buildAutoRouterTestTargets({ tiers, semanticMatchingEnabled: false, embeddingModel: undefined }); + expect(targets).toEqual([ + { labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }, + { labels: ["MEDIUM", "COMPLEX"], modelGroup: "claude-sonnet-4", mode: "chat" }, + { labels: ["REASONING"], modelGroup: "o3", mode: "chat" }, + ]); + }); + + it("drops empty/whitespace tiers", () => { + const targets = buildAutoRouterTestTargets({ + tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: " ", REASONING: "" }, + semanticMatchingEnabled: false, + embeddingModel: undefined, + }); + expect(targets).toEqual([{ labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }]); + }); + + it("returns [] when no tier is configured", () => { + expect( + buildAutoRouterTestTargets({ + tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" }, + semanticMatchingEnabled: false, + embeddingModel: undefined, + }), + ).toEqual([]); + }); + + it("appends an embedding target only when semantic matching is on and a model is set", () => { + const targets = buildAutoRouterTestTargets({ + tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" }, + semanticMatchingEnabled: true, + embeddingModel: "voyage-3-5", + }); + expect(targets).toEqual([ + { labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }, + { labels: ["Embedding"], modelGroup: "voyage-3-5", mode: "embedding" }, + ]); + }); + + it("omits the embedding target when semantic matching is on but no model is chosen", () => { + const targets = buildAutoRouterTestTargets({ + tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" }, + semanticMatchingEnabled: true, + embeddingModel: undefined, + }); + expect(targets).toEqual([{ labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }]); + }); + + it("omits the embedding target when a model is set but semantic matching is off", () => { + const targets = buildAutoRouterTestTargets({ + tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" }, + semanticMatchingEnabled: false, + embeddingModel: "voyage-3-5", + }); + expect(targets).toEqual([{ labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }]); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts new file mode 100644 index 00000000000..0104b6bc9c5 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts @@ -0,0 +1,42 @@ +import { ComplexityTiers } from "./ComplexityRouterConfig"; + +export type AutoRouterTestMode = "chat" | "embedding"; + +export interface AutoRouterTestTarget { + labels: string[]; + modelGroup: string; + mode: AutoRouterTestMode; +} + +export interface BuildAutoRouterTestTargetsParams { + tiers: ComplexityTiers; + semanticMatchingEnabled: boolean; + embeddingModel: string | undefined; +} + +const TIER_ORDER: (keyof ComplexityTiers)[] = ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]; + +export const buildAutoRouterTestTargets = ({ + tiers, + semanticMatchingEnabled, + embeddingModel, +}: BuildAutoRouterTestTargetsParams): AutoRouterTestTarget[] => { + const groupedByModel = TIER_ORDER.reduce>((acc, tier) => { + const modelGroup = tiers[tier]?.trim(); + if (!modelGroup) return acc; + return { ...acc, [modelGroup]: [...(acc[modelGroup] ?? []), tier] }; + }, {}); + + const tierTargets: AutoRouterTestTarget[] = Object.entries(groupedByModel).map(([modelGroup, labels]) => ({ + labels, + modelGroup, + mode: "chat" as const, + })); + + const embeddingTarget: AutoRouterTestTarget[] = + semanticMatchingEnabled && embeddingModel?.trim() + ? [{ labels: ["Embedding"], modelGroup: embeddingModel.trim(), mode: "embedding" as const }] + : []; + + return [...tierTargets, ...embeddingTarget]; +}; From ddc13b331a83d49eb7c101eebde4d830ae86ea71 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 16:04:04 -0700 Subject: [PATCH 012/123] fix(ui): probe auto-router tiers via real proxy routing, not /health/test_connection Live testing showed the first cut was broken: /health/test_connection merges {...configParams, ...requestParams}, so passing the public model_group name as the request model overrode the resolved provider model and every tier failed with "LLM Provider NOT provided". The frontend only has the public group name, not the underlying litellm_params, so it cannot build the request that endpoint needs. Switch to testing each model group the way production actually routes it: send a minimal request to /v1/chat/completions (or /v1/embeddings for the embedding model) by public group name through the shared apiClient. The router resolves the group, credentials, and provider itself, so a green row means the tier is genuinely reachable. Verified live: voyage embedding returns 200, a tier with a bad key returns the real provider auth error. Also address Greptile feedback: rows now update progressively as each probe settles instead of all at once, and TIER_ORDER is derived through a `satisfies Record` guard so adding a tier without listing it is a compile error. --- .../auto_router_connection_test.test.tsx | 32 ++++++----- .../add_model/auto_router_connection_test.tsx | 55 +++++-------------- .../build_auto_router_test_targets.ts | 9 ++- .../src/components/networking.tsx | 28 ++++++++++ 4 files changed, 69 insertions(+), 55 deletions(-) diff --git a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx index 9b872d5edee..b07270b5ced 100644 --- a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.test.tsx @@ -7,11 +7,11 @@ vi.mock("../networking", async () => { const actual = await vi.importActual("../networking"); return { ...actual, - testConnectionRequest: vi.fn(), + testModelGroupConnection: vi.fn(), }; }); -const getMock = async () => vi.mocked((await import("../networking")).testConnectionRequest); +const getMock = async () => vi.mocked((await import("../networking")).testModelGroupConnection); const targets: AutoRouterTestTarget[] = [ { labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }, @@ -32,12 +32,12 @@ describe("AutoRouterConnectionTest", () => { await waitFor(() => expect(mock).toHaveBeenCalledTimes(3)); - expect(mock).toHaveBeenCalledWith("sk-test", { model: "gpt-4o-mini" }, {}, "chat"); - expect(mock).toHaveBeenCalledWith("sk-test", { model: "claude-sonnet-4" }, {}, "chat"); - expect(mock).toHaveBeenCalledWith("sk-test", { model: "voyage-3-5" }, {}, "embedding"); + expect(mock).toHaveBeenCalledWith("sk-test", "gpt-4o-mini", "chat"); + expect(mock).toHaveBeenCalledWith("sk-test", "claude-sonnet-4", "chat"); + expect(mock).toHaveBeenCalledWith("sk-test", "voyage-3-5", "embedding"); }); - it("shows a success indicator per target when the health check passes", async () => { + it("shows a success indicator per target when the routing probe passes", async () => { const mock = await getMock(); mock.mockResolvedValue({ status: "success" }); @@ -48,12 +48,14 @@ describe("AutoRouterConnectionTest", () => { expect(screen.getByText("MEDIUM, COMPLEX")).toBeInTheDocument(); }); - it("renders the provider error message for a failing target while others pass", async () => { + it("renders the provider error message (litellm prefix stripped) for a failing target while others pass", async () => { const mock = await getMock(); - mock.mockImplementation((_token, litellmParams) => - litellmParams.model === "claude-sonnet-4" - ? Promise.resolve({ status: "error", result: { error: "litellm.AuthenticationError: invalid api key" } }) - : Promise.resolve({ status: "success" }), + mock.mockImplementation((_token, modelGroup) => + Promise.resolve( + modelGroup === "claude-sonnet-4" + ? { status: "error", error: "litellm.AuthenticationError: invalid api key" } + : { status: "success" }, + ), ); renderWithProviders(); @@ -64,9 +66,9 @@ describe("AutoRouterConnectionTest", () => { expect(screen.getAllByTestId("test-status-success")).toHaveLength(2); }); - it("surfaces a thrown network error as a failing row", async () => { + it("renders a non-litellm error string verbatim", async () => { const mock = await getMock(); - mock.mockRejectedValue(new Error("Network request failed")); + mock.mockResolvedValue({ status: "error", error: "Connection test failed: 404 Not Found" }); renderWithProviders( { />, ); - await waitFor(() => expect(screen.getByTestId("test-error-message")).toHaveTextContent("Network request failed")); + await waitFor(() => + expect(screen.getByTestId("test-error-message")).toHaveTextContent("Connection test failed: 404 Not Found"), + ); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx index 77588006e35..5badd155da8 100644 --- a/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/auto_router_connection_test.tsx @@ -1,7 +1,7 @@ import React from "react"; import { Typography } from "antd"; import { CheckCircleTwoTone, CloseCircleTwoTone, LoadingOutlined } from "@ant-design/icons"; -import { testConnectionRequest } from "../networking"; +import { testModelGroupConnection, ModelGroupConnectionResult } from "../networking"; import { AutoRouterTestTarget } from "./build_auto_router_test_targets"; const { Text } = Typography; @@ -12,45 +12,13 @@ interface AutoRouterConnectionTestProps { onTestComplete?: () => void; } -type TargetResult = { status: "pending" } | { status: "success" } | { status: "error"; error: string }; - -interface NormalizedResponse { - ok: boolean; - error?: string; -} - -const normalizeTestConnectionResponse = (response: unknown): NormalizedResponse => { - if (typeof response !== "object" || response === null) { - return { ok: false, error: "Unexpected response from connection test" }; - } - const record = response as Record; - if (record.status === "success") { - return { ok: true }; - } - const result = - typeof record.result === "object" && record.result !== null ? (record.result as Record) : {}; - const resultError = typeof result.error === "string" ? result.error : undefined; - const recordMessage = typeof record.message === "string" ? record.message : undefined; - return { ok: false, error: resultError ?? recordMessage ?? "Unknown error" }; -}; +type TargetResult = { status: "pending" } | ModelGroupConnectionResult; const cleanErrorMessage = (error: string): string => { const mainError = error.split("stack trace:")[0].trim(); return mainError.replace(/^litellm\.(.*?)Error: /, ""); }; -const runTarget = async (accessToken: string, target: AutoRouterTestTarget): Promise => { - try { - const response = await testConnectionRequest(accessToken, { model: target.modelGroup }, {}, target.mode); - const normalized = normalizeTestConnectionResponse(response); - return normalized.ok - ? { status: "success" } - : { status: "error", error: cleanErrorMessage(normalized.error ?? "Unknown error") }; - } catch (error) { - return { status: "error", error: cleanErrorMessage(error instanceof Error ? error.message : String(error)) }; - } -}; - const AutoRouterConnectionTest: React.FC = ({ accessToken, targets, @@ -61,16 +29,22 @@ const AutoRouterConnectionTest: React.FC = ({ React.useEffect(() => { let cancelled = false; const run = async () => { - const settled = await Promise.all(targets.map((target) => runTarget(accessToken, target))); - if (cancelled) return; - setResults(settled); - if (onTestComplete) onTestComplete(); + await Promise.all( + targets.map(async (target, index) => { + const result = await testModelGroupConnection(accessToken, target.modelGroup, target.mode); + if (cancelled) return; + const cleaned: TargetResult = + result.status === "error" ? { status: "error", error: cleanErrorMessage(result.error) } : result; + setResults((prev) => prev.map((r, i) => (i === index ? cleaned : r))); + }), + ); + if (!cancelled && onTestComplete) onTestComplete(); }; run(); return () => { cancelled = true; }; - // eslint-disable-next-line react-hooks/exhaustive-deps -- probes run once per mount; the parent remounts via `key` to start a fresh test, and re-running on prop identity changes would refire paid health checks + // eslint-disable-next-line react-hooks/exhaustive-deps -- probes run once per mount; the parent remounts via `key` to start a fresh test, and re-running on prop identity changes would refire paid requests }, []); if (targets.length === 0) { @@ -80,7 +54,8 @@ const AutoRouterConnectionTest: React.FC = ({ return (
- Each configured tier routes to a saved model group. Test Connection runs a live health check against each one. + Each configured tier routes to a saved model group. Test Connection sends a minimal request through the proxy to + each one, exactly as the auto router would. {targets.map((target, index) => { const result = results[index] ?? { status: "pending" }; diff --git a/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts b/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts index 0104b6bc9c5..b2a3cc10012 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_auto_router_test_targets.ts @@ -14,7 +14,14 @@ export interface BuildAutoRouterTestTargetsParams { embeddingModel: string | undefined; } -const TIER_ORDER: (keyof ComplexityTiers)[] = ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]; +// Keys drive iteration order; `satisfies Record` makes it a +// compile error to add a tier to ComplexityTiers without listing it here (and vice versa). +const TIER_ORDER = Object.keys({ + SIMPLE: null, + MEDIUM: null, + COMPLEX: null, + REASONING: null, +} satisfies Record) as (keyof ComplexityTiers)[]; export const buildAutoRouterTestTargets = ({ tiers, diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index da6bb079876..ae263727569 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2315,6 +2315,34 @@ export const testConnectionRequest = async ( } }; +export type ModelGroupConnectionResult = { status: "success" } | { status: "error"; error: string }; + +/** + * Test an existing model group by routing a minimal request through the proxy + * exactly as production would (by public model_group name). Unlike + * /health/test_connection, this needs no litellm_params resolution: the router + * resolves the group, credentials, and provider. Used by the auto-router Test + * Connection to probe each tier's model group and the embedding model. + */ +export const testModelGroupConnection = async ( + accessToken: string, + modelGroup: string, + mode: "chat" | "embedding", +): Promise => { + const path = mode === "embedding" ? "/v1/embeddings" : "/v1/chat/completions"; + const body = + mode === "embedding" + ? { model: modelGroup, input: "test from litellm" } + : { model: modelGroup, messages: [{ role: "user", content: "test from litellm" }], max_tokens: 1 }; + + try { + await apiClient.post(path, { accessToken, body }); + return { status: "success" }; + } catch (error) { + return { status: "error", error: error instanceof Error ? error.message : String(error) }; + } +}; + // ... existing code ... export const keyInfoV1Call = async (accessToken: string, key: string) => { try { From e0463a38fffdeb32dec9297039ede25a113ae543 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 16:10:39 -0700 Subject: [PATCH 013/123] fix(completion): forward aws credential kwargs into litellm_params so the responses bridge keeps WIF auth Chat-completions requests to responses-only Bedrock Mantle models are bridged to the Responses API, but completion() forwarded only aws_bedrock_project_id into get_litellm_params, so aws_role_name, aws_web_identity_token, aws_session_name and the other SigV4 credential kwargs never reached sign_request and botocore fell back to the default credential chain ("Bedrock Mantle auth failed: no Bearer token and no usable AWS credentials"). Forward the whole AWS credential kwarg family, extracted from the OPTIONAL_KWARGS_KEYS set get_litellm_params already supports. --- .../litellm_core_utils/get_litellm_params.py | 56 ++++++++++------- litellm/main.py | 7 ++- tests/test_litellm/test_main.py | 62 +++++++++++++++++++ 3 files changed, 99 insertions(+), 26 deletions(-) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 352e55e9c23..b8ef9d8cca7 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -2,26 +2,8 @@ from typing import Optional from litellm.llms.openai.data_residency import infer_openai_data_residency -# Pre-define optional kwargs keys as frozenset for O(1) lookups -# These are extracted from kwargs only if present, avoiding unnecessary .get() calls -OPTIONAL_KWARGS_KEYS = frozenset( +AWS_CREDENTIAL_KWARGS_KEYS = frozenset( { - "azure_ad_token", - "tenant_id", - "client_id", - "client_secret", - "azure_username", - "azure_password", - "azure_scope", - "timeout", - "gcs_bucket_name", - "bucket_name", - "vertex_credentials", - "vertex_project", - "vertex_location", - "vertex_ai_project", - "vertex_ai_location", - "vertex_ai_credentials", "aws_region_name", "aws_access_key_id", "aws_secret_access_key", @@ -34,14 +16,40 @@ OPTIONAL_KWARGS_KEYS = frozenset( "aws_external_id", "aws_bedrock_runtime_endpoint", "aws_bedrock_project_id", - "tpm", - "rpm", - "itpm", - "otpm", - "use_xai_oauth", } ) +# Pre-define optional kwargs keys as frozenset for O(1) lookups +# These are extracted from kwargs only if present, avoiding unnecessary .get() calls +OPTIONAL_KWARGS_KEYS = ( + frozenset( + { + "azure_ad_token", + "tenant_id", + "client_id", + "client_secret", + "azure_username", + "azure_password", + "azure_scope", + "timeout", + "gcs_bucket_name", + "bucket_name", + "vertex_credentials", + "vertex_project", + "vertex_location", + "vertex_ai_project", + "vertex_ai_location", + "vertex_ai_credentials", + "tpm", + "rpm", + "itpm", + "otpm", + "use_xai_oauth", + } + ) + | AWS_CREDENTIAL_KWARGS_KEYS +) + # Backward-compatible alias for existing imports/tests. _OPTIONAL_KWARGS_KEYS = OPTIONAL_KWARGS_KEYS diff --git a/litellm/main.py b/litellm/main.py index 7d457d9cdd1..6fd68921fb0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -92,7 +92,10 @@ from litellm.litellm_core_utils.completion_timeout import CompletionTimeout from litellm.litellm_core_utils.request_timeout_resolver import ( get_configured_request_timeout, ) -from litellm.litellm_core_utils.get_litellm_params import OPTIONAL_KWARGS_KEYS +from litellm.litellm_core_utils.get_litellm_params import ( + AWS_CREDENTIAL_KWARGS_KEYS, + OPTIONAL_KWARGS_KEYS, +) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_provider_specific_headers import ( ProviderSpecificHeaderUtils, @@ -5322,7 +5325,7 @@ def completion( # type: ignore tpm=kwargs.get("tpm"), rpm=kwargs.get("rpm"), use_xai_oauth=kwargs.get("use_xai_oauth", False), - aws_bedrock_project_id=kwargs.get("aws_bedrock_project_id"), + **{key: kwargs[key] for key in AWS_CREDENTIAL_KWARGS_KEYS if key in kwargs}, ) cast(LiteLLMLoggingObj, logging).update_environment_variables( model=model, diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 28cf4fa0744..c5e70e0fabf 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -2081,3 +2081,65 @@ def test_stream_chunk_builder_text_completion_combines_text_and_usage(): assert response.usage.prompt_tokens > 0 assert response.usage.completion_tokens > 0 assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens + + +@pytest.mark.asyncio +async def test_acompletion_forwards_aws_credentials_through_responses_bridge( + respx_mock: respx.MockRouter, monkeypatch +): + from botocore.credentials import Credentials + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + original_disable_aiohttp = litellm.disable_aiohttp_transport + try: + litellm.disable_aiohttp_transport = True + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + litellm.in_memory_llm_clients_cache.flush_cache() + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + get_credentials_mock = MagicMock(return_value=Credentials("fake-key", "fake-secret")) + monkeypatch.setattr(BaseAWSLLM, "get_credentials", get_credentials_mock) + + respx_mock.post("https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses").respond( + json={ + "id": "resp_123", + "object": "response", + "created_at": 1760144904, + "status": "completed", + "model": "openai.gpt-5.4", + "output": [ + { + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "ok", "annotations": []}], + } + ], + } + ) + + response = await litellm.acompletion( + model="bedrock_mantle/openai.gpt-5.4", + messages=[{"role": "user", "content": "hi"}], + api_base="https://bedrock-mantle.us-east-2.api.aws/v1", + aws_region_name="us-east-2", + aws_session_name="litellm-gcp", + aws_role_name="arn:aws:iam::123456789012:role/litellm-bedrock-role", + aws_web_identity_token="oidc/google/108963886734710037768", + num_retries=0, + ) + + assert response.choices[0].message.content == "ok" + credential_kwargs = get_credentials_mock.call_args.kwargs + assert credential_kwargs["aws_role_name"] == "arn:aws:iam::123456789012:role/litellm-bedrock-role" + assert credential_kwargs["aws_web_identity_token"] == "oidc/google/108963886734710037768" + assert credential_kwargs["aws_session_name"] == "litellm-gcp" + authorization = respx_mock.calls.last.request.headers["Authorization"] + assert authorization.startswith("AWS4-HMAC-SHA256") + assert "fake-key" in authorization + finally: + litellm.disable_aiohttp_transport = original_disable_aiohttp + litellm.in_memory_llm_clients_cache.flush_cache() From b9aef1b8106dd514372ffd82b5a2a9aee5f5def6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 11:45:20 -0700 Subject: [PATCH 014/123] test(e2e): cover key rpm/tpm rate limiting, window reset, and pacing headers --- tests/e2e/CLAUDE.md | 2 +- tests/e2e/e2e_http.py | 11 +- tests/e2e/management/test_management_e2e.py | 5 +- tests/e2e/models.py | 1 + .../quota_management/ratelimit/conftest.py | 15 ++ .../ratelimit/quota_client.py | 32 ++++ .../ratelimit/test_rate_limit_e2e.py | 149 ++++++++++++++++++ 7 files changed, 210 insertions(+), 5 deletions(-) create mode 100644 tests/e2e/quota_management/ratelimit/conftest.py create mode 100644 tests/e2e/quota_management/ratelimit/quota_client.py create mode 100644 tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 88992038cb7..6b1b9950656 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -11,7 +11,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family - `embeddings/` - the `/embeddings` endpoint across providers - `batches/` - the `/batches` endpoint (placeholder until the first test lands) - `realtime/` - realtime websocket sessions, including the pipecat audio path -- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window) and `spend_tracking/` (spend logging and cost attribution on `/spend/*`) +- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`) - `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials; also the dashboard UI behavior on top of them, driven through the proxy-served UI at /ui with playwright (optional dep behind importorskip) - `logging/` - logging-integration delivery (datadog and friends) - `security/` - secret handling and log-leak protection diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index 32005faed4b..ad53b2b4aa2 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -109,14 +109,16 @@ class StreamingResponse(BaseModel): """Raw outcome for calls whose body is provider-native or streamed: status, the x-litellm-call-id header, the x-litellm-response-cost header (StandardLogging response_cost), the content-type (which tells streaming `text/event-stream` from - non-streaming `application/json`), and the body. SpendLogs.request_id is the - completion body id, not call_id. Used by passthrough and streaming, where one - validated JSON model does not fit.""" + non-streaming `application/json`), the response headers (lowercased names, e.g. + the x-ratelimit-* pacing headers and retry-after on a 429), and the body. + SpendLogs.request_id is the completion body id, not call_id. Used by passthrough + and streaming, where one validated JSON model does not fit.""" status_code: int call_id: str | None = None # x-litellm-call-id header response_cost: float | None = None # x-litellm-response-cost header content_type: str | None = None + headers: dict[str, str] = {} body: str chunks: int = 0 # streamed events (0 for non-streaming) @@ -276,12 +278,14 @@ def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingRespon call_id = _hdr(resp, "x-litellm-call-id") response_cost = _parse_response_cost(resp) content_type = _hdr(resp, "content-type") + headers = {name.lower(): value for name, value in resp.headers.items()} if not stream or not (200 <= resp.status_code < 300): return StreamingResponse( status_code=resp.status_code, call_id=call_id, response_cost=response_cost, content_type=content_type, + headers=headers, body=resp.text, ) lines = cast("Iterator[bytes]", resp.iter_lines()) @@ -291,6 +295,7 @@ def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingRespon call_id=call_id, response_cost=response_cost, content_type=content_type, + headers=headers, body="", chunks=chunks, ) diff --git a/tests/e2e/management/test_management_e2e.py b/tests/e2e/management/test_management_e2e.py index cbd5db0d59f..7a05a5c520f 100644 --- a/tests/e2e/management/test_management_e2e.py +++ b/tests/e2e/management/test_management_e2e.py @@ -111,7 +111,7 @@ class TestKeyRoutes: key = _generate_key( client, resources, - KeyGenerateBody(models=["gemini-2.5-flash"], key_alias=alias, tpm_limit=424242), + KeyGenerateBody(models=["gemini-2.5-flash"], key_alias=alias, tpm_limit=424242, rpm_limit=424243), ) info = client.gateway.key_info(key) @@ -122,6 +122,9 @@ class TestKeyRoutes: assert info.tpm_limit == 424242, ( f"/key/info reports tpm_limit {info.tpm_limit}, configured 424242" ) + assert info.rpm_limit == 424243, ( + f"/key/info reports rpm_limit {info.rpm_limit}, configured 424243" + ) _poll_chat_ok(client, key, "gemini-2.5-flash") _assert_model_denied( diff --git a/tests/e2e/models.py b/tests/e2e/models.py index e32f2709181..ebd082161e0 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -82,6 +82,7 @@ class KeyInfo(BaseModel): key_alias: str | None = None models: list[str] = [] tpm_limit: int | None = None + rpm_limit: int | None = None team_id: str | None = None spend: float | None = None max_budget: float | None = None diff --git a/tests/e2e/quota_management/ratelimit/conftest.py b/tests/e2e/quota_management/ratelimit/conftest.py new file mode 100644 index 00000000000..4a5a73bb5e4 --- /dev/null +++ b/tests/e2e/quota_management/ratelimit/conftest.py @@ -0,0 +1,15 @@ +"""Quota-management suite's `client` fixture. + +The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker +live in the parent tests/e2e/conftest.py. QuotaClient holds the shared Gateway, +so the `resources` fixture cleans up keys through it. +""" + +import pytest + +from quota_client import QuotaClient, build_client + + +@pytest.fixture(scope="session") +def client() -> QuotaClient: + return build_client() diff --git a/tests/e2e/quota_management/ratelimit/quota_client.py b/tests/e2e/quota_management/ratelimit/quota_client.py new file mode 100644 index 00000000000..806ab1d1557 --- /dev/null +++ b/tests/e2e/quota_management/ratelimit/quota_client.py @@ -0,0 +1,32 @@ +"""Client for the quota-management suite: the shared Gateway plus raw chat +calls judged by HTTP status, body, and headers (a rate-limit block is a 429 +whose body and retry-after header carry the contract, not a typed success +model).""" + +from __future__ import annotations + +from dataclasses import dataclass + +from e2e_gateway import Gateway, build_gateway +from e2e_http import StreamingResponse +from models import ChatBody, ChatMessage + + +@dataclass(frozen=True, slots=True) +class QuotaClient: + gateway: Gateway + + def chat(self, key: str, model: str, content: str, *, max_tokens: int = 16) -> StreamingResponse: + return self.gateway.transport.send( + "/chat/completions", + headers=self.gateway.transport.bearer(key), + json=ChatBody( + model=model, + messages=[ChatMessage(role="user", content=content)], + max_tokens=max_tokens, + ), + ) + + +def build_client() -> QuotaClient: + return QuotaClient(gateway=build_gateway()) diff --git a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py new file mode 100644 index 00000000000..a66d7974fc5 --- /dev/null +++ b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py @@ -0,0 +1,149 @@ +"""Live e2e: key-level rpm/tpm rate limits on the gateway. + +Covers quota_management.ratelimit.*: a key generated with rpm_limit/tpm_limit gets a +429 once the limit is crossed inside one window (blocks_over_limit), serves +again once the window rolls (resets_after_window), and successful responses +report x-ratelimit-* limit/remaining headers so clients can pace +(headers_report_remaining). Each test asserts both halves of the contract: the +recorded state (/key/info echoes the configured limit) and the enforced +behavior (the 429, the recovery, or the headers on live traffic). + +The v3 limiter counts a request against the rpm budget at the pre-call hook, +before model routing, so every call that clears auth consumes budget whether or +not it ultimately succeeds. All calls of one test must land inside a single +window (LITELLM_RATE_LIMIT_WINDOW_SIZE, 60s default), which real chat latency +comfortably allows. +""" + +from __future__ import annotations + +import time + +import pytest + +from e2e_config import unique_marker +from e2e_http import StreamingResponse, require_successful_call +from lifecycle import ResourceManager +from models import KeyGenerateBody +from quota_client import QuotaClient + +pytestmark = pytest.mark.e2e + +MODEL = "claude-haiku-4-5" + + +def _limited_key( + client: QuotaClient, + resources: ResourceManager, + *, + rpm_limit: int | None = None, + tpm_limit: int | None = None, +) -> str: + key = client.gateway.generate_key(KeyGenerateBody(models=[MODEL], rpm_limit=rpm_limit, tpm_limit=tpm_limit)) + resources.defer(lambda: client.gateway.delete_key(key)) + return key + + +def _chat(client: QuotaClient, key: str) -> StreamingResponse: + return client.chat(key, MODEL, f"reply with one word {unique_marker()}") + + +def _first_ok(client: QuotaClient, key: str) -> StreamingResponse: + """First successful call on a fresh key, which opens the rate-limit window. + A fresh key may briefly 401 until the data plane's auth cache picks it up, so + retry on 401 to a deadline; a 401 never reaches the rate limiter, so only the + successful call consumes budget. Any other failure is behavior under test and + fails hard.""" + deadline = time.monotonic() + client.gateway.poll_timeout + while True: + outcome = _chat(client, key) + if outcome.ok: + return outcome + if outcome.status_code != 401 or time.monotonic() >= deadline: + require_successful_call(outcome) + time.sleep(client.gateway.poll_interval) + + +def _assert_rate_limited(outcome: StreamingResponse, limit_type: str) -> None: + assert outcome.status_code == 429, ( + f"expected a 429 {limit_type} rate-limit block, got {outcome.status_code}: {outcome.body[:300]}" + ) + assert "Rate limit exceeded for api_key" in outcome.body, ( + f"429 body must name the api_key scope, got: {outcome.body[:300]}" + ) + assert f"Limit type: {limit_type}" in outcome.body, ( + f"429 body must carry 'Limit type: {limit_type}', got: {outcome.body[:300]}" + ) + retry_after = outcome.headers.get("retry-after") + assert retry_after is not None and retry_after.isdigit() and int(retry_after) > 0, ( + f"429 must carry a positive integer retry-after header, got {retry_after!r}" + ) + + +class TestKeyRateLimits: + @pytest.mark.covers("quota_management.ratelimit.rpm.blocks_over_limit") + def test_rpm_limit_blocks_over_limit(self, client: QuotaClient, resources: ResourceManager) -> None: + key = _limited_key(client, resources, rpm_limit=3) + info = client.gateway.key_info(key) + assert info.rpm_limit == 3, f"/key/info reports rpm_limit {info.rpm_limit}, configured 3" + + _ = _first_ok(client, key) + for _ in range(2): + require_successful_call(_chat(client, key)) + + _assert_rate_limited(_chat(client, key), "requests") + + @pytest.mark.covers("quota_management.ratelimit.tpm.blocks_over_limit") + def test_tpm_limit_blocks_over_limit(self, client: QuotaClient, resources: ResourceManager) -> None: + key = _limited_key(client, resources, tpm_limit=60) + info = client.gateway.key_info(key) + assert info.tpm_limit == 60, f"/key/info reports tpm_limit {info.tpm_limit}, configured 60" + + _ = _first_ok(client, key) + for _ in range(8): + outcome = _chat(client, key) + if outcome.status_code == 429: + _assert_rate_limited(outcome, "tokens") + return + require_successful_call(outcome) + pytest.fail("tpm_limit=60 was never enforced with a 429 within 8 calls of ~25 tokens each") + + @pytest.mark.covers("quota_management.ratelimit.rpm.resets_after_window") + def test_rpm_limit_resets_after_window(self, client: QuotaClient, resources: ResourceManager) -> None: + key = _limited_key(client, resources, rpm_limit=1) + + _ = _first_ok(client, key) + _assert_rate_limited(_chat(client, key), "requests") + + deadline = time.monotonic() + client.gateway.poll_timeout + while time.monotonic() < deadline: + outcome = _chat(client, key) + if outcome.ok: + return + assert outcome.status_code == 429, ( + f"while the window drains only 429s are acceptable, got {outcome.status_code}: {outcome.body[:300]}" + ) + time.sleep(client.gateway.poll_interval) + pytest.fail("a blocked key never recovered after the rate-limit window elapsed") + + @pytest.mark.covers("quota_management.ratelimit.rpm.headers_report_remaining") + def test_headers_report_limit_and_remaining(self, client: QuotaClient, resources: ResourceManager) -> None: + key = _limited_key(client, resources, rpm_limit=5, tpm_limit=100000) + + first = _first_ok(client, key) + assert first.headers.get("x-ratelimit-api_key-limit-requests") == "5", ( + f"success response must report the key's request limit, headers: " + f"{ {k: v for k, v in first.headers.items() if 'ratelimit' in k} }" + ) + assert first.headers.get("x-ratelimit-api_key-remaining-requests") == "4", ( + f"first call against rpm_limit=5 must leave 4 remaining, got " + f"{first.headers.get('x-ratelimit-api_key-remaining-requests')!r}" + ) + assert first.headers.get("x-ratelimit-api_key-limit-tokens") == "100000", ( + f"success response must report the key's token limit, got " + f"{first.headers.get('x-ratelimit-api_key-limit-tokens')!r}" + ) + remaining_tokens = first.headers.get("x-ratelimit-api_key-remaining-tokens") + assert remaining_tokens is not None and remaining_tokens.isdigit() and int(remaining_tokens) < 100000, ( + f"one call must leave remaining tokens reported and below the limit, got {remaining_tokens!r}" + ) From 1ca1a9cd03bb740afe680b7e8834f89d81f6c045 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 13:10:40 -0700 Subject: [PATCH 015/123] test(e2e): assert the tpm block at its exact token crossing instead of a call-count heuristic --- .../ratelimit/test_rate_limit_e2e.py | 74 +++++++++++++++++-- 1 file changed, 66 insertions(+), 8 deletions(-) diff --git a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py index a66d7974fc5..089a189a1e7 100644 --- a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py @@ -10,16 +10,24 @@ behavior (the 429, the recovery, or the headers on live traffic). The v3 limiter counts a request against the rpm budget at the pre-call hook, before model routing, so every call that clears auth consumes budget whether or -not it ultimately succeeds. All calls of one test must land inside a single -window (LITELLM_RATE_LIMIT_WINDOW_SIZE, 60s default), which real chat latency +not it ultimately succeeds. The tpm budget is reserved pre-call from an estimate +(message chars // 4 + max_tokens) and reconciled to the body's actual +usage.total_tokens after the call, so a block may legitimately fire before the +actual spend crosses the limit; the tpm test asserts the exact contract on both +sides (a 429 only once the blocked call's reservation exceeds the remaining +budget, and no later than the first call after actual spend reaches the limit). +All calls of one test must land inside a single window +(LITELLM_RATE_LIMIT_WINDOW_SIZE, 60s default), which real chat latency comfortably allows. """ from __future__ import annotations +import re import time import pytest +from pydantic import BaseModel, ConfigDict, ValidationError from e2e_config import unique_marker from e2e_http import StreamingResponse, require_successful_call @@ -30,6 +38,40 @@ from quota_client import QuotaClient pytestmark = pytest.mark.e2e MODEL = "claude-haiku-4-5" +TPM_LIMIT = 60 +CHAT_MAX_TOKENS = 16 +RESERVATION_CHARS_PER_TOKEN = 4 +WINDOW_SECONDS = 60 + + +class _ChatUsage(BaseModel): + model_config = ConfigDict(extra="ignore") + + total_tokens: int + + +class _ChatBodyWithUsage(BaseModel): + model_config = ConfigDict(extra="ignore") + + usage: _ChatUsage + + +def _total_tokens(outcome: StreamingResponse) -> int: + try: + return _ChatBodyWithUsage.model_validate_json(outcome.body).usage.total_tokens + except ValidationError: + pytest.fail(f"successful chat body must report usage.total_tokens, got: {outcome.body[:300]}") + + +def _reserved_tokens(content: str) -> int: + return max(1, len(content) // RESERVATION_CHARS_PER_TOKEN) + CHAT_MAX_TOKENS + + +def _remaining_from_429(body: str) -> int: + found = re.search(r"Remaining: (\d+)", body) + if found is None: + pytest.fail(f"429 body must report the remaining budget, got: {body[:300]}") + return int(found.group(1)) def _limited_key( @@ -95,18 +137,34 @@ class TestKeyRateLimits: @pytest.mark.covers("quota_management.ratelimit.tpm.blocks_over_limit") def test_tpm_limit_blocks_over_limit(self, client: QuotaClient, resources: ResourceManager) -> None: - key = _limited_key(client, resources, tpm_limit=60) + key = _limited_key(client, resources, tpm_limit=TPM_LIMIT) info = client.gateway.key_info(key) - assert info.tpm_limit == 60, f"/key/info reports tpm_limit {info.tpm_limit}, configured 60" + assert info.tpm_limit == TPM_LIMIT, f"/key/info reports tpm_limit {info.tpm_limit}, configured {TPM_LIMIT}" - _ = _first_ok(client, key) - for _ in range(8): - outcome = _chat(client, key) + first = _first_ok(client, key) + window_deadline = time.monotonic() + WINDOW_SECONDS - 10 + spent = _total_tokens(first) + + while spent < TPM_LIMIT: + assert time.monotonic() < window_deadline, ( + f"spent only {spent} of {TPM_LIMIT} tokens before the {WINDOW_SECONDS}s window could roll; " + "the exact-crossing assertion needs every call inside one window" + ) + content = f"reply with one word {unique_marker()}" + outcome = client.chat(key, MODEL, content, max_tokens=CHAT_MAX_TOKENS) if outcome.status_code == 429: _assert_rate_limited(outcome, "tokens") + remaining = _remaining_from_429(outcome.body) + reserved = _reserved_tokens(content) + assert reserved > remaining, ( + f"blocked while the call still fit: {remaining} of {TPM_LIMIT} tokens remained but the call " + f"reserved only {reserved} ({spent} actual tokens spent so far)" + ) return require_successful_call(outcome) - pytest.fail("tpm_limit=60 was never enforced with a 429 within 8 calls of ~25 tokens each") + spent += _total_tokens(outcome) + + _assert_rate_limited(_chat(client, key), "tokens") @pytest.mark.covers("quota_management.ratelimit.rpm.resets_after_window") def test_rpm_limit_resets_after_window(self, client: QuotaClient, resources: ResourceManager) -> None: From 6f2435b54aa4e7a62a23140e2b59b9bdb8933d1b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 13:17:54 -0700 Subject: [PATCH 016/123] test(e2e): fail the rpm reset test when the limiter resets early --- .../ratelimit/test_rate_limit_e2e.py | 48 ++++++++++++++----- 1 file changed, 36 insertions(+), 12 deletions(-) diff --git a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py index 089a189a1e7..a6e35c5f3d5 100644 --- a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py @@ -2,7 +2,7 @@ Covers quota_management.ratelimit.*: a key generated with rpm_limit/tpm_limit gets a 429 once the limit is crossed inside one window (blocks_over_limit), serves -again once the window rolls (resets_after_window), and successful responses +again once the window rolls and no sooner (resets_after_window), and successful responses report x-ratelimit-* limit/remaining headers so clients can pace (headers_report_remaining). Each test asserts both halves of the contract: the recorded state (/key/info echoes the configured limit) and the enforced @@ -19,12 +19,20 @@ budget, and no later than the first call after actual spend reaches the limit). All calls of one test must land inside a single window (LITELLM_RATE_LIMIT_WINDOW_SIZE, 60s default), which real chat latency comfortably allows. + +The window opens at the pre-call hook of the first counted call, which happens +after that call is sent, so the send timestamp of the winning first call is a +lower bound on the window start. The reset test uses it to reject an early +reset: recovery must not arrive before the full window has elapsed from that +send, less a small tolerance for the limiter's integer-second window +arithmetic. """ from __future__ import annotations import re import time +from dataclasses import dataclass import pytest from pydantic import BaseModel, ConfigDict, ValidationError @@ -42,6 +50,13 @@ TPM_LIMIT = 60 CHAT_MAX_TOKENS = 16 RESERVATION_CHARS_PER_TOKEN = 4 WINDOW_SECONDS = 60 +RESET_TOLERANCE_SECONDS = 5 + + +@dataclass(frozen=True, slots=True) +class _FirstOk: + sent_at: float + response: StreamingResponse class _ChatUsage(BaseModel): @@ -90,17 +105,19 @@ def _chat(client: QuotaClient, key: str) -> StreamingResponse: return client.chat(key, MODEL, f"reply with one word {unique_marker()}") -def _first_ok(client: QuotaClient, key: str) -> StreamingResponse: - """First successful call on a fresh key, which opens the rate-limit window. - A fresh key may briefly 401 until the data plane's auth cache picks it up, so - retry on 401 to a deadline; a 401 never reaches the rate limiter, so only the - successful call consumes budget. Any other failure is behavior under test and - fails hard.""" +def _first_ok(client: QuotaClient, key: str) -> _FirstOk: + """First successful call on a fresh key, which opens the rate-limit window; + `sent_at` is captured just before the winning send, so the window opened no + earlier than it. A fresh key may briefly 401 until the data plane's auth + cache picks it up, so retry on 401 to a deadline; a 401 never reaches the + rate limiter, so only the successful call consumes budget. Any other failure + is behavior under test and fails hard.""" deadline = time.monotonic() + client.gateway.poll_timeout while True: + sent_at = time.monotonic() outcome = _chat(client, key) if outcome.ok: - return outcome + return _FirstOk(sent_at=sent_at, response=outcome) if outcome.status_code != 401 or time.monotonic() >= deadline: require_successful_call(outcome) time.sleep(client.gateway.poll_interval) @@ -142,8 +159,8 @@ class TestKeyRateLimits: assert info.tpm_limit == TPM_LIMIT, f"/key/info reports tpm_limit {info.tpm_limit}, configured {TPM_LIMIT}" first = _first_ok(client, key) - window_deadline = time.monotonic() + WINDOW_SECONDS - 10 - spent = _total_tokens(first) + window_deadline = first.sent_at + WINDOW_SECONDS - 10 + spent = _total_tokens(first.response) while spent < TPM_LIMIT: assert time.monotonic() < window_deadline, ( @@ -170,13 +187,20 @@ class TestKeyRateLimits: def test_rpm_limit_resets_after_window(self, client: QuotaClient, resources: ResourceManager) -> None: key = _limited_key(client, resources, rpm_limit=1) - _ = _first_ok(client, key) + first = _first_ok(client, key) _assert_rate_limited(_chat(client, key), "requests") deadline = time.monotonic() + client.gateway.poll_timeout while time.monotonic() < deadline: + attempt_sent_at = time.monotonic() outcome = _chat(client, key) if outcome.ok: + window_age = attempt_sent_at - first.sent_at + assert window_age >= WINDOW_SECONDS - RESET_TOLERANCE_SECONDS, ( + f"the key recovered {window_age:.1f}s after the window opened, before the " + f"{WINDOW_SECONDS}s window (less {RESET_TOLERANCE_SECONDS}s tolerance) elapsed; " + "the limiter reset early instead of after the window" + ) return assert outcome.status_code == 429, ( f"while the window drains only 429s are acceptable, got {outcome.status_code}: {outcome.body[:300]}" @@ -188,7 +212,7 @@ class TestKeyRateLimits: def test_headers_report_limit_and_remaining(self, client: QuotaClient, resources: ResourceManager) -> None: key = _limited_key(client, resources, rpm_limit=5, tpm_limit=100000) - first = _first_ok(client, key) + first = _first_ok(client, key).response assert first.headers.get("x-ratelimit-api_key-limit-requests") == "5", ( f"success response must report the key's request limit, headers: " f"{ {k: v for k, v in first.headers.items() if 'ratelimit' in k} }" From 9f55925c31b66494391c20a372bd4f1da40819d7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 13:19:08 -0700 Subject: [PATCH 017/123] test(e2e): name the tpm window deadline's latency margin --- tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py index a6e35c5f3d5..c326f5019a3 100644 --- a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py @@ -51,6 +51,7 @@ CHAT_MAX_TOKENS = 16 RESERVATION_CHARS_PER_TOKEN = 4 WINDOW_SECONDS = 60 RESET_TOLERANCE_SECONDS = 5 +LAST_CALL_LATENCY_MARGIN_SECONDS = 10 @dataclass(frozen=True, slots=True) @@ -159,7 +160,7 @@ class TestKeyRateLimits: assert info.tpm_limit == TPM_LIMIT, f"/key/info reports tpm_limit {info.tpm_limit}, configured {TPM_LIMIT}" first = _first_ok(client, key) - window_deadline = first.sent_at + WINDOW_SECONDS - 10 + window_deadline = first.sent_at + WINDOW_SECONDS - LAST_CALL_LATENCY_MARGIN_SECONDS spent = _total_tokens(first.response) while spent < TPM_LIMIT: From ca55fc2deb0f974992c29e581f205a357f668b44 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 13:24:59 -0700 Subject: [PATCH 018/123] refactor(e2e): model the tpm spend loop's two outcomes as values --- .../ratelimit/test_rate_limit_e2e.py | 52 +++++++++++++------ 1 file changed, 36 insertions(+), 16 deletions(-) diff --git a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py index c326f5019a3..c837bec19f8 100644 --- a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py @@ -90,6 +90,38 @@ def _remaining_from_429(body: str) -> int: return int(found.group(1)) +@dataclass(frozen=True, slots=True) +class _BlockedByReservation: + outcome: StreamingResponse + content: str + spent: int + + +@dataclass(frozen=True, slots=True) +class _CrossedLimit: + spent: int + + +def _spend_until_blocked_or_crossed(client: QuotaClient, key: str, first: _FirstOk) -> _BlockedByReservation | _CrossedLimit: + """Drive chat traffic, summing each body's actual usage.total_tokens, until + the limiter blocks (which the reservation may do before the actual spend + crosses the limit) or the actual spend reaches the limit.""" + window_deadline = first.sent_at + WINDOW_SECONDS - LAST_CALL_LATENCY_MARGIN_SECONDS + spent = _total_tokens(first.response) + while spent < TPM_LIMIT: + assert time.monotonic() < window_deadline, ( + f"spent only {spent} of {TPM_LIMIT} tokens before the {WINDOW_SECONDS}s window could roll; " + "the exact-crossing assertion needs every call inside one window" + ) + content = f"reply with one word {unique_marker()}" + outcome = client.chat(key, MODEL, content, max_tokens=CHAT_MAX_TOKENS) + if outcome.status_code == 429: + return _BlockedByReservation(outcome=outcome, content=content, spent=spent) + require_successful_call(outcome) + spent += _total_tokens(outcome) + return _CrossedLimit(spent=spent) + + def _limited_key( client: QuotaClient, resources: ResourceManager, @@ -160,17 +192,8 @@ class TestKeyRateLimits: assert info.tpm_limit == TPM_LIMIT, f"/key/info reports tpm_limit {info.tpm_limit}, configured {TPM_LIMIT}" first = _first_ok(client, key) - window_deadline = first.sent_at + WINDOW_SECONDS - LAST_CALL_LATENCY_MARGIN_SECONDS - spent = _total_tokens(first.response) - - while spent < TPM_LIMIT: - assert time.monotonic() < window_deadline, ( - f"spent only {spent} of {TPM_LIMIT} tokens before the {WINDOW_SECONDS}s window could roll; " - "the exact-crossing assertion needs every call inside one window" - ) - content = f"reply with one word {unique_marker()}" - outcome = client.chat(key, MODEL, content, max_tokens=CHAT_MAX_TOKENS) - if outcome.status_code == 429: + match _spend_until_blocked_or_crossed(client, key, first): + case _BlockedByReservation(outcome=outcome, content=content, spent=spent): _assert_rate_limited(outcome, "tokens") remaining = _remaining_from_429(outcome.body) reserved = _reserved_tokens(content) @@ -178,11 +201,8 @@ class TestKeyRateLimits: f"blocked while the call still fit: {remaining} of {TPM_LIMIT} tokens remained but the call " f"reserved only {reserved} ({spent} actual tokens spent so far)" ) - return - require_successful_call(outcome) - spent += _total_tokens(outcome) - - _assert_rate_limited(_chat(client, key), "tokens") + case _CrossedLimit(): + _assert_rate_limited(_chat(client, key), "tokens") @pytest.mark.covers("quota_management.ratelimit.rpm.resets_after_window") def test_rpm_limit_resets_after_window(self, client: QuotaClient, resources: ResourceManager) -> None: From 36dff0f63ad549ed1517faf92d62f590a4ce1d1a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 16:10:33 -0700 Subject: [PATCH 019/123] test(e2e): source the ratelimit suite's model from E2E_CHEAP_ANTHROPIC_MODEL --- tests/e2e/e2e_config.py | 2 ++ tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py | 4 ++-- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 75bd715a23a..cde30bf7841 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -24,6 +24,8 @@ CONTROL_PLANE_BASE_URL = os.environ.get( UI_USERNAME = os.environ.get("E2E_UI_USERNAME", "admin") UI_PASSWORD = os.environ.get("E2E_UI_PASSWORD", MASTER_KEY) +CHEAP_ANTHROPIC_MODEL = os.environ.get("E2E_CHEAP_ANTHROPIC_MODEL", "claude-haiku-4-5") + # Writes on the proxy are eventually consistent (e.g. spend rows flush on # proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once. POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120")) diff --git a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py index c837bec19f8..b60b0f37ed2 100644 --- a/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py +++ b/tests/e2e/quota_management/ratelimit/test_rate_limit_e2e.py @@ -37,7 +37,7 @@ from dataclasses import dataclass import pytest from pydantic import BaseModel, ConfigDict, ValidationError -from e2e_config import unique_marker +from e2e_config import CHEAP_ANTHROPIC_MODEL, unique_marker from e2e_http import StreamingResponse, require_successful_call from lifecycle import ResourceManager from models import KeyGenerateBody @@ -45,7 +45,7 @@ from quota_client import QuotaClient pytestmark = pytest.mark.e2e -MODEL = "claude-haiku-4-5" +MODEL = CHEAP_ANTHROPIC_MODEL TPM_LIMIT = 60 CHAT_MAX_TOKENS = 16 RESERVATION_CHARS_PER_TOKEN = 4 From 201730efe573efb7f41941173b74d08e55b77b80 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 16:35:13 -0700 Subject: [PATCH 020/123] fix(bedrock): allow bedrock-mantle:CreateInference in the web identity session policy --- litellm/llms/bedrock/base_aws_llm.py | 9 ++++ .../test_web_identity_session_policy.py | 48 +++++++++++++++++++ 2 files changed, 57 insertions(+) diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index f449851b76f..df811f8d262 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -877,6 +877,15 @@ class BaseAWSLLM: "Resource": "*", "Condition": {"Bool": {"aws:SecureTransport": "true"}}, }, + { + "Sid": "BedrockMantleLiteLLM", + "Effect": "Allow", + "Action": [ + "bedrock-mantle:CreateInference", + ], + "Resource": "*", + "Condition": {"Bool": {"aws:SecureTransport": "true"}}, + }, ], } assume_role_params = { diff --git a/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py b/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py index 7e9c8a273ae..0cbdc518cc2 100644 --- a/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py +++ b/tests/test_litellm/llms/bedrock/test_web_identity_session_policy.py @@ -158,6 +158,54 @@ class TestClaudePlatformActionsCovered: ) +class TestBedrockMantleActionsCovered: + """LIT-3859: bedrock_mantle inference authorizes against the + ``bedrock-mantle`` action namespace, so the session-policy ceiling + must include it or every Mantle request via OIDC/WIF auth denies + with "no session policy allows the bedrock-mantle:CreateInference + action" even when the role's identity policy grants it.""" + + def test_bedrock_mantle_create_inference_present(self): + policy = _captured_policy() + all_actions: set = set() + for stmt in policy["Statement"]: + stmt_actions = stmt.get("Action") + if isinstance(stmt_actions, str): + all_actions.add(stmt_actions) + elif isinstance(stmt_actions, list): + all_actions.update(stmt_actions) + assert "bedrock-mantle:CreateInference" in all_actions, ( + "bedrock-mantle:CreateInference missing from session policy — " + "bedrock_mantle/* requests will 403 on OIDC/WIF auth" + ) + + def test_bedrock_mantle_statement_allows(self): + policy = _captured_policy() + stmt = _statement_by_sid(policy, "BedrockMantleLiteLLM") + assert stmt["Effect"] == "Allow" + assert stmt["Resource"] == "*" + + def test_no_bedrock_mantle_wildcard(self): + policy = _captured_policy() + stmt = _statement_by_sid(policy, "BedrockMantleLiteLLM") + actions = stmt["Action"] + if isinstance(actions, str): + actions = [actions] + assert "bedrock-mantle:*" not in actions, ( + "session policy must not grant bedrock-mantle:* — " + "the ceiling should match the documented action set" + ) + + def test_bedrock_mantle_statement_carries_secure_transport_condition(self): + policy = _captured_policy() + stmt = _statement_by_sid(policy, "BedrockMantleLiteLLM") + cond = stmt.get("Condition") or {} + assert cond.get("Bool", {}).get("aws:SecureTransport") == "true", ( + "BedrockMantleLiteLLM must require aws:SecureTransport=true " + "to keep parity with the bedrock statement" + ) + + def _make_jwt(payload: dict) -> str: def _segment(data: dict) -> str: return base64.urlsafe_b64encode(json.dumps(data).encode()).rstrip(b"=").decode() From a80f85692b8feec7f25c1459f0da0ab51632bd10 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 16:43:36 -0700 Subject: [PATCH 021/123] fix(ui): drop max_tokens from the auto-router connection probe max_tokens=1 makes reasoning models (o1/o3/...) return a 400 "max_tokens reached" because reasoning tokens count against the cap, so a reachable reasoning tier showed a false failure in Test Connection. Live-verified: o3 400s with the cap and succeeds without it. Extract the request shape into a pure buildModelGroupTestRequest and cover it with a test asserting the chat body carries no max_tokens (or max_completion_tokens), so this regression is caught in unit tests instead of only against a live reasoning model. --- .../src/components/networking.test.ts | 16 +++++++++++++ .../src/components/networking.tsx | 24 ++++++++++++++----- 2 files changed, 34 insertions(+), 6 deletions(-) diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts index 43f85cc8674..b9d04a00f61 100644 --- a/ui/litellm-dashboard/src/components/networking.test.ts +++ b/ui/litellm-dashboard/src/components/networking.test.ts @@ -514,3 +514,19 @@ describe("sessionSpendLogsCall", () => { expect(parsed.searchParams.get("page_size")).toBe("100"); }); }); + +describe("buildModelGroupTestRequest", () => { + it("builds a chat completion request with NO max_tokens (reasoning models 400 on a tiny cap)", () => { + const { path, body } = Networking.buildModelGroupTestRequest("o3", "chat"); + expect(path).toBe("/v1/chat/completions"); + expect(body).toEqual({ model: "o3", messages: [{ role: "user", content: "test from litellm" }] }); + expect(body).not.toHaveProperty("max_tokens"); + expect(body).not.toHaveProperty("max_completion_tokens"); + }); + + it("builds an embeddings request for embedding mode", () => { + const { path, body } = Networking.buildModelGroupTestRequest("text-embedding-3-small", "embedding"); + expect(path).toBe("/v1/embeddings"); + expect(body).toEqual({ model: "text-embedding-3-small", input: "test from litellm" }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index ae263727569..10d7f12604d 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2324,17 +2324,29 @@ export type ModelGroupConnectionResult = { status: "success" } | { status: "erro * resolves the group, credentials, and provider. Used by the auto-router Test * Connection to probe each tier's model group and the embedding model. */ +/** + * Build the minimal request that probes a model group by public name. No + * max_tokens: reasoning models (o1/o3/...) reject a tiny cap with "max_tokens + * reached" because reasoning tokens count against it, which would show a false + * failure for a reachable tier. + */ +export const buildModelGroupTestRequest = ( + modelGroup: string, + mode: "chat" | "embedding", +): { path: string; body: Record } => + mode === "embedding" + ? { path: "/v1/embeddings", body: { model: modelGroup, input: "test from litellm" } } + : { + path: "/v1/chat/completions", + body: { model: modelGroup, messages: [{ role: "user", content: "test from litellm" }] }, + }; + export const testModelGroupConnection = async ( accessToken: string, modelGroup: string, mode: "chat" | "embedding", ): Promise => { - const path = mode === "embedding" ? "/v1/embeddings" : "/v1/chat/completions"; - const body = - mode === "embedding" - ? { model: modelGroup, input: "test from litellm" } - : { model: modelGroup, messages: [{ role: "user", content: "test from litellm" }], max_tokens: 1 }; - + const { path, body } = buildModelGroupTestRequest(modelGroup, mode); try { await apiClient.post(path, { accessToken, body }); return { status: "success" }; From cc27528d4f248b5b6e908f424caa270a2df84513 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 16:47:00 -0700 Subject: [PATCH 022/123] test(main): assert the responses bridge forwards static aws keys as well as web identity params --- tests/test_litellm/test_main.py | 28 +++++++++++++++++++++------- 1 file changed, 21 insertions(+), 7 deletions(-) diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index c5e70e0fabf..4611aafa3c1 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -2084,8 +2084,24 @@ def test_stream_chunk_builder_text_completion_combines_text_and_usage(): @pytest.mark.asyncio +@pytest.mark.parametrize( + "aws_credential_kwargs", + [ + { + "aws_session_name": "litellm-gcp", + "aws_role_name": "arn:aws:iam::123456789012:role/litellm-bedrock-role", + "aws_web_identity_token": "oidc/google/108963886734710037768", + }, + { + "aws_access_key_id": "AKIASTATICKEYFORTEST", + "aws_secret_access_key": "static-secret-key", + "aws_session_token": "static-session-token", + }, + ], + ids=["web_identity", "static_keys"], +) async def test_acompletion_forwards_aws_credentials_through_responses_bridge( - respx_mock: respx.MockRouter, monkeypatch + respx_mock: respx.MockRouter, monkeypatch, aws_credential_kwargs: dict ): from botocore.credentials import Credentials @@ -2126,17 +2142,15 @@ async def test_acompletion_forwards_aws_credentials_through_responses_bridge( messages=[{"role": "user", "content": "hi"}], api_base="https://bedrock-mantle.us-east-2.api.aws/v1", aws_region_name="us-east-2", - aws_session_name="litellm-gcp", - aws_role_name="arn:aws:iam::123456789012:role/litellm-bedrock-role", - aws_web_identity_token="oidc/google/108963886734710037768", num_retries=0, + **aws_credential_kwargs, ) assert response.choices[0].message.content == "ok" credential_kwargs = get_credentials_mock.call_args.kwargs - assert credential_kwargs["aws_role_name"] == "arn:aws:iam::123456789012:role/litellm-bedrock-role" - assert credential_kwargs["aws_web_identity_token"] == "oidc/google/108963886734710037768" - assert credential_kwargs["aws_session_name"] == "litellm-gcp" + assert credential_kwargs["aws_region_name"] == "us-east-2" + for key, value in aws_credential_kwargs.items(): + assert credential_kwargs[key] == value authorization = respx_mock.calls.last.request.headers["Authorization"] assert authorization.startswith("AWS4-HMAC-SHA256") assert "fake-key" in authorization From 2deeeb73a7678f8f802fc320a8966250bdfd8e92 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 17:12:38 -0700 Subject: [PATCH 023/123] docs(github): add QA runbook section to the PR template --- .github/pull_request_template.md | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index bd9fc2285d1..fcf894a4dca 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -41,3 +41,7 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac ✅ Test ## Changes + +## QA runbook + + From 1d589832c7b5468f629cf9d028e479a9468f684e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 17:18:22 -0700 Subject: [PATCH 024/123] docs(github): scope the QA runbook to tests/e2e edits and add example checklists --- .github/pull_request_template.md | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index fcf894a4dca..0c72186d959 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -44,4 +44,18 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac ## QA runbook - + From a6c8fb2d1c8b9c8c0c4fc3215e00f3f8d7eee7c8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 17:21:36 -0700 Subject: [PATCH 025/123] docs(github): shape QA runbook examples as node id plus behavior bullets --- .github/pull_request_template.md | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 0c72186d959..df44a4063a1 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -46,16 +46,16 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac From f947ef14a2a7db3e2d81914977358636463bea71 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Sat, 11 Jul 2026 17:50:13 -0700 Subject: [PATCH 026/123] fix(xecguard): use StandardLoggingGuardrailInformation in logging hook (#32911) XecGuard's async_logging_hook wrote a bare dict to standard_logging_object["guardrail_information"] while the typed contract is Optional[List[StandardLoggingGuardrailInformation]]. Readers that iterated the field walked dict keys, raised on info.get, or silently dropped the entry from guardrail usage tracking and spend-log writes Construct the typed entry and append it to the existing list or create a new one, matching the shared helper pattern. Record the configured guardrail name instead of a hardcoded "xecguard" and pass the GuardrailEventHooks enum for guardrail_mode --- .../guardrail_hooks/xecguard/xecguard.py | 31 ++++++++++------- .../guardrail_hooks/test_xecguard.py | 33 +++++++++++++++++-- 2 files changed, 50 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index 1b663c16d5b..f4a6f0aeb3b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -49,7 +49,11 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus +from litellm.types.utils import ( + GenericGuardrailAPIInputs, + GuardrailStatus, + StandardLoggingGuardrailInformation, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import ( @@ -246,16 +250,21 @@ class XecGuardGuardrail(CustomGuardrail): "guardrail_intervened" if scan_result.get("decision") == "UNSAFE" else "success" ) end_time = datetime.now() - kwargs["standard_logging_object"]["guardrail_information"] = { - "duration": (end_time - start_time).total_seconds(), - "end_time": end_time.timestamp(), - "guardrail_mode": "logging_only", - "guardrail_name": "xecguard", - "guardrail_response": scan_result, - "guardrail_status": guardrail_status, - "masked_entity_count": None, - "start_time": start_time.timestamp(), - } + slg = StandardLoggingGuardrailInformation( + guardrail_name=self.guardrail_name or "xecguard", + guardrail_mode=GuardrailEventHooks.logging_only, + guardrail_response=scan_result, + guardrail_status=guardrail_status, + start_time=start_time.timestamp(), + end_time=end_time.timestamp(), + duration=(end_time - start_time).total_seconds(), + masked_entity_count=None, + ) + existing = kwargs["standard_logging_object"].get("guardrail_information") + if isinstance(existing, list): + existing.append(slg) + else: + kwargs["standard_logging_object"]["guardrail_information"] = [slg] except Exception as exc: verbose_proxy_logger.debug( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py index 60c595c64e5..f64967abbb0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py @@ -1644,12 +1644,37 @@ class TestXecGuardLoggingHook: ) assert out_kwargs is kwargs assert out_result is result - info = kwargs["standard_logging_object"]["guardrail_information"] + info_list = kwargs["standard_logging_object"]["guardrail_information"] + assert isinstance(info_list, list), "guardrail_information must be a list" + assert len(info_list) == 1 + info = info_list[0] assert info["guardrail_mode"] == "logging_only" - assert info["guardrail_name"] == "xecguard" + assert info["guardrail_name"] == "test-xecguard" assert info["guardrail_status"] == "success" assert info["guardrail_response"]["trace_id"] == "lg-1" + @pytest.mark.asyncio + async def test_async_logging_hook_appends_to_existing_guardrail_info( + self, xecguard_guardrail, mock_request_data + ): + resp = _make_response({"decision": "SAFE", "trace_id": "lg-4"}) + prior_entry = {"guardrail_name": "other-guardrail"} + with patch.object(xecguard_guardrail.async_handler, "post", return_value=resp): + kwargs = { + **mock_request_data, + "standard_logging_object": {"guardrail_information": [prior_entry]}, + } + await xecguard_guardrail.async_logging_hook( + kwargs=kwargs, + result=_build_model_response("some answer"), + call_type="acompletion", + ) + info_list = kwargs["standard_logging_object"]["guardrail_information"] + assert len(info_list) == 2 + assert info_list[0] is prior_entry + assert info_list[1]["guardrail_name"] == "test-xecguard" + assert info_list[1]["guardrail_response"]["trace_id"] == "lg-4" + @pytest.mark.asyncio async def test_async_logging_hook_without_response_records_info( self, xecguard_guardrail, mock_request_data @@ -1680,7 +1705,9 @@ class TestXecGuardLoggingHook: result=_build_model_response("x"), call_type="acompletion", ) - info = kwargs["standard_logging_object"]["guardrail_information"] + info_list = kwargs["standard_logging_object"]["guardrail_information"] + assert isinstance(info_list, list), "guardrail_information must be a list" + info = info_list[0] assert info["guardrail_status"] == "guardrail_intervened" @pytest.mark.asyncio From 8c5473f198b71213d9b703c68e6ab2115036acd8 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 11 Jul 2026 18:01:33 -0700 Subject: [PATCH 027/123] feat(ui): adopt openapi-react-query ($api) and convert useCustomers (#32949) * feat(ui): adopt openapi-react-query and convert useCustomers to $api Add openapi-react-query and expose $api = createQueryClient(fetchClient) alongside fetchClient. Rewrite useCustomers as $api.useQuery("get", "/customer/list", {}, { enabled, select }), which derives the query key from method + path (dropping the hand-written createQueryKeys entry and the manual key) and forwards the request signal for cancellation. The response type still flows from schema.d.ts as CustomerResponse[]. Tests assert the path, the admin/token enabled gate, and the empty-body select fallback. * test(ui): read the last render's options in useCustomers helper The lastCallOptions helper was named for the last call but read mock.calls[0]. Harmless while each test renders once, but it would silently assert against first-render options if a test ever re-renders. Read the final call instead. --- ui/litellm-dashboard/package-lock.json | 14 +++ ui/litellm-dashboard/package.json | 1 + .../hooks/customers/useCustomers.test.ts | 108 ++++++------------ .../hooks/customers/useCustomers.ts | 20 ++-- ui/litellm-dashboard/src/lib/http/api.ts | 9 ++ 5 files changed, 69 insertions(+), 83 deletions(-) diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 648503abb27..b8f6265441b 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -29,6 +29,7 @@ "next": "16.2.6", "openai": "4.104.0", "openapi-fetch": "^0.17.0", + "openapi-react-query": "^0.5.4", "papaparse": "5.5.3", "react": "18.3.1", "react-copy-to-clipboard": "5.1.1", @@ -10544,6 +10545,19 @@ "openapi-typescript-helpers": "^0.1.0" } }, + "node_modules/openapi-react-query": { + "version": "0.5.4", + "resolved": "https://registry.npmjs.org/openapi-react-query/-/openapi-react-query-0.5.4.tgz", + "integrity": "sha512-V9lRiozjHot19/BYSgXYoyznDxDJQhEBSdi26+SJ0UqjMANLQhkni4XG+Z7e3Ag7X46ZLMrL9VxYkghU3QvbWg==", + "license": "MIT", + "dependencies": { + "openapi-typescript-helpers": "^0.1.0" + }, + "peerDependencies": { + "@tanstack/react-query": "^5.80.0", + "openapi-fetch": "^0.17.0" + } + }, "node_modules/openapi-typescript": { "version": "7.13.0", "resolved": "https://registry.npmjs.org/openapi-typescript/-/openapi-typescript-7.13.0.tgz", diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 40c495cebf1..a14349c7239 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -44,6 +44,7 @@ "next": "16.2.6", "openai": "4.104.0", "openapi-fetch": "^0.17.0", + "openapi-react-query": "^0.5.4", "papaparse": "5.5.3", "react": "18.3.1", "react-copy-to-clipboard": "5.1.1", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts index 1e614b709e2..b09ab4498f7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.test.ts @@ -1,12 +1,10 @@ -import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { renderHook, waitFor } from "@testing-library/react"; -import React, { ReactNode } from "react"; +import { renderHook } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { useCustomers, type EndUser } from "./useCustomers"; -const mockGet = vi.fn(); +const useQueryMock = vi.fn(); vi.mock("@/lib/http/api", () => ({ - fetchClient: { GET: (...args: unknown[]) => mockGet(...args) }, + $api: { useQuery: (...args: unknown[]) => useQueryMock(...args) }, })); const mockUseAuthorized = vi.fn(); @@ -14,91 +12,55 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => mockUseAuthorized(), })); -const mockCustomers: EndUser[] = [ - { user_id: "customer-1", alias: "Test Customer 1", spend: 150.5, blocked: false }, - { user_id: "customer-2", alias: null, spend: 0, blocked: true }, -]; +const authorized = { accessToken: "test-access-token", userRole: "Admin" }; -const authorized = { - accessToken: "test-access-token", - userRole: "Admin", - userId: "test-user-id", - token: "test-token", - userEmail: "test@example.com", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, +type QueryOptions = { enabled: boolean; select: (data: EndUser[] | undefined) => EndUser[] }; + +const lastCallOptions = (): QueryOptions => { + const calls = useQueryMock.mock.calls; + return calls[calls.length - 1][3] as QueryOptions; }; describe("useCustomers", () => { - let queryClient: QueryClient; - beforeEach(() => { - queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); vi.clearAllMocks(); + useQueryMock.mockReturnValue({ data: [] }); mockUseAuthorized.mockReturnValue(authorized); }); - const wrapper = ({ children }: { children: ReactNode }) => - React.createElement(QueryClientProvider, { client: queryClient }, children); - - it("fetches /customer/list and returns the typed list on success", async () => { - mockGet.mockResolvedValue({ data: mockCustomers }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - expect(result.current.isLoading).toBe(true); - - await waitFor(() => { - expect(result.current.isSuccess).toBe(true); - }); - - expect(result.current.data).toEqual(mockCustomers); - expect(mockGet).toHaveBeenCalledWith("/customer/list"); - expect(mockGet).toHaveBeenCalledTimes(1); + it("queries GET /customer/list with a derived key (no hand-written queryKey)", () => { + renderHook(() => useCustomers()); + expect(useQueryMock).toHaveBeenCalledWith("get", "/customer/list", {}, expect.any(Object)); }); - it("surfaces an error when the request rejects", async () => { - const testError = new Error("Failed to fetch customers"); - mockGet.mockRejectedValue(testError); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - await waitFor(() => { - expect(result.current.isError).toBe(true); - }); - - expect(result.current.error).toEqual(testError); - expect(result.current.data).toBeUndefined(); + it("enables the query only for an admin holding an access token", () => { + renderHook(() => useCustomers()); + expect(lastCallOptions().enabled).toBe(true); }); - it("falls back to an empty list when the response has no body", async () => { - mockGet.mockResolvedValue({ data: undefined }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - await waitFor(() => { - expect(result.current.isSuccess).toBe(true); - }); - - expect(result.current.data).toEqual([]); + it("disables the query when the access token is missing", () => { + mockUseAuthorized.mockReturnValue({ ...authorized, accessToken: null }); + renderHook(() => useCustomers()); + expect(lastCallOptions().enabled).toBe(false); }); - it("does not fetch when the access token is missing", () => { - mockUseAuthorized.mockReturnValue({ ...authorized, accessToken: null, token: null }); - - const { result } = renderHook(() => useCustomers(), { wrapper }); - - expect(result.current.isFetched).toBe(false); - expect(mockGet).not.toHaveBeenCalled(); - }); - - it("does not fetch when the user is not an admin", () => { + it("disables the query for a non-admin role", () => { mockUseAuthorized.mockReturnValue({ ...authorized, userRole: "member" }); + renderHook(() => useCustomers()); + expect(lastCallOptions().enabled).toBe(false); + }); - const { result } = renderHook(() => useCustomers(), { wrapper }); + it("selects an empty list when the response body is missing", () => { + renderHook(() => useCustomers()); + expect(lastCallOptions().select(undefined)).toEqual([]); + }); - expect(result.current.isFetched).toBe(false); - expect(mockGet).not.toHaveBeenCalled(); + it("selects the customer list through unchanged", () => { + const customers: EndUser[] = [ + { user_id: "customer-1", alias: "Test Customer 1", spend: 150.5, blocked: false }, + { user_id: "customer-2", alias: null, spend: 0, blocked: true }, + ]; + renderHook(() => useCustomers()); + expect(lastCallOptions().select(customers)).toEqual(customers); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts index 25e2e3f5e90..ebea4618b7f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/customers/useCustomers.ts @@ -1,19 +1,19 @@ -import { useQuery } from "@tanstack/react-query"; -import { createQueryKeys } from "../common/queryKeysFactory"; -import { fetchClient } from "@/lib/http/api"; +import { $api } from "@/lib/http/api"; import { all_admin_roles } from "@/utils/roles"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import type { components } from "@/lib/http/schema"; export type EndUser = components["schemas"]["CustomerResponse"]; -const customersKeys = createQueryKeys("customers"); - export const useCustomers = () => { const { accessToken, userRole } = useAuthorized(); - return useQuery({ - queryKey: customersKeys.list({}), - queryFn: async () => (await fetchClient.GET("/customer/list")).data ?? [], - enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!), - }); + return $api.useQuery( + "get", + "/customer/list", + {}, + { + enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!), + select: (data) => data ?? [], + }, + ); }; diff --git a/ui/litellm-dashboard/src/lib/http/api.ts b/ui/litellm-dashboard/src/lib/http/api.ts index 2e8e1d36c5b..aa6c2d6c0fd 100644 --- a/ui/litellm-dashboard/src/lib/http/api.ts +++ b/ui/litellm-dashboard/src/lib/http/api.ts @@ -1,4 +1,5 @@ import createFetchClient, { type Middleware } from "openapi-fetch"; +import createQueryClient from "openapi-react-query"; import type { paths } from "./schema"; import { ApiError, deriveErrorMessage } from "./client"; import { getAuthHeaderName, getAuthToken, getRequestBaseUrl, reportError } from "./runtime"; @@ -46,3 +47,11 @@ const middleware: Middleware = { */ export const fetchClient = createFetchClient({ baseUrl: globalThis.location?.origin ?? "" }); fetchClient.use(middleware); + +/** + * TanStack Query bound to the typed client. Callers write + * `$api.useQuery("get", "/path", init, options)`; the query key is derived from + * method + path + init (no hand-maintained key), the request signal is + * forwarded for cancellation, and the response type comes from schema.d.ts. + */ +export const $api = createQueryClient(fetchClient); From b6dbda48f9e88ed0c3e1918488e0545946e9d627 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 11 Jul 2026 18:01:41 -0700 Subject: [PATCH 028/123] refactor(ui): colocate the mcp-servers view, keeping the shared mcp_tools surface (#32968) * refactor(ui): colocate the usage view, keeping the shared usage components Split for the usage (UsagePage) segment. Most of the folder is the usage page's own view, but four pieces are reused elsewhere and stay in @/components/UsagePage: TopKeyView (old-usage), KeyModelUsageView and value_formatters (activity_metrics), and the shared types (activity_metrics, chartUtils). The other 21 files move into usage/_components, preserving the folder structure. The external consumers import only the retained files, so they are untouched. The moved files' imports of the retained files become @/components/UsagePage paths, other escaping relative imports are absolutized, and lint suppressions are re-keyed for moved files only. No behavior change. * refactor(ui): colocate the mcp-servers view, keeping the shared mcp_tools surface --- ui/litellm-dashboard/eslint-suppressions.json | 40 +++++++++---------- .../_components}/DcrBridgeToggle.tsx | 2 +- .../_components}/EnvVarsSection.tsx | 0 .../_components}/MCPLogoSelector.test.tsx | 0 .../_components}/MCPLogoSelector.tsx | 0 .../_components}/MCPNetworkSettings.tsx | 4 +- .../MCPPermissionManagement.test.tsx | 0 .../_components}/MCPPermissionManagement.tsx | 2 +- .../_components}/MCPServerCard.test.tsx | 2 +- .../_components}/MCPServerCard.tsx | 2 +- .../MCPStandardsSettings.test.tsx | 2 +- .../_components}/MCPStandardsSettings.tsx | 2 +- .../_components}/MCPSubmissionsTab.tsx | 2 +- .../_components}/MCPToolsetsTab.tsx | 12 ++++-- .../_components}/OAuthFormFields.test.tsx | 0 .../_components}/OAuthFormFields.tsx | 2 +- .../_components}/OpenAPIFormSection.tsx | 2 +- .../_components}/OpenAPIQuickPicker.tsx | 2 +- .../PassthroughAuthorizeSection.test.tsx | 0 .../PassthroughAuthorizeSection.tsx | 2 +- .../_components}/StdioConfiguration.tsx | 0 .../TokenEndpointAuthMethodField.tsx | 0 .../_components}/TokenExchangeFormFields.tsx | 0 .../_components}/ToolTestPanel.test.tsx | 4 +- .../_components}/ToolTestPanel.tsx | 4 +- .../_components}/TruePassthroughWarning.tsx | 2 +- .../_components}/UserEnvVarsModal.tsx | 6 +-- .../_components}/create_mcp_server.test.tsx | 4 +- .../_components}/create_mcp_server.tsx | 6 +-- .../mcp-servers/_components}/index.tsx | 0 .../mcp-servers/_components}/mcp_connect.tsx | 4 +- .../mcp_connection_status.test.tsx | 0 .../_components}/mcp_connection_status.tsx | 0 .../_components}/mcp_discovery.tsx | 4 +- .../_components}/mcp_server_cost_config.tsx | 2 +- .../_components}/mcp_server_cost_display.tsx | 2 +- .../_components}/mcp_server_edit.test.tsx | 8 ++-- .../_components}/mcp_server_edit.tsx | 11 +++-- .../_components}/mcp_server_view.tsx | 2 +- .../_components}/mcp_servers.test.tsx | 6 +-- .../mcp-servers/_components}/mcp_servers.tsx | 24 ++++++----- .../mcp_tool_configuration.test.tsx | 0 .../_components}/mcp_tool_configuration.tsx | 2 +- .../_components}/mcp_tools.test.tsx | 4 +- .../mcp-servers/_components}/mcp_tools.tsx | 4 +- .../mcp-servers/_components}/testUtils.ts | 0 .../mcp-servers/_components}/utils.test.tsx | 0 .../mcp-servers/_components}/utils.tsx | 2 +- .../src/app/(dashboard)/mcp-servers/page.tsx | 2 +- .../tests/CreateKeyPage.expiredToken.test.tsx | 2 +- 50 files changed, 100 insertions(+), 83 deletions(-) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/DcrBridgeToggle.tsx (95%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/EnvVarsSection.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPLogoSelector.test.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPLogoSelector.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPNetworkSettings.tsx (97%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPPermissionManagement.test.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPPermissionManagement.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPServerCard.test.tsx (96%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPServerCard.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPStandardsSettings.test.tsx (97%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPStandardsSettings.tsx (96%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPSubmissionsTab.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/MCPToolsetsTab.tsx (98%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/OAuthFormFields.test.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/OAuthFormFields.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/OpenAPIFormSection.tsx (98%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/OpenAPIQuickPicker.tsx (98%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/PassthroughAuthorizeSection.test.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/PassthroughAuthorizeSection.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/StdioConfiguration.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/TokenEndpointAuthMethodField.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/TokenExchangeFormFields.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/ToolTestPanel.test.tsx (97%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/ToolTestPanel.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/TruePassthroughWarning.tsx (94%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/UserEnvVarsModal.tsx (95%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/create_mcp_server.test.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/create_mcp_server.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/index.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_connect.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_connection_status.test.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_connection_status.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_discovery.tsx (98%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_server_cost_config.tsx (98%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_server_cost_display.tsx (97%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_server_edit.test.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_server_edit.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_server_view.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_servers.test.tsx (98%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_servers.tsx (97%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_tool_configuration.test.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_tool_configuration.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_tools.test.tsx (98%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/mcp_tools.tsx (99%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/testUtils.ts (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/utils.test.tsx (100%) rename ui/litellm-dashboard/src/{components/mcp_tools => app/(dashboard)/mcp-servers/_components}/utils.tsx (97%) diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 48b400536bd..eb8bdee0ea2 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1878,17 +1878,17 @@ "count": 1 } }, - "src/components/mcp_tools/MCPLogoSelector.test.tsx": { + "src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx": { "unused-imports/no-unused-imports": { "count": 1 } }, - "src/components/mcp_tools/MCPNetworkSettings.tsx": { + "src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx": { "react-hooks/immutability": { "count": 2 } }, - "src/components/mcp_tools/MCPSubmissionsTab.tsx": { + "src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx": { "react-hooks/set-state-in-effect": { "count": 1 } @@ -1898,7 +1898,7 @@ "count": 5 } }, - "src/components/mcp_tools/MCPToolsetsTab.tsx": { + "src/app/(dashboard)/mcp-servers/_components/MCPToolsetsTab.tsx": { "no-nested-ternary": { "count": 1 }, @@ -1920,7 +1920,7 @@ "count": 1 } }, - "src/components/mcp_tools/OAuthFormFields.tsx": { + "src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx": { "no-nested-ternary": { "count": 1 }, @@ -1928,12 +1928,12 @@ "count": 1 } }, - "src/components/mcp_tools/OpenAPIQuickPicker.tsx": { + "src/app/(dashboard)/mcp-servers/_components/OpenAPIQuickPicker.tsx": { "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/components/mcp_tools/ToolTestPanel.tsx": { + "src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx": { "no-nested-ternary": { "count": 3 }, @@ -1944,12 +1944,12 @@ "count": 1 } }, - "src/components/mcp_tools/UserEnvVarsModal.tsx": { + "src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx": { "no-nested-ternary": { "count": 2 } }, - "src/components/mcp_tools/create_mcp_server.tsx": { + "src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx": { "no-nested-ternary": { "count": 1 }, @@ -1960,7 +1960,7 @@ "count": 4 } }, - "src/components/mcp_tools/mcp_connect.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx": { "no-restricted-imports": { "count": 1 }, @@ -1968,7 +1968,7 @@ "count": 4 } }, - "src/components/mcp_tools/mcp_connection_status.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.tsx": { "no-nested-ternary": { "count": 3 }, @@ -1976,22 +1976,22 @@ "count": 1 } }, - "src/components/mcp_tools/mcp_discovery.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_discovery.tsx": { "react-hooks/set-state-in-effect": { "count": 2 } }, - "src/components/mcp_tools/mcp_server_cost_config.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_config.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/components/mcp_tools/mcp_server_cost_display.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_display.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/components/mcp_tools/mcp_server_edit.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx": { "no-nested-ternary": { "count": 1 }, @@ -2005,12 +2005,12 @@ "count": 5 } }, - "src/components/mcp_tools/mcp_server_view.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/components/mcp_tools/mcp_servers.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx": { "no-nested-ternary": { "count": 1 }, @@ -2021,12 +2021,12 @@ "count": 2 } }, - "src/components/mcp_tools/mcp_tool_configuration.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/components/mcp_tools/mcp_tools.tsx": { + "src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx": { "no-nested-ternary": { "count": 1 }, @@ -2558,4 +2558,4 @@ "count": 1 } } -} \ No newline at end of file +} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/DcrBridgeToggle.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx similarity index 95% rename from ui/litellm-dashboard/src/components/mcp_tools/DcrBridgeToggle.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx index f1b642293c0..49c182aa6be 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/DcrBridgeToggle.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx @@ -1,7 +1,7 @@ import React from "react"; import { Form, Switch, Tooltip } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; -import { isClientForwardedTokenMode } from "./types"; +import { isClientForwardedTokenMode } from "@/components/mcp_tools/types"; /** * DCR-bridge toggle for the client-forwarded token modes (true_passthrough / diff --git a/ui/litellm-dashboard/src/components/mcp_tools/EnvVarsSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/EnvVarsSection.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/EnvVarsSection.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/EnvVarsSection.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPLogoSelector.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPLogoSelector.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPLogoSelector.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPLogoSelector.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPNetworkSettings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPNetworkSettings.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx index 00323465731..7ab240389f3 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPNetworkSettings.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx @@ -1,13 +1,13 @@ import React, { useState, useEffect } from "react"; import { Select, Button, Card, Typography, Spin, Tag } from "antd"; import { SaveOutlined, PlusOutlined } from "@ant-design/icons"; -import { DeprecationBanner } from "../DeprecationBanner"; +import { DeprecationBanner } from "@/components/DeprecationBanner"; import { getGeneralSettingsCall, updateConfigFieldSetting, deleteConfigFieldSetting, fetchMCPClientIp, -} from "../networking"; +} from "@/components/networking"; const { Text } = Typography; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.test.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx index 27cbdf2ea34..aae13d4b467 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx @@ -1,7 +1,7 @@ import React, { useEffect } from "react"; import { Alert, Form, Select, Tooltip, Collapse, Input, Space, Button, Switch } from "antd"; import { InfoCircleOutlined, MinusCircleOutlined, PlusOutlined } from "@ant-design/icons"; -import { MCPServer, AUTH_TYPE } from "./types"; +import { MCPServer, AUTH_TYPE } from "@/components/mcp_tools/types"; const { Panel } = Collapse; interface MCPPermissionManagementProps { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx similarity index 96% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPServerCard.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index 8463a23cc16..a0998b587fb 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -2,7 +2,7 @@ import React from "react"; import { render, screen } from "@testing-library/react"; import { describe, it, expect, vi } from "vitest"; import MCPServerCard from "./MCPServerCard"; -import type { MCPServer } from "./types"; +import type { MCPServer } from "@/components/mcp_tools/types"; const baseServer: MCPServer = { server_id: "srv-1", diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPServerCard.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index c16ad87980e..4282cdba278 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -8,7 +8,7 @@ import { MoreOutlined, ThunderboltOutlined, } from "@ant-design/icons"; -import { AUTH_TYPE, type MCPServer } from "./types"; +import { AUTH_TYPE, type MCPServer } from "@/components/mcp_tools/types"; import { getMaskedAndFullUrl } from "./utils"; const { Text } = Typography; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPStandardsSettings.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPStandardsSettings.test.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPStandardsSettings.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPStandardsSettings.test.tsx index fd8d9f92c48..08f65a72537 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPStandardsSettings.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPStandardsSettings.test.tsx @@ -1,6 +1,6 @@ import { describe, it, expect } from "vitest"; import { FIELD_GROUPS, MCP_REQUIRED_FIELD_DEFS, SETTINGS_KEY } from "./MCPStandardsSettings"; -import { MCPServer } from "./types"; +import { MCPServer } from "@/components/mcp_tools/types"; const makeServer = (overrides: Partial = {}): MCPServer => ({ server_id: "s1", diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPStandardsSettings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPStandardsSettings.tsx similarity index 96% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPStandardsSettings.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPStandardsSettings.tsx index fb38e392631..13a9b7c171e 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPStandardsSettings.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPStandardsSettings.tsx @@ -1,6 +1,6 @@ "use client"; -import { MCPServer } from "./types"; +import { MCPServer } from "@/components/mcp_tools/types"; export interface RequiredFieldDef { key: string; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPSubmissionsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPSubmissionsTab.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx index 50c88a63582..de030420bc3 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPSubmissionsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx @@ -18,7 +18,7 @@ import { getGeneralSettingsCall, updateConfigFieldSetting, } from "@/components/networking"; -import { MCPServer, MCPSubmissionsSummary } from "./types"; +import { MCPServer, MCPSubmissionsSummary } from "@/components/mcp_tools/types"; import { FIELD_GROUPS, MCP_REQUIRED_FIELD_DEFS, SETTINGS_KEY } from "./MCPStandardsSettings"; import NotificationsManager from "@/components/molecules/notifications_manager"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPToolsetsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolsetsTab.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/mcp_tools/MCPToolsetsTab.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolsetsTab.tsx index 5e7e99ee8ec..fa44694887b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPToolsetsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPToolsetsTab.tsx @@ -7,9 +7,15 @@ import { useMCPToolsets } from "@/app/(dashboard)/hooks/mcpServers/useMCPToolset import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers"; import { useQueryClient } from "@tanstack/react-query"; import { DateCell, IdCell } from "@/components/shared/table_cells"; -import { DataTable } from "../view_logs/table"; -import { createMCPToolset, updateMCPToolset, deleteMCPToolset, listMCPTools, getProxyBaseUrl } from "../networking"; -import { MCPToolset, MCPToolsetTool } from "./types"; +import { DataTable } from "@/components/view_logs/table"; +import { + createMCPToolset, + updateMCPToolset, + deleteMCPToolset, + listMCPTools, + getProxyBaseUrl, +} from "@/components/networking"; +import { MCPToolset, MCPToolsetTool } from "@/components/mcp_tools/types"; const { Text: AntdText } = Typography; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/OAuthFormFields.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/OAuthFormFields.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.test.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/OAuthFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/OAuthFormFields.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx index 93afefe4358..f359bdee065 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/OAuthFormFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx @@ -2,7 +2,7 @@ import React from "react"; import { Form, Input, InputNumber, Select, Tooltip } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; import { Button, TextInput } from "@tremor/react"; -import { OAUTH_FLOW } from "./types"; +import { OAUTH_FLOW } from "@/components/mcp_tools/types"; import TokenEndpointAuthMethodField from "./TokenEndpointAuthMethodField"; interface OAuthFlowStatus { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/OpenAPIFormSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/mcp_tools/OpenAPIFormSection.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx index d606b8ba1fc..073780b359f 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/OpenAPIFormSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx @@ -2,7 +2,7 @@ import React, { useState } from "react"; import { Form, Input, Tooltip } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; import { FormInstance } from "antd/es/form"; -import { AUTH_TYPE, OAUTH_FLOW } from "./types"; +import { AUTH_TYPE, OAUTH_FLOW } from "@/components/mcp_tools/types"; import OpenAPIQuickPicker, { OpenAPIRegistryEntry, OpenAPIKeyTool } from "./OpenAPIQuickPicker"; interface OpenAPIFormSectionProps { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/OpenAPIQuickPicker.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIQuickPicker.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/mcp_tools/OpenAPIQuickPicker.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIQuickPicker.tsx index a8208a09286..0aec81fdf4b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/OpenAPIQuickPicker.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIQuickPicker.tsx @@ -1,6 +1,6 @@ import React, { useEffect, useState } from "react"; import { Spin } from "antd"; -import { fetchOpenAPIRegistry } from "../networking"; +import { fetchOpenAPIRegistry } from "@/components/networking"; export interface OpenAPIKeyTool { name: string; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.test.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.tsx index 0ed4ee555d1..dc10f0f1392 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/PassthroughAuthorizeSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.tsx @@ -1,7 +1,7 @@ import React from "react"; import { Button, Checkbox, Form, Input } from "antd"; import DcrBridgeToggle from "./DcrBridgeToggle"; -import { credentialAuthClass, isClientForwardedTokenMode } from "./types"; +import { credentialAuthClass, isClientForwardedTokenMode } from "@/components/mcp_tools/types"; interface PassthroughOAuthFlow { startOAuthFlow: () => void | Promise; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/StdioConfiguration.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/StdioConfiguration.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/StdioConfiguration.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/StdioConfiguration.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/TokenEndpointAuthMethodField.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenEndpointAuthMethodField.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/TokenEndpointAuthMethodField.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenEndpointAuthMethodField.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/TokenExchangeFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenExchangeFormFields.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/TokenExchangeFormFields.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenExchangeFormFields.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.test.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.test.tsx index 4bd351216b8..0613e84feed 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.test.tsx @@ -3,9 +3,9 @@ import { render, screen } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import { ToolTestPanel } from "./ToolTestPanel"; -import { InputSchema, MCPTool } from "./types"; +import { InputSchema, MCPTool } from "@/components/mcp_tools/types"; -vi.mock("../molecules/notifications_manager", () => ({ +vi.mock("@/components/molecules/notifications_manager", () => ({ default: { success: vi.fn(), fromBackend: vi.fn(), diff --git a/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx index c15150eee58..8f042445c01 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx @@ -1,10 +1,10 @@ import React from "react"; import { Button, TextInput } from "@tremor/react"; -import { MCPTool, InputSchema, InputSchemaProperty } from "./types"; +import { MCPTool, InputSchema, InputSchemaProperty } from "@/components/mcp_tools/types"; import { resolveLogoSrc } from "@/lib/assetPaths"; import { Form, Select, Tooltip } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; -import NotificationsManager from "../molecules/notifications_manager"; +import NotificationsManager from "@/components/molecules/notifications_manager"; const isPlainObject = (value: unknown): value is Record => typeof value === "object" && value !== null && !Array.isArray(value); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/TruePassthroughWarning.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TruePassthroughWarning.tsx similarity index 94% rename from ui/litellm-dashboard/src/components/mcp_tools/TruePassthroughWarning.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TruePassthroughWarning.tsx index b52d3f4c672..9c57cbd7d14 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/TruePassthroughWarning.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TruePassthroughWarning.tsx @@ -1,6 +1,6 @@ import React from "react"; import { Alert } from "antd"; -import { AUTH_TYPE } from "./types"; +import { AUTH_TYPE } from "@/components/mcp_tools/types"; /** * Warning shown in the create/edit MCP server forms when auth_type diff --git a/ui/litellm-dashboard/src/components/mcp_tools/UserEnvVarsModal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx similarity index 95% rename from ui/litellm-dashboard/src/components/mcp_tools/UserEnvVarsModal.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx index 08a285cd56b..f76aef365f7 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/UserEnvVarsModal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx @@ -1,9 +1,9 @@ import React from "react"; import { Modal, Form, Input, Button, Alert, Spin, Tag, Typography } from "antd"; import { useMutation, useQuery } from "@tanstack/react-query"; -import { MCPServer, MCPUserEnvVarsStatus } from "./types"; -import { getMCPUserEnvVars, storeMCPUserEnvVars } from "../networking"; -import NotificationsManager from "../molecules/notifications_manager"; +import { MCPServer, MCPUserEnvVarsStatus } from "@/components/mcp_tools/types"; +import { getMCPUserEnvVars, storeMCPUserEnvVars } from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; const { Text, Title } = Typography; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.test.tsx index 32fe439a316..6ce7f5c75ed 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.test.tsx @@ -1,12 +1,12 @@ import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import * as networking from "../networking"; +import * as networking from "@/components/networking"; import { setToken } from "@/utils/mcpTokenStore"; import CreateMCPServer from "./create_mcp_server"; import { selectAntOption } from "./testUtils"; -vi.mock("../networking", () => ({ +vi.mock("@/components/networking", () => ({ createMCPServer: vi.fn(), fetchOpenAPIRegistry: vi.fn().mockResolvedValue({ apis: [] }), registerMCPServer: vi.fn(), diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx index b9a6f5d229d..70838c592c8 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx @@ -2,7 +2,7 @@ import React, { useState } from "react"; import { Modal, Tooltip, Form, Select, Input, InputNumber, Switch, Collapse } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; import { Button, TextInput } from "@tremor/react"; -import { createMCPServer, registerMCPServer, storeMCPOAuthUserCredential } from "../networking"; +import { createMCPServer, registerMCPServer, storeMCPOAuthUserCredential } from "@/components/networking"; import { setToken } from "@/utils/mcpTokenStore"; import { AUTH_TYPE, @@ -20,7 +20,7 @@ import { isHeldOAuthTokenStale, preservedDeclaredAppCredentials, withoutMintedTokenCredentials, -} from "./types"; +} from "@/components/mcp_tools/types"; import OAuthFormFields from "./OAuthFormFields"; import TruePassthroughWarning from "./TruePassthroughWarning"; import PassthroughAuthorizeSection from "./PassthroughAuthorizeSection"; @@ -35,7 +35,7 @@ import MCPLogoSelector from "./MCPLogoSelector"; import EnvVarsSection from "./EnvVarsSection"; import { isAdminRole } from "@/utils/roles"; import { validateMCPServerUrl, validateMCPServerName, normalizeEnvVars, TOOL_DISPLAY_NAME_PATTERN } from "./utils"; -import NotificationsManager from "../molecules/notifications_manager"; +import NotificationsManager from "@/components/molecules/notifications_manager"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; import { useTestMCPConnection } from "@/hooks/useTestMCPConnection"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/index.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/index.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/index.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/index.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx index 42c557f142f..7bdfd9c6b8f 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connect.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx @@ -4,8 +4,8 @@ import React, { useState } from "react"; import { Card, Typography, Space, Alert, Button, Switch, Form, Collapse } from "antd"; import { TabPanel, TabPanels, TabGroup, TabList, Tab, Title as TremorTitle, Text as TremorText } from "@tremor/react"; import { CopyIcon, Code, Terminal, Globe, CheckIcon, ExternalLinkIcon, KeyIcon, ServerIcon, Zap } from "lucide-react"; -import { getProxyBaseUrl } from "../networking"; -import { copyToClipboard as utilCopyToClipboard } from "../../utils/dataUtils"; +import { getProxyBaseUrl } from "@/components/networking"; +import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; const { Title, Text } = Typography; const { Panel } = Collapse; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connection_status.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_connection_status.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.test.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_connection_status.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_connection_status.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_discovery.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_discovery.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_discovery.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_discovery.tsx index 189b52dd6b2..6fcff011ba6 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_discovery.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_discovery.tsx @@ -1,7 +1,7 @@ import React, { useState, useMemo, useEffect } from "react"; import { Modal, Input, Typography } from "antd"; -import { fetchDiscoverableMCPServers } from "../networking"; -import { DiscoverableMCPServer, DiscoverMCPServersResponse } from "./types"; +import { fetchDiscoverableMCPServers } from "@/components/networking"; +import { DiscoverableMCPServer, DiscoverMCPServersResponse } from "@/components/mcp_tools/types"; import { mcpLogoImg } from "./create_mcp_server"; import { resolveLogoSrc } from "@/lib/assetPaths"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_cost_config.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_config.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_server_cost_config.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_config.tsx index 3f3986d5a2e..89c41693a4c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_cost_config.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_config.tsx @@ -2,7 +2,7 @@ import React from "react"; import { Tooltip, InputNumber, Collapse, Badge } from "antd"; import { InfoCircleOutlined, DollarOutlined, ToolOutlined } from "@ant-design/icons"; import { Card, Title, Text } from "@tremor/react"; -import { MCPServerCostInfo } from "./types"; +import { MCPServerCostInfo } from "@/components/mcp_tools/types"; interface MCPServerCostConfigProps { value?: MCPServerCostInfo; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_cost_display.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_display.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_server_cost_display.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_display.tsx index e41a06b879e..f26f7ba2320 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_cost_display.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_display.tsx @@ -1,6 +1,6 @@ import React from "react"; import { Text } from "@tremor/react"; -import { MCPServerCostInfo } from "./types"; +import { MCPServerCostInfo } from "@/components/mcp_tools/types"; interface MCPServerCostDisplayProps { costConfig?: MCPServerCostInfo | null; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx index 5b8d0aac8c0..278bf6f7e13 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx @@ -4,18 +4,18 @@ import { render, screen, waitFor, fireEvent, act } from "@testing-library/react" import userEvent from "@testing-library/user-event"; import MCPServerEdit, { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit"; import { setSecureItem } from "@/utils/secureStorage"; -import * as networking from "../networking"; -import NotificationsManager from "../molecules/notifications_manager"; +import * as networking from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; import { selectAntOption } from "./testUtils"; -vi.mock("../networking", () => ({ +vi.mock("@/components/networking", () => ({ updateMCPServer: vi.fn(), listMCPTools: vi.fn().mockResolvedValue({ tools: [], error: null }), storeMCPOAuthUserCredential: vi.fn().mockResolvedValue({}), testMCPToolsListRequest: vi.fn().mockResolvedValue({ tools: [], error: null }), })); -vi.mock("../molecules/notifications_manager", () => ({ +vi.mock("@/components/molecules/notifications_manager", () => ({ default: { success: vi.fn(), fromBackend: vi.fn(), diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index 6709eb02c65..3b184cc6f1e 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -18,8 +18,13 @@ import { TRANSPORT, getMcpOAuthMode, oauth2FlowToFormValue, -} from "./types"; -import { updateMCPServer, listMCPTools, storeMCPOAuthUserCredential, testMCPToolsListRequest } from "../networking"; +} from "@/components/mcp_tools/types"; +import { + updateMCPServer, + listMCPTools, + storeMCPOAuthUserCredential, + testMCPToolsListRequest, +} from "@/components/networking"; import { getToken, isTokenValid, removeToken, setToken } from "@/utils/mcpTokenStore"; import { buildMcpPassthroughAuthHeader } from "@/utils/mcpHeaderUtils"; import MCPServerCostConfig from "./mcp_server_cost_config"; @@ -39,7 +44,7 @@ import { normalizeToolOverrideMap, TOOL_DISPLAY_NAME_PATTERN, } from "./utils"; -import NotificationsManager from "../molecules/notifications_manager"; +import NotificationsManager from "@/components/molecules/notifications_manager"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx index 7ed00f74225..620f76a739a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx @@ -2,7 +2,7 @@ import React, { useState } from "react"; import { ArrowLeftIcon, EyeIcon, EyeOffIcon } from "@heroicons/react/outline"; import { Title, Card, Button, Text, Grid, TabGroup, TabList, TabPanel, TabPanels, Tab, Icon } from "@tremor/react"; -import { MCPServer, handleTransport, handleAuth } from "./types"; +import { MCPServer, handleTransport, handleAuth } from "@/components/mcp_tools/types"; // TODO: Move Tools viewer from index file import { MCPToolsViewer } from "."; import MCPServerEdit, { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx index 446fdd8c22d..d61bc23c757 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx @@ -3,10 +3,10 @@ import { render, waitFor, screen, fireEvent, act } from "@testing-library/react" import { describe, it, expect, vi, beforeEach } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import MCPServers from "./mcp_servers"; -import * as networking from "../networking"; +import * as networking from "@/components/networking"; // Mock the networking module -vi.mock("../networking", () => ({ +vi.mock("@/components/networking", () => ({ fetchMCPServers: vi.fn(), fetchMCPServerHealth: vi.fn(), deleteMCPServer: vi.fn(), @@ -19,7 +19,7 @@ vi.mock("../networking", () => ({ })); // Mock NotificationsManager -vi.mock("../molecules/notifications_manager", () => ({ +vi.mock("@/components/molecules/notifications_manager", () => ({ default: { success: vi.fn(), fromBackend: vi.fn(), diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index f0d3e60b32e..0afb4bd9314 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -1,29 +1,35 @@ import { isAdminRole } from "@/utils/roles"; import { QuestionCircleOutlined, SearchOutlined } from "@ant-design/icons"; import { Button, Tab, TabGroup, TabList, TabPanel, TabPanels, Text, Title } from "@tremor/react"; -import NewBadge from "../common_components/NewBadge"; +import NewBadge from "@/components/common_components/NewBadge"; import { Descriptions, Empty, Input, Modal, Select, Spin, Tooltip, Typography } from "antd"; import React, { useEffect, useState, useMemo, useCallback } from "react"; import { useQuery } from "@tanstack/react-query"; -import { useMCPServers } from "../../app/(dashboard)/hooks/mcpServers/useMCPServers"; -import { useMCPServerHealth } from "../../app/(dashboard)/hooks/mcpServers/useMCPServerHealth"; -import NotificationsManager from "../molecules/notifications_manager"; -import { deleteMCPServer } from "../networking"; +import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers"; +import { useMCPServerHealth } from "@/app/(dashboard)/hooks/mcpServers/useMCPServerHealth"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import { deleteMCPServer } from "@/components/networking"; import { MCPSubmissionsTab } from "./MCPSubmissionsTab"; import { MCPToolsetsTab } from "./MCPToolsetsTab"; import CreateMCPServer from "./create_mcp_server"; import MCPConnect from "./mcp_connect"; import MCPServerCard from "./MCPServerCard"; import { MCPServerView } from "./mcp_server_view"; -import type { DiscoverableMCPServer, MCPServer, MCPServerProps, MCPUserEnvVarsStatus, Team } from "./types"; -import MCPSemanticFilterSettings from "../Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterSettings"; +import type { + DiscoverableMCPServer, + MCPServer, + MCPServerProps, + MCPUserEnvVarsStatus, + Team, +} from "@/components/mcp_tools/types"; +import MCPSemanticFilterSettings from "@/components/Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterSettings"; import MCPNetworkSettings from "./MCPNetworkSettings"; import MCPDiscovery from "./mcp_discovery"; -import { ByokCredentialModal } from "./ByokCredentialModal"; +import { ByokCredentialModal } from "@/components/mcp_tools/ByokCredentialModal"; import { getSecureItem } from "@/utils/secureStorage"; import { TOOLS_OAUTH_UI_STATE_KEY } from "@/hooks/mcpOAuthUtils"; import UserEnvVarsModal from "./UserEnvVarsModal"; -import { listMCPUserEnvVarStatus } from "../networking"; +import { listMCPUserEnvVarStatus } from "@/components/networking"; type SortKey = "created_desc" | "updated_desc" | "name_asc" | "health"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.test.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.tsx index 1ebc07eac86..60c4c264c3c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tool_configuration.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.tsx @@ -2,7 +2,7 @@ import React, { useEffect, useMemo, useRef, useState } from "react"; import { Card, Title, Text } from "@tremor/react"; import { ToolOutlined, CheckCircleOutlined, SearchOutlined, EditOutlined } from "@ant-design/icons"; import { Badge, Spin, Checkbox, Input, Radio } from "antd"; -import McpCrudPermissionPanel from "./McpCrudPermissionPanel"; +import McpCrudPermissionPanel from "@/components/mcp_tools/McpCrudPermissionPanel"; import { TOOL_DISPLAY_NAME_PATTERN } from "./utils"; interface KeyTool { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx similarity index 98% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx index 2e8fa901f6d..8b0e6d62f66 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx @@ -2,10 +2,10 @@ import { render, screen, waitFor } from "@testing-library/react"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { describe, expect, it, vi, beforeEach } from "vitest"; import MCPToolsViewer from "./mcp_tools"; -import { listMCPTools, getMCPOAuthUserCredentialStatus } from "../networking"; +import { listMCPTools, getMCPOAuthUserCredentialStatus } from "@/components/networking"; import { isTokenValid, getToken } from "@/utils/mcpTokenStore"; -vi.mock("../networking", () => ({ +vi.mock("@/components/networking", () => ({ listMCPTools: vi.fn(), callMCPTool: vi.fn(), getMCPOAuthUserCredentialStatus: vi.fn(), diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx similarity index 99% rename from ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx index 928a2e3c6bb..428c10da284 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx @@ -9,8 +9,8 @@ import { MCPContent, CallMCPToolResponse, getMcpOAuthMode, -} from "./types"; -import { listMCPTools, callMCPTool, getMCPOAuthUserCredentialStatus } from "../networking"; +} from "@/components/mcp_tools/types"; +import { listMCPTools, callMCPTool, getMCPOAuthUserCredentialStatus } from "@/components/networking"; import { isTokenValid, getToken, removeToken } from "@/utils/mcpTokenStore"; import { sanitizeMcpAliasForHeader, buildMcpPassthroughAuthHeader } from "@/utils/mcpHeaderUtils"; import { useToolsOAuthFlow } from "@/hooks/useToolsOAuthFlow"; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/testUtils.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/testUtils.ts similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/testUtils.ts rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/testUtils.ts diff --git a/ui/litellm-dashboard/src/components/mcp_tools/utils.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/mcp_tools/utils.test.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx diff --git a/ui/litellm-dashboard/src/components/mcp_tools/utils.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx similarity index 97% rename from ui/litellm-dashboard/src/components/mcp_tools/utils.tsx rename to ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx index 7d6e24fc480..4738e1e8fba 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/utils.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx @@ -1,4 +1,4 @@ -import { MCPEnvVar, MCPEnvVarScope } from "./types"; +import { MCPEnvVar, MCPEnvVarScope } from "@/components/mcp_tools/types"; export const extractMCPToken = (url: string): { token: string | null; baseUrl: string } => { try { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx index dfc7ca15896..462c48360cd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx @@ -1,6 +1,6 @@ "use client"; -import { MCPServers } from "@/components/mcp_tools"; +import { MCPServers } from "./_components"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; export default function McpServers() { diff --git a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx index e871bbe8272..62346eff057 100644 --- a/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx +++ b/ui/litellm-dashboard/tests/CreateKeyPage.expiredToken.test.tsx @@ -130,7 +130,7 @@ vi.mock("@/components/cache_dashboard", () => ({ default: stub("cache-dashboard" vi.mock("@/app/(dashboard)/guardrails/_components", () => ({ default: stub("guardrails") })); vi.mock("@/components/prompts", () => ({ default: stub("prompts") })); vi.mock("@/components/transform_request", () => ({ default: stub("transform-request") })); -vi.mock("@/components/mcp_tools", () => ({ MCPServers: stub("mcp-servers") })); +vi.mock("@/app/(dashboard)/mcp-servers/_components", () => ({ MCPServers: stub("mcp-servers") })); vi.mock("@/app/(dashboard)/tag-management/_components", () => ({ default: stub("tag-management") })); vi.mock("@/app/(dashboard)/vector-stores/_components", () => ({ default: stub("vector-stores") })); vi.mock("@/components/ui_theme_settings", () => ({ default: stub("ui-theme-settings") })); From 0008d96af4d969cc497d829f1a635f7cd3a7df5c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 18:07:37 -0700 Subject: [PATCH 029/123] docs(github): add Final Attestation and per-test sanity-check step to QA runbook --- .github/pull_request_template.md | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index df44a4063a1..d7e80b32749 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -46,7 +46,7 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac + +### Final Attestation + +- [ ] The tests check the right things, including the edge cases, and regressions in the respective real-world customer use-cases are not possible after this PR From 5fc1a3c671bf9f25e1410efd4ebb9c2617a9219a Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 11 Jul 2026 18:16:18 -0700 Subject: [PATCH 030/123] refactor(ui): convert endpoint usage charts to shadcn/recharts (#32723) * refactor(ui): convert endpoint usage charts to shadcn/recharts Adds a LineChart wrapper to the shared charts kit, mirroring the BarChart/AreaChart composition with connectNulls and curveType props, and converts EndpointUsageBarChart and EndpointUsageLineChart from tremor to the shared wrappers. Both endpoint chart tests now assert on real recharts SVG output instead of tremor mocks. * refactor(ui): drop unused endpointData prop from EndpointUsageLineChart * fix(ui): point endpoint chart test type imports at the UsagePage types alias after colocation move --- ui/litellm-dashboard/eslint-suppressions.json | 103 +++++++-------- .../EndpointUsage/EndpointUsage.tsx | 2 +- .../components/EndpointUsageBarChart.test.tsx | 89 ++++++++----- .../components/EndpointUsageBarChart.tsx | 45 +++---- .../EndpointUsageLineChart.test.tsx | 119 ++++++++++++++---- .../components/EndpointUsageLineChart.tsx | 53 +++++--- .../src/components/shared/charts/index.ts | 1 + .../shared/charts/line_chart.test.tsx | 117 +++++++++++++++++ .../components/shared/charts/line_chart.tsx | 99 +++++++++++++++ 9 files changed, 473 insertions(+), 155 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/shared/charts/line_chart.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/charts/line_chart.tsx diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index eb8bdee0ea2..3f6f4da25d0 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -858,9 +858,6 @@ "src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx": { "no-nested-ternary": { "count": 3 - }, - "no-restricted-imports": { - "count": 1 } }, "src/app/(dashboard)/projects/_components/ProjectKeysSection.tsx": { @@ -1079,6 +1076,51 @@ "count": 1 } }, + "src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, + "src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/usage/_components/components/EntityUsage/SpendByProvider.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "react-hooks/immutability": { + "count": 1 + } + }, + "src/app/(dashboard)/usage/_components/components/UsagePageView.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/purity": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 3 + } + }, + "src/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity.ts": { + "react-hooks/refs": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 1 + } + }, "src/app/(dashboard)/users/_components/DefaultUserSettings.tsx": { "no-restricted-imports": { "count": 1 @@ -1502,61 +1544,6 @@ "count": 1 } }, - "src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx": { - "no-nested-ternary": { - "count": 1 - } - }, - "src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/usage/_components/components/EntityUsage/SpendByProvider.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/usage/_components/components/UsageAIChatPanel.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "react-hooks/immutability": { - "count": 1 - } - }, - "src/app/(dashboard)/usage/_components/components/UsagePageView.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/purity": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 3 - } - }, - "src/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity.ts": { - "react-hooks/refs": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, "src/components/VirtualKeysPage/VirtualKeysTable.tsx": { "no-nested-ternary": { "count": 2 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx index 64fdb13a0b9..51e451ca770 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx @@ -59,7 +59,7 @@ const EndpointUsage: React.FC = ({ userSpendData }) => {
- +
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx index f200b8a6a18..a9e65b21f4b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx @@ -1,39 +1,70 @@ -import { render, screen } from "@testing-library/react"; -import { describe, expect, it, vi } from "vitest"; +import { screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; +import { renderWithProviders } from "@/../tests/test-utils"; +import { MetricWithMetadata } from "@/components/UsagePage/types"; import EndpointUsageBarChart from "./EndpointUsageBarChart"; -vi.mock("@tremor/react", async () => { - const React = await import("react"); - - function Card({ children }: any) { - return React.createElement("div", { "data-testid": "tremor-card" }, children); - } - (Card as any).displayName = "Card"; - - function Title({ children }: any) { - return React.createElement("h2", { "data-testid": "tremor-title" }, children); - } - (Title as any).displayName = "Title"; - - function BarChart(_props: any) { - return React.createElement("div", { "data-testid": "tremor-bar-chart" }, "Bar Chart"); - } - (BarChart as any).displayName = "BarChart"; - - return { Card, Title, BarChart }; +const metric = (successful: number, failed: number): MetricWithMetadata => ({ + metrics: { + spend: 0, + prompt_tokens: 0, + completion_tokens: 0, + total_tokens: 0, + api_requests: successful + failed, + successful_requests: successful, + failed_requests: failed, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + }, + metadata: {}, + api_key_breakdown: {}, }); -vi.mock("@/components/common_components/chartUtils", () => ({ - CustomLegend: ({ categories }: any) =>
{categories.join(", ")}
, - CustomTooltip: () =>
Tooltip
, -})); +const endpointData = { + "/chat/completions": metric(120, 5), + "/embeddings": metric(40, 2), +}; describe("EndpointUsageBarChart", () => { - it("should render", () => { - render(); + it("renders the title and the header legend labels", () => { + renderWithProviders(); - expect(screen.getByTestId("tremor-card")).toBeInTheDocument(); expect(screen.getByText("Success vs Failed Requests by Endpoint")).toBeInTheDocument(); - expect(screen.getByTestId("tremor-bar-chart")).toBeInTheDocument(); + expect(screen.getByText("Successful Requests")).toBeInTheDocument(); + expect(screen.getByText("Failed Requests")).toBeInTheDocument(); + }); + + it("renders stacked green and red bars per endpoint", () => { + const { container } = renderWithProviders(); + + expect(container.querySelectorAll(".recharts-bar")).toHaveLength(2); + const rectangles = Array.from(container.querySelectorAll("path.recharts-rectangle")); + expect(rectangles).toHaveLength(4); + const fills = new Set(rectangles.map((rect) => rect.getAttribute("fill"))); + expect(fills).toEqual(new Set(["var(--color-green-500, #22c55e)", "var(--color-red-500, #ef4444)"])); + + const xPositions = rectangles.map((rect) => rect.getAttribute("d")?.split(",")[0]); + expect(new Set(xPositions).size).toBe(2); + }); + + it("labels the x axis with endpoint names", () => { + renderWithProviders(); + + expect(screen.getAllByText("/chat/completions").length).toBeGreaterThan(0); + expect(screen.getAllByText("/embeddings").length).toBeGreaterThan(0); + }); + + it("keeps the chart's own legend off; only the header legend is shown", () => { + const { container } = renderWithProviders(); + + expect(container.querySelector(".recharts-legend-wrapper")).toBeNull(); + expect(screen.queryByText("metrics.successful_requests")).not.toBeInTheDocument(); + }); + + it("renders an empty chart without bars when endpointData is absent", () => { + const { container } = renderWithProviders(); + + expect(screen.getByText("Success vs Failed Requests by Endpoint")).toBeInTheDocument(); + expect(container.querySelectorAll("path.recharts-rectangle")).toHaveLength(0); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx index 2badbe30868..bf9868d77cf 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx @@ -1,6 +1,6 @@ import React from "react"; -import { BarChart, Card, Title } from "@tremor/react"; -import { CustomLegend, CustomTooltip } from "@/components/common_components/chartUtils"; +import { BarChart, CustomLegend, CustomTooltip } from "@/components/shared/charts"; +import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { MetricWithMetadata } from "@/components/UsagePage/types"; interface EndpointUsageBarChartProps { @@ -8,11 +8,9 @@ interface EndpointUsageBarChartProps { } const EndpointUsageBarChart: React.FC = ({ endpointData }) => { - const dataToUse = endpointData || {}; - // Transform endpoint data into chart format const chartData = React.useMemo(() => { - return Object.entries(dataToUse).map(([endpoint, data]) => ({ + return Object.entries(endpointData || {}).map(([endpoint, data]) => ({ endpoint, "metrics.successful_requests": data.metrics.successful_requests, "metrics.failed_requests": data.metrics.failed_requests, @@ -21,31 +19,34 @@ const EndpointUsageBarChart: React.FC = ({ endpointD failed_requests: data.metrics.failed_requests, }, })); - }, [dataToUse]); + }, [endpointData]); const valueFormatter = (value: number) => value.toLocaleString(); return ( -
- Success vs Failed Requests by Endpoint - +
+ Success vs Failed Requests by Endpoint + +
+ + + -
- +
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx index ec99825289c..33914e627dc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx @@ -1,34 +1,103 @@ -import { render, screen } from "@testing-library/react"; -import { describe, expect, it, vi } from "vitest"; +import { screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; +import { renderWithProviders } from "@/../tests/test-utils"; +import { DailyData, MetricWithMetadata, SpendMetrics } from "@/components/UsagePage/types"; import EndpointUsageLineChart from "./EndpointUsageLineChart"; -vi.mock("@tremor/react", async () => { - const React = await import("react"); - - function Card({ children }: any) { - return React.createElement("div", { "data-testid": "tremor-card" }, children); - } - (Card as any).displayName = "Card"; - - function Title({ children }: any) { - return React.createElement("h2", { "data-testid": "tremor-title" }, children); - } - (Title as any).displayName = "Title"; - - function LineChart(_props: any) { - return React.createElement("div", { "data-testid": "tremor-line-chart" }, "Line Chart"); - } - (LineChart as any).displayName = "LineChart"; - - return { Card, Title, LineChart }; +const spendMetrics = (apiRequests: number): SpendMetrics => ({ + spend: 0, + prompt_tokens: 0, + completion_tokens: 0, + total_tokens: 0, + api_requests: apiRequests, + successful_requests: apiRequests, + failed_requests: 0, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, }); +const endpointMetric = (apiRequests: number): MetricWithMetadata => ({ + metrics: spendMetrics(apiRequests), + metadata: {}, + api_key_breakdown: {}, +}); + +const day = (date: string, endpoints: Record): DailyData => ({ + date, + metrics: spendMetrics(0), + breakdown: { + models: {}, + model_groups: {}, + mcp_servers: {}, + providers: {}, + api_keys: {}, + entities: {}, + endpoints: Object.fromEntries( + Object.entries(endpoints).map(([name, requests]) => [name, endpointMetric(requests)]), + ), + }, +}); + +const dailyData = { + results: [ + day("2026-06-03T12:00:00", { "/chat/completions": 4000, "/embeddings": 900 }), + day("2026-06-02T12:00:00", { "/chat/completions": 2500, "/embeddings": 700 }), + day("2026-06-01T12:00:00", { "/chat/completions": 1200 }), + ], +}; + describe("EndpointUsageLineChart", () => { - it("should render", () => { - render(); + it("renders the title", () => { + renderWithProviders(); - expect(screen.getByTestId("tremor-card")).toBeInTheDocument(); expect(screen.getByText("Endpoint Usage Trends")).toBeInTheDocument(); - expect(screen.getByTestId("tremor-line-chart")).toBeInTheDocument(); + }); + + it("renders one line per endpoint with the tremor palette strokes", () => { + const { container } = renderWithProviders(); + + const curves = Array.from(container.querySelectorAll("path.recharts-line-curve")); + expect(curves).toHaveLength(2); + expect(new Set(curves.map((curve) => curve.getAttribute("stroke")))).toEqual( + new Set(["var(--color-blue-500, #3b82f6)", "var(--color-cyan-500, #06b6d4)"]), + ); + }); + + it("shows a legend with the endpoint names", () => { + const { container } = renderWithProviders(); + + const legend = container.querySelector(".recharts-legend-wrapper"); + expect(legend).not.toBeNull(); + expect(legend!.textContent).toContain("/chat/completions"); + expect(legend!.textContent).toContain("/embeddings"); + }); + + it("orders formatted dates oldest to newest on the x axis", () => { + const { container } = renderWithProviders(); + + const tickLabels = Array.from(container.querySelectorAll(".recharts-xAxis-tick-labels text")).map( + (tick) => tick.textContent, + ); + expect(tickLabels).toEqual(["Jun 1", "Jun 2", "Jun 3"]); + }); + + it("formats y axis ticks with toLocaleString", () => { + renderWithProviders(); + + expect(screen.getAllByText(/^\d,\d{3}$/).length).toBeGreaterThan(0); + }); + + it("draws smooth natural curves", () => { + const { container } = renderWithProviders(); + + const path = container.querySelector("path.recharts-line-curve")?.getAttribute("d") ?? ""; + expect(path).toContain("C"); + }); + + it("renders an empty chart without lines when dailyData is absent", () => { + const { container } = renderWithProviders(); + + expect(screen.getByText("Endpoint Usage Trends")).toBeInTheDocument(); + expect(container.querySelectorAll("path.recharts-line-curve")).toHaveLength(0); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx index 9f838d6156a..483a30b1639 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx @@ -1,10 +1,10 @@ -import { Card, LineChart, Title } from "@tremor/react"; import { useMemo } from "react"; +import { LineChart, type ChartColor } from "@/components/shared/charts"; +import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { DailyData } from "@/components/UsagePage/types"; interface EndpointUsageLineChartProps { dailyData?: { results: DailyData[] }; - endpointData?: Record; } // Transform daily data into chart format @@ -42,7 +42,7 @@ function transformDailyDataToChart(dailyData: DailyData[]): Array { if (!dailyData?.results || dailyData.results.length === 0) { return []; @@ -59,26 +59,39 @@ export function EndpointUsageLineChart({ dailyData, endpointData }: EndpointUsag }, [chartData]); // Tremor color palette for multiple lines - const colors = ["blue", "cyan", "indigo", "violet", "purple", "fuchsia", "pink", "rose", "red", "orange"]; + const colors: readonly ChartColor[] = [ + "blue", + "cyan", + "indigo", + "violet", + "purple", + "fuchsia", + "pink", + "rose", + "red", + "orange", + ]; return ( -
- Endpoint Usage Trends -
- value.toLocaleString()} - showLegend={true} - showGridLines={true} - yAxisWidth={60} - connectNulls={true} - curveType="natural" - /> + + Endpoint Usage Trends + + + value.toLocaleString()} + showLegend={true} + showGridLines={true} + yAxisWidth={60} + connectNulls={true} + curveType="natural" + /> +
); } diff --git a/ui/litellm-dashboard/src/components/shared/charts/index.ts b/ui/litellm-dashboard/src/components/shared/charts/index.ts index ba0a7544ddb..8383c767064 100644 --- a/ui/litellm-dashboard/src/components/shared/charts/index.ts +++ b/ui/litellm-dashboard/src/components/shared/charts/index.ts @@ -10,3 +10,4 @@ export { } from "./chart_tooltip"; export { CHART_COLOR_HEX, DEFAULT_COLOR_CYCLE, categoryFills, chartColorValue, type ChartColor } from "./colors"; export { DonutChart, type DonutChartProps } from "./donut_chart"; +export { LineChart, type LineChartCurveType, type LineChartProps } from "./line_chart"; diff --git a/ui/litellm-dashboard/src/components/shared/charts/line_chart.test.tsx b/ui/litellm-dashboard/src/components/shared/charts/line_chart.test.tsx new file mode 100644 index 00000000000..9385dc49494 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/charts/line_chart.test.tsx @@ -0,0 +1,117 @@ +import { render, screen } from "@testing-library/react"; +import React from "react"; +import { describe, expect, it } from "vitest"; +import { LineChart } from "./line_chart"; + +const data = [ + { date: "Jun 1", "/chat/completions": 10, "/embeddings": 4 }, + { date: "Jun 2", "/chat/completions": 15, "/embeddings": 6 }, + { date: "Jun 3", "/chat/completions": 12, "/embeddings": 9 }, +]; + +describe("LineChart", () => { + it("renders one line per category with the mapped tremor stroke colors", () => { + const { container } = render( + , + ); + + const curves = Array.from(container.querySelectorAll("path.recharts-line-curve")); + expect(curves).toHaveLength(2); + expect(curves.map((curve) => curve.getAttribute("stroke"))).toEqual([ + "var(--color-blue-500, #3b82f6)", + "var(--color-cyan-500, #06b6d4)", + ]); + }); + + it("falls back to the tremor default color cycle when no colors are passed", () => { + const { container } = render( + , + ); + + const strokes = Array.from(container.querySelectorAll("path.recharts-line-curve")).map((curve) => + curve.getAttribute("stroke"), + ); + expect(strokes).toEqual(["var(--color-blue-500, #3b82f6)", "var(--color-cyan-500, #06b6d4)"]); + }); + + it("applies valueFormatter to the value axis ticks", () => { + render( + `${v} req`} + />, + ); + + expect(screen.getAllByText(/ req$/).length).toBeGreaterThan(0); + }); + + it("renders a legend by default, matching tremor, and hides it when showLegend is false", () => { + const { container, rerender } = render( + , + ); + expect(screen.getByText("/chat/completions")).toBeInTheDocument(); + expect(container.querySelector(".recharts-legend-wrapper")).not.toBeNull(); + + rerender( + , + ); + expect(screen.queryByText("/chat/completions")).not.toBeInTheDocument(); + }); + + it("draws straight segments by default and curved segments for curveType natural", () => { + const { container: linear } = render( + , + ); + const { container: natural } = render( + , + ); + + const linearPath = linear.querySelector("path.recharts-line-curve")?.getAttribute("d") ?? ""; + const naturalPath = natural.querySelector("path.recharts-line-curve")?.getAttribute("d") ?? ""; + expect(linearPath).not.toContain("C"); + expect(naturalPath).toContain("C"); + }); + + it("bridges gaps over null values only when connectNulls is set", () => { + const gappedData = [ + { date: "Jun 1", "/chat/completions": 10 }, + { date: "Jun 2", "/chat/completions": null }, + { date: "Jun 3", "/chat/completions": 12 }, + { date: "Jun 4", "/chat/completions": 15 }, + ]; + + const { container: broken } = render( + , + ); + const { container: bridged } = render( + , + ); + + const brokenPath = broken.querySelector("path.recharts-line-curve")?.getAttribute("d") ?? ""; + const bridgedPath = bridged.querySelector("path.recharts-line-curve")?.getAttribute("d") ?? ""; + expect((brokenPath.match(/M/g) ?? []).length).toBeGreaterThan(1); + expect((bridgedPath.match(/M/g) ?? []).length).toBe(1); + }); + + it("renders an empty chart without lines when there are no categories", () => { + const { container } = render(); + + expect(container.querySelector("[data-slot='chart']")).not.toBeNull(); + expect(container.querySelectorAll("path.recharts-line-curve")).toHaveLength(0); + }); + + it("emits no per-chart style tag; colors flow through strokes, not CSS vars", () => { + const { container } = render( + , + ); + expect(container.querySelector("style")).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/charts/line_chart.tsx b/ui/litellm-dashboard/src/components/shared/charts/line_chart.tsx new file mode 100644 index 00000000000..2dc8747118d --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/charts/line_chart.tsx @@ -0,0 +1,99 @@ +"use client"; + +import * as React from "react"; +import { CartesianGrid, Line, LineChart as RechartsLineChart, XAxis, YAxis } from "recharts"; +import { ChartContainer, ChartLegend, ChartLegendContent, ChartTooltip, type ChartConfig } from "@/components/ui/chart"; +import { cn } from "@/lib/cva.config"; +import { ValueTooltip, type ChartTooltipComponent } from "./chart_tooltip"; +import { categoryFills, type ChartColor } from "./colors"; + +export type LineChartCurveType = "linear" | "natural" | "monotone" | "step"; + +export type LineChartProps> = { + data: readonly TDatum[]; + index: string; + categories: readonly string[]; + colors?: readonly ChartColor[]; + valueFormatter?: (value: number) => string; + yAxisWidth?: number; + tickGap?: number; + showLegend?: boolean; + showXAxis?: boolean; + showGridLines?: boolean; + showTooltip?: boolean; + customTooltip?: ChartTooltipComponent; + connectNulls?: boolean; + curveType?: LineChartCurveType; + className?: string; + style?: React.CSSProperties; +}; + +export function LineChart>({ + data, + index, + categories, + colors, + valueFormatter, + yAxisWidth = 56, + tickGap = 5, + showLegend = true, + showXAxis = true, + showGridLines = true, + showTooltip = true, + customTooltip, + connectNulls = false, + curveType = "linear", + className, + style, +}: LineChartProps) { + const fills = categoryFills(categories.length, colors); + const config: ChartConfig = Object.fromEntries(categories.map((category) => [category, { label: category }])); + const TooltipContent = customTooltip ?? ValueTooltip; + + return ( + + + {showGridLines && } + + + {showTooltip && ( + ( + + )} + /> + )} + {showLegend && ( + } + /> + )} + {categories.map((category, i) => ( + + ))} + + + ); +} From 0de308a5b8f09af6d653bee4cdb9765cdf2c55ff Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 18:18:37 -0700 Subject: [PATCH 031/123] fix(auto_router): filter embedding models out of tier selects, require all tiers, add inline validation The Add Auto Router complexity tab let chat models fill the embedding-model slot (and vice versa) since neither dropdown filtered on ModelGroup.mode, and submit only required at least one of the four tiers instead of all four. Adds getMissingTiersError alongside the existing getSemanticConfigError, and highlights unfilled tier/embedding selects inline once a submit attempt fails. --- .../add_model/ComplexityRouterConfig.test.tsx | 35 ++++++++++-- .../add_model/ComplexityRouterConfig.tsx | 22 ++++++-- .../SemanticKeywordMatching.test.tsx | 53 +++++++++++++++++++ .../add_model/SemanticKeywordMatching.tsx | 12 ++++- .../add_model/add_auto_router_tab.tsx | 17 ++++-- .../build_complexity_router_config.test.ts | 26 +++++++++ .../build_complexity_router_config.ts | 8 +++ 7 files changed, 160 insertions(+), 13 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.test.tsx diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index a433808085b..0613b0c02ae 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -4,9 +4,10 @@ import { vi } from "vitest"; import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; const mockModelInfo = [ - { model_group: "gpt-4" }, - { model_group: "gpt-3.5-turbo" }, - { model_group: "claude-3-opus" }, + { model_group: "gpt-4", mode: "chat" }, + { model_group: "gpt-3.5-turbo", mode: "chat" }, + { model_group: "claude-3-opus", mode: "chat" }, + { model_group: "text-embedding-3-small", mode: "embedding" }, ] as any[]; const defaultValue: ComplexityRouterConfigValue = { @@ -207,4 +208,32 @@ describe("ComplexityRouterConfig", () => { await user.click(screen.getByRole("switch")); expect(onSemanticMatchingEnabledChange).toHaveBeenCalledWith(true, expect.anything()); }); + + it("excludes embedding-mode models from the tier and classifier dropdowns", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const simpleTierSection = screen.getByText("Simple Tier").closest(".mb-4") as HTMLElement; + const combobox = within(simpleTierSection).getByRole("combobox"); + await user.click(combobox); + + expect((await screen.findAllByText("gpt-3.5-turbo")).length).toBeGreaterThan(0); + expect(screen.queryAllByText("text-embedding-3-small")).toHaveLength(0); + }); + + it("does not show tier validation errors by default", () => { + renderWithProviders(); + expect(screen.queryByText("This tier is required")).not.toBeInTheDocument(); + }); + + it("shows a validation error only under unfilled tiers when showValidationErrors is true", () => { + renderWithProviders( + , + ); + expect(screen.getAllByText("This tier is required")).toHaveLength(1); + }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 555648db8ad..18c575c7c4c 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -45,6 +45,7 @@ interface ComplexityRouterConfigProps { onEmbeddingModelChange?: (model: string) => void; matchThreshold?: number; onMatchThresholdChange?: (threshold: number) => void; + showValidationErrors?: boolean; } const TIER_DESCRIPTIONS: Record = { @@ -84,12 +85,15 @@ const ComplexityRouterConfig: React.FC = ({ onEmbeddingModelChange = () => {}, matchThreshold = 0.5, onMatchThresholdChange = () => {}, + showValidationErrors = false, }) => { - // Prepare model options for dropdowns - const modelOptions = modelInfo.map((model) => ({ - value: model.model_group, - label: model.model_group, - })); + // Embedding models can't serve a chat-completion role, so they're excluded here. + const modelOptions = modelInfo + .filter((model) => model.mode !== "embedding") + .map((model) => ({ + value: model.model_group, + label: model.model_group, + })); const handleTierChange = (tier: keyof ComplexityTiers, model: string) => { onChange({ @@ -148,6 +152,7 @@ const ComplexityRouterConfig: React.FC = ({ {(Object.keys(TIER_DESCRIPTIONS) as Array).map((tier, index) => { const tierInfo = TIER_DESCRIPTIONS[tier]; + const tierMissing = showValidationErrors && !value.tiers[tier]; return (
{index > 0 && } @@ -170,7 +175,13 @@ const ComplexityRouterConfig: React.FC = ({ showSearch style={{ width: "100%" }} options={modelOptions} + status={tierMissing ? "error" : undefined} /> + {tierMissing && ( + + This tier is required + + )}
); @@ -323,6 +334,7 @@ const ComplexityRouterConfig: React.FC = ({ matchThreshold={matchThreshold} onMatchThresholdChange={onMatchThresholdChange} modelInfo={modelInfo} + showValidationErrors={showValidationErrors} /> )} diff --git a/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.test.tsx b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.test.tsx new file mode 100644 index 00000000000..2336e6faf43 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.test.tsx @@ -0,0 +1,53 @@ +import { renderWithProviders, screen } from "../../../tests/test-utils"; +import userEvent from "@testing-library/user-event"; +import { vi } from "vitest"; +import SemanticKeywordMatching from "./SemanticKeywordMatching"; + +const mockModelInfo = [ + { model_group: "gpt-4", mode: "chat" }, + { model_group: "text-embedding-3-small", mode: "embedding" }, + { model_group: "voyage-3-5", mode: "embedding" }, + { model_group: "legacy-model" }, +] as any[]; + +const baseProps = { + enabled: true, + onEnabledChange: vi.fn(), + embeddingModel: undefined, + onEmbeddingModelChange: vi.fn(), + matchThreshold: 0.5, + onMatchThresholdChange: vi.fn(), + modelInfo: mockModelInfo, +}; + +describe("SemanticKeywordMatching", () => { + it("only lists embedding-mode models in the embedding model dropdown", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const combobox = screen.getByRole("combobox"); + await user.click(combobox); + + expect((await screen.findAllByText("text-embedding-3-small")).length).toBeGreaterThan(0); + expect(screen.getAllByText("voyage-3-5").length).toBeGreaterThan(0); + expect(screen.queryAllByText("gpt-4")).toHaveLength(0); + expect(screen.queryAllByText("legacy-model")).toHaveLength(0); + }); + + it("does not show a validation error by default", () => { + renderWithProviders(); + expect(screen.queryByText("An embedding model is required")).not.toBeInTheDocument(); + }); + + it("shows a validation error when showValidationErrors is true and no embedding model is set", () => { + renderWithProviders(); + expect(screen.getByText("An embedding model is required")).toBeInTheDocument(); + }); + + it("hides the validation error once an embedding model is set", () => { + renderWithProviders( + , + ); + expect(screen.queryByText("An embedding model is required")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx index 0f9907ac6c9..c7583427af6 100644 --- a/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx +++ b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx @@ -15,6 +15,7 @@ interface SemanticKeywordMatchingProps { matchThreshold: number; onMatchThresholdChange: (threshold: number) => void; modelInfo: ModelGroup[]; + showValidationErrors?: boolean; } const SemanticKeywordMatching: React.FC = ({ @@ -25,11 +26,14 @@ const SemanticKeywordMatching: React.FC = ({ matchThreshold, onMatchThresholdChange, modelInfo, + showValidationErrors = false, }) => { - const modelOptions = Array.from(new Set(modelInfo.map((model) => model.model_group))).map((model_group) => ({ + const embeddingModels = modelInfo.filter((model) => model.mode === "embedding"); + const modelOptions = Array.from(new Set(embeddingModels.map((model) => model.model_group))).map((model_group) => ({ value: model_group, label: model_group, })); + const embeddingModelMissing = showValidationErrors && !embeddingModel; return ( @@ -60,7 +64,13 @@ const SemanticKeywordMatching: React.FC = ({ showSearch style={{ width: "100%" }} options={modelOptions} + status={embeddingModelMissing ? "error" : undefined} /> + {embeddingModelMissing && ( + + An embedding model is required + + )}
Minimum match score diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 55c62224a48..ce69f8f7ae3 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -11,7 +11,11 @@ import RouterConfigBuilder from "./RouterConfigBuilder"; import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; import { KeywordTierRule } from "./KeywordTierRules"; import { DEFAULT_MATCH_THRESHOLD } from "./SemanticKeywordMatching"; -import { buildComplexityRouterConfig, getSemanticConfigError } from "./build_complexity_router_config"; +import { + buildComplexityRouterConfig, + getMissingTiersError, + getSemanticConfigError, +} from "./build_complexity_router_config"; import { buildAutoRouterTestTargets, AutoRouterTestTarget } from "./build_auto_router_test_targets"; import AutoRouterConnectionTest from "./auto_router_connection_test"; import NotificationManager from "../molecules/notifications_manager"; @@ -43,6 +47,7 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc const [semanticMatchingEnabled, setSemanticMatchingEnabled] = useState(false); const [embeddingModel, setEmbeddingModel] = useState(undefined); const [matchThreshold, setMatchThreshold] = useState(DEFAULT_MATCH_THRESHOLD); + const [showValidationErrors, setShowValidationErrors] = useState(false); // Semantic router config (existing) const [routerConfig, setRouterConfig] = useState(null); @@ -86,19 +91,22 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc classifier_llm_config: classifierLlmConfig, } = complexityRouterConfig; - const filledTiers = Object.values(tiers).filter(Boolean); - if (filledTiers.length === 0) { - NotificationManager.fromBackend("Please select at least one model for a complexity tier"); + const missingTiersError = getMissingTiersError(tiers); + if (missingTiersError) { + setShowValidationErrors(true); + NotificationManager.fromBackend(missingTiersError); return; } if (classifierType === "llm" && !classifierLlmConfig?.model) { + setShowValidationErrors(true); NotificationManager.fromBackend("Please select a classifier model, or switch back to Heuristic"); return; } const semanticError = getSemanticConfigError({ semanticMatchingEnabled, embeddingModel, keywordTierRules }); if (semanticError) { + setShowValidationErrors(true); NotificationManager.fromBackend(semanticError); return; } @@ -296,6 +304,7 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc onEmbeddingModelChange={setEmbeddingModel} matchThreshold={matchThreshold} onMatchThresholdChange={setMatchThreshold} + showValidationErrors={showValidationErrors} />
) : ( diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index 3c252646b57..e5b547d8240 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -1,5 +1,6 @@ import { buildComplexityRouterConfig, + getMissingTiersError, getSemanticConfigError, BuildComplexityRouterConfigParams, } from "./build_complexity_router_config"; @@ -124,6 +125,31 @@ describe("buildComplexityRouterConfig", () => { }); }); +describe("getMissingTiersError", () => { + it("returns null when all four tiers have a model", () => { + expect(getMissingTiersError(tiers)).toBeNull(); + }); + + it("names the specific missing tier when only one is blank", () => { + expect(getMissingTiersError({ ...tiers, REASONING: "" })).toBe( + "Select a model for the following tier(s): REASONING", + ); + }); + + it("names multiple missing tiers in SIMPLE/MEDIUM/COMPLEX/REASONING order", () => { + expect(getMissingTiersError({ ...tiers, SIMPLE: "", REASONING: "" })).toBe( + "Select a model for the following tier(s): SIMPLE, REASONING", + ); + }); + + it("names all four tiers when none are filled", () => { + const noTiers = { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" }; + expect(getMissingTiersError(noTiers)).toBe( + "Select a model for the following tier(s): SIMPLE, MEDIUM, COMPLEX, REASONING", + ); + }); +}); + describe("getSemanticConfigError", () => { const rule = { id: "r1", keywords: ["k8s"], tier: "REASONING" as const }; diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 82ea4f8c12f..3eddca8c35b 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -30,6 +30,14 @@ export interface ComplexityRouterConfigPayload { match_threshold?: number; } +const TIER_KEYS: Array = ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]; + +export const getMissingTiersError = (tiers: ComplexityTiers): string | null => { + const missing = TIER_KEYS.filter((tier) => !tiers[tier]); + if (missing.length === 0) return null; + return `Select a model for the following tier(s): ${missing.join(", ")}`; +}; + export const getSemanticConfigError = ({ semanticMatchingEnabled, embeddingModel, From 2ed4ceb12e5ec5867d150a0d0eb5b9a97196cd4a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 19:00:03 -0700 Subject: [PATCH 032/123] fix(model-cost-map): anchor the bedrock-claude-ids routing rule to the start of the id --- ...odel_prices_and_context_window_backup.json | 4 +-- model_prices_and_context_window.json | 4 +-- .../test_get_model_cost_map.py | 36 +++++++++++++++++++ 3 files changed, 40 insertions(+), 4 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b4ae842c4e6..773ffb92b59 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -45051,8 +45051,8 @@ "rules": [ { "name": "bedrock-claude-ids", - "pattern": "anthropic\\.claude-", - "description": "Any Bedrock-syntax Claude id: the dotted anthropic.claude- segment appears in bare (anthropic.claude-...), region-prefixed (us./eu./au./jp./apac.) and global.-prefixed ids, for every version. Routes these to bedrock before the bare-id Anthropic rule is consulted.", + "pattern": "^(?:[a-z-]+\\.)?anthropic\\.claude-", + "description": "A Bedrock-syntax Claude id, for every version: anthropic.claude- at the start of the name, optionally behind a single dotted geo segment (us./eu./au./jp./apac./global./us-gov.). Anchored to the start because routing rules see the raw request string and provider inference feeds the proxy's provider/* wildcard access checks: an id under an unrecognized namespace such as bedrockz/anthropic.claude-... must stay unroutable rather than resolve to bedrock and slip through a bedrock/* key. Routes to bedrock before the bare-id Anthropic rule is consulted.", "model_info": { "litellm_provider": "bedrock" } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 94d9f6496bd..6a770998331 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -45284,8 +45284,8 @@ "rules": [ { "name": "bedrock-claude-ids", - "pattern": "anthropic\\.claude-", - "description": "Any Bedrock-syntax Claude id: the dotted anthropic.claude- segment appears in bare (anthropic.claude-...), region-prefixed (us./eu./au./jp./apac.) and global.-prefixed ids, for every version. Routes these to bedrock before the bare-id Anthropic rule is consulted.", + "pattern": "^(?:[a-z-]+\\.)?anthropic\\.claude-", + "description": "A Bedrock-syntax Claude id, for every version: anthropic.claude- at the start of the name, optionally behind a single dotted geo segment (us./eu./au./jp./apac./global./us-gov.). Anchored to the start because routing rules see the raw request string and provider inference feeds the proxy's provider/* wildcard access checks: an id under an unrecognized namespace such as bedrockz/anthropic.claude-... must stay unroutable rather than resolve to bedrock and slip through a bedrock/* key. Routes to bedrock before the bare-id Anthropic rule is consulted.", "model_info": { "litellm_provider": "bedrock" } diff --git a/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py b/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py index bdd71f28b1a..1a38b5dc769 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py +++ b/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py @@ -138,6 +138,42 @@ def test_shipped_backup_carries_the_claude_routing_rules(): set_fallback_generalizations(previous) +def test_shipped_routing_rules_never_match_through_an_unrecognized_namespace(): + """Routing rules decide ``litellm_provider`` for otherwise-unknown ids, and the + proxy's wildcard access check (``can_key_call_model`` with a ``bedrock/*`` key) + trusts that inference: it rebuilds ``{provider}/{model}`` and matches it against + the key's patterns. A routing pattern that matches as a substring lets + ``bedrockz/anthropic.claude-...`` resolve to bedrock and slip through a + ``bedrock/*`` key, so every shipped routing rule must anchor to the start of + the name and never match an id carrying an unrecognized namespace prefix.""" + backup = GetModelCostMap.load_local_model_cost_map() + rules = backup[FALLBACK_GENERALIZATIONS_KEY]["rules"] + + routing_rules = [r for r in rules if "litellm_provider" in r["model_info"]] + assert routing_rules + assert all(r["pattern"].startswith("^") for r in routing_rules) + + previous = list(get_fallback_generalization_rules()) + try: + set_fallback_generalizations(rules) + for bedrock_id in [ + "anthropic.claude-3-5-sonnet-20240620-v1:0", + "anthropic.claude-v2:1", + "us.anthropic.claude-sonnet-4-5-20250929-v1:0", + "us-gov.anthropic.claude-3-5-sonnet-20240620-v1:0", + "global.anthropic.claude-fable-5-20260120-v1:0", + ]: + assert match_routing_generalization(bedrock_id) == "bedrock", bedrock_id + for namespaced in [ + "bedrockz/anthropic.claude-3-5-sonnet-20240620", + "bedrockz/us.anthropic.claude-3-5-sonnet-20240620-v1:0", + "bedrockz/claude-3-5-sonnet-20240620", + ]: + assert match_routing_generalization(namespaced) is None, namespaced + finally: + set_fallback_generalizations(previous) + + def test_shipped_backup_marks_claude_4_6_plus_adaptive_not_4_0(): """Adaptive thinking is data, not code. The bundled backup must carry supports_adaptive_thinking on genuine Claude >= 4.6 entries (every provider From d0d1c0e346fdb5f907666b1fb4aa0b26f21beb9c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 19:00:03 -0700 Subject: [PATCH 033/123] fix(proxy-auth): deny provider-wildcard access inferred through an unrecognized model namespace --- litellm/proxy/auth/auth_checks.py | 13 +++++++++++-- tests/proxy_unit_tests/test_auth_checks.py | 2 ++ 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 230b9b70ff0..93811812901 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -4383,14 +4383,23 @@ def _model_custom_llm_provider_matches_wildcard_pattern(model: str, allowed_mode or - `model=claude-3-5-sonnet-20240620` - `allowed_model_pattern=anthropic/*` + + A model that already carries a namespace get_llm_provider did not consume + (e.g. `bedrockz/anthropic.claude-...`) is never granted here: its provider was + inferred from a fragment of the full string, so rebuilding + `{provider}/{model}` would produce `bedrock/bedrockz/...` and slip an + unrecognized namespace through a `bedrock/*` key. """ try: - model, custom_llm_provider, _, _ = get_llm_provider(model=model) + stripped_model, custom_llm_provider, _, _ = get_llm_provider(model=model) except Exception: return False + if stripped_model == model and "/" in model: + return False + return is_model_allowed_by_pattern( - model=f"{custom_llm_provider}/{model}", + model=f"{custom_llm_provider}/{stripped_model}", allowed_model_pattern=allowed_model_pattern, ) diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index e7136ecb195..e58e6c9694b 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -236,6 +236,8 @@ async def test_can_team_call_model(model, expect_to_work): (["bedrock/*"], "bedrock/anthropic.claude-3-5-sonnet-20240620", True), (["bedrock/*"], "bedrockz/anthropic.claude-3-5-sonnet-20240620", False), (["bedrock/us.*"], "bedrock/us.amazon.nova-micro-v1:0", True), + (["openai/*"], "ft:gpt-4-0613", True), + (["openai/*"], "bedrockz/ft:gpt-4-0613", False), ], ) @pytest.mark.asyncio From 3464b5e7dfaf027e366f95f60c55e89265d6601e Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 19:03:48 -0700 Subject: [PATCH 034/123] fix(auto_router): reset inline validation errors when switching router type --- .../src/components/add_model/add_auto_router_tab.tsx | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index ce69f8f7ae3..4e7a72435bf 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -238,7 +238,14 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc
Router Type - setRouterType(e.target.value)} className="w-full"> + { + setRouterType(e.target.value); + setShowValidationErrors(false); + }} + className="w-full" + >
From d6883d15b0feac1bfe07eaa18414ef14b1a5a7f4 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 19:08:07 -0700 Subject: [PATCH 035/123] fix(auto_router): flag name field and tier fields together on empty submit Clicking Add Auto Router with the name empty returned early with only a toast, so blank tier selects never got their inline error state. The empty-name branch now sets showValidationErrors and triggers antd validation on the name field, so every unfilled mandatory field is flagged at once. Adds a regression test for the tab component. --- .../add_model/add_auto_router_tab.test.tsx | 40 +++++++++++++++++++ .../add_model/add_auto_router_tab.tsx | 2 + 2 files changed, 42 insertions(+) create mode 100644 ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx new file mode 100644 index 00000000000..4713f8c6869 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx @@ -0,0 +1,40 @@ +import { renderWithProviders, screen } from "../../../tests/test-utils"; +import userEvent from "@testing-library/user-event"; +import { vi } from "vitest"; +import { Form } from "antd"; +import AddAutoRouterTab from "./add_auto_router_tab"; +import NotificationManager from "../molecules/notifications_manager"; + +vi.mock("../networking", () => ({ + modelAvailableCall: vi.fn().mockResolvedValue({ data: [] }), +})); + +vi.mock("@/components/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn().mockResolvedValue([]), +})); + +vi.mock("./handle_add_auto_router_submit", () => ({ + handleAddAutoRouterSubmit: vi.fn(), +})); + +vi.mock("../molecules/notifications_manager", () => ({ + default: { fromBackend: vi.fn() }, +})); + +const Harness = () => { + const [form] = Form.useForm(); + return ; +}; + +describe("AddAutoRouterTab", () => { + it("flags every mandatory field when Add Auto Router is clicked with nothing filled", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /add auto router/i })); + + expect(await screen.findByText("Auto router name is required")).toBeInTheDocument(); + expect(screen.getAllByText("This tier is required")).toHaveLength(4); + expect(NotificationManager.fromBackend).toHaveBeenCalledWith("Please enter an Auto Router Name"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 4e7a72435bf..a74eab0abdd 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -198,6 +198,8 @@ const AddAutoRouterTab: React.FC = ({ form, handleOk, acc const handleAutoRouterSubmit = () => { const name = form.getFieldValue("auto_router_name"); if (!name) { + setShowValidationErrors(true); + form.validateFields(["auto_router_name"]).catch(() => undefined); NotificationManager.fromBackend("Please enter an Auto Router Name"); return; } From 34c6cce70562ded634b27771a11f698c46c184dd Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 10 Jul 2026 14:36:20 -0700 Subject: [PATCH 036/123] feat(mcp): mint gateway-bound envelope at the token endpoint for dcr_bridge oauth_delegate --- .../mcp_server/discoverable_endpoints.py | 91 +++++++++++++- .../mcp_server/test_discoverable_endpoints.py | 116 ++++++++++++++++++ 2 files changed, 206 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ebceb320906..41ed49a7508 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -10,7 +10,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse import httpx from fastapi import APIRouter, Form, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response -from pydantic import BaseModel, ValidationError +from pydantic import BaseModel, SecretStr, ValidationError from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( @@ -37,6 +37,9 @@ from litellm.types.mcp import MCPAuth, MCPCredentials from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( + UpstreamTokenGrant, + ) from litellm.proxy._types import LiteLLM_MCPServerTable, UserAPIKeyAuth # TTL cache for upstream OAuth metadata fetched from pass-through MCP servers. @@ -654,6 +657,86 @@ async def authorize_with_server( return response +def _bridge_grant_from_token_response(token_response: object) -> Optional["UpstreamTokenGrant"]: + """Validate an upstream OAuth token response into a typed grant, or None when it lacks a usable + access token. Each field is isinstance-checked so nothing untyped from ``response.json()`` flows + into the grant.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import + UpstreamTokenGrant, + ) + + if not isinstance(token_response, dict): + return None + access = token_response.get("access_token") + if not isinstance(access, str) or not access: + return None + token_type = token_response.get("token_type") + refresh = token_response.get("refresh_token") + scope = token_response.get("scope") + expires_in = token_response.get("expires_in") + return UpstreamTokenGrant( + access_token=SecretStr(access), + token_type=token_type if isinstance(token_type, str) and token_type else "Bearer", + refresh_token=SecretStr(refresh) if isinstance(refresh, str) and refresh else None, + scope=scope if isinstance(scope, str) and scope else None, + expires_in=expires_in if isinstance(expires_in, int) and expires_in > 0 else None, + ) + + +async def _mint_bridge_delegate_token_response( + request: Request, mcp_server: MCPServer, token_response: object +) -> JSONResponse: + """Return the client-held envelope bearer for a DCR-bridge ``oauth_delegate`` token exchange. + + The envelope binds the caller's litellm identity (resolved from the token request) to the + upstream grant, so the client holds one bearer that later admits it and forwards the upstream + token, with nothing stored server-side. Fails closed with an OAuth ``invalid_request`` when no + litellm identity accompanies the token request rather than minting an identity-less credential. + """ + from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import + build_bridge_token_response, + envelope_keys_from_master_key, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import + EnvelopeIdentity, + SealedEnvelope, + ) + from litellm.proxy.proxy_server import ( + master_key, # noqa: PLC0415 # inline import avoids a module-load circular import + ) + + if not master_key: + raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set") + + user_id = await _extract_user_id_from_request(request) + if not user_id: + raise HTTPException( + status_code=400, + detail={ + "error": "invalid_request", + "error_description": ( + "this server issues a gateway-bound credential; send a litellm credential " + "(x-litellm-api-key or Authorization) on the token request" + ), + }, + ) + + grant = _bridge_grant_from_token_response(token_response) + if grant is None: + raise HTTPException(status_code=502, detail="Upstream token response has no usable access_token") + + now = datetime.now(timezone.utc) + keys = envelope_keys_from_master_key(master_key) + identity = EnvelopeIdentity(user_id=user_id, server_id=mcp_server.server_id) + sealed = build_bridge_token_response(identity, grant, keys, now) + if not isinstance(sealed, SealedEnvelope): + raise HTTPException(status_code=500, detail="Failed to mint the gateway-bound credential") + + expires_in = max(1, int((sealed.expires_at - now).total_seconds())) + body = {"access_token": sealed.token.get_secret_value(), "token_type": "Bearer", "expires_in": expires_in} + return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS) + + async def exchange_token_with_server( request: Request, mcp_server: MCPServer, @@ -791,6 +874,12 @@ async def exchange_token_with_server( mcp_server.server_id, ) + # A DCR-bridge oauth_delegate server hands the client a gateway-bound envelope (identity plus the + # upstream token) instead of the raw upstream token, so the one bearer both admits the caller and + # forwards the upstream credential. Only this mode mints; every other server returns the raw token. + if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: + return await _mint_bridge_delegate_token_response(request, mcp_server, token_response) + result = { "access_token": access_token, "token_type": token_response.get("token_type", "Bearer"), diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index c55a631c7b3..a2e8d693fab 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4362,6 +4362,122 @@ async def test_register_bridge_relay_never_persists(): mock_persist.assert_not_called() +_BRIDGE_MASTER_KEY = "sk-bridge-producer-master-key-0123456789abcdef" + + +async def _exchange_for_bridge_server(server, upstream_body, user_id): + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + + fake_http_response = MagicMock() + fake_http_response.json.return_value = upstream_body + fake_http_response.raise_for_status = MagicMock() + fake_http_client = MagicMock() + fake_http_client.post = AsyncMock(return_value=fake_http_response) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=fake_http_client, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", + new=AsyncMock(return_value=user_id), + ), + patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), + ): + return await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="https://claude.ai/api/mcp/auth_callback", + client_id="dcr-client-123", + client_secret=None, + code_verifier="verifier", + ) + + +@pytest.mark.asyncio +async def test_oauth_delegate_bridge_token_exchange_mints_envelope_not_raw_token(): + """A dcr_bridge oauth_delegate token exchange returns a gateway-bound envelope, not the raw + upstream token: the response access_token opens (under the same master-key-derived keys and the + server_id) to the caller's identity and the upstream Authorization, and the raw upstream token + never appears in the bearer the client receives.""" + from datetime import datetime, timezone + + from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( + BridgeEnvelopeAdmitted, + envelope_keys_from_master_key, + resolve_bridge_envelope, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} + response = await _exchange_for_bridge_server(server, upstream, user_id="user-77") + + body = json.loads(response.body) + token = body["access_token"] + assert body["token_type"] == "Bearer" + assert body["expires_in"] > 0 + assert token.startswith("llm_env_") + assert "UPSTREAM-SECRET-TOKEN" not in token + assert "refresh_token" not in body + + keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY) + opened = resolve_bridge_envelope(token, keys, datetime.now(timezone.utc), server.server_id) + assert isinstance(opened, BridgeEnvelopeAdmitted) + assert opened.identity.user_id == "user-77" + assert opened.upstream_authorization.get_secret_value() == "Bearer UPSTREAM-SECRET-TOKEN" + + +@pytest.mark.asyncio +async def test_oauth_delegate_bridge_token_exchange_fails_closed_without_litellm_identity(): + """Without a resolvable litellm identity on the token request, the exchange must not mint an + identity-less envelope; it returns an OAuth invalid_request so the client sends a credential.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} + + with pytest.raises(HTTPException) as exc: + await _exchange_for_bridge_server(server, upstream, user_id=None) + + assert exc.value.status_code == 400 + assert exc.value.detail["error"] == "invalid_request" + + +@pytest.mark.asyncio +async def test_true_passthrough_bridge_token_exchange_returns_raw_upstream_token(): + """Only oauth_delegate mints. A true_passthrough dcr_bridge server relays the raw upstream token + to the client, since that mode has no litellm identity to bind and the caller owns the token.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.true_passthrough) + upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} + response = await _exchange_for_bridge_server(server, upstream, user_id="user-77") + + body = json.loads(response.body) + assert body["access_token"] == "UPSTREAM-SECRET-TOKEN" + assert not body["access_token"].startswith("llm_env_") + + +@pytest.mark.asyncio +async def test_non_bridge_oauth_delegate_token_exchange_returns_raw_upstream_token(): + """An oauth_delegate server without dcr_bridge keeps the pre-change contract: the raw upstream + token is returned, so flag-off behavior is byte-identical.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate, dcr_bridge=None) + upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} + response = await _exchange_for_bridge_server(server, upstream, user_id="user-77") + + body = json.loads(response.body) + assert body["access_token"] == "UPSTREAM-SECRET-TOKEN" + + async def _exchange_persistence_attempted_for_auth_type(auth_type) -> bool: """Run exchange_token_with_server for a server of ``auth_type`` and report whether it attempted to persist the exchanged token server-side. The client-forwarded token modes must not persist: From 85255c96fb45fc2473fa5b4ad1939102c8ab9db0 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 10 Jul 2026 17:22:59 -0700 Subject: [PATCH 037/123] feat(mcp): seal the authorizing key hash in the dcr_bridge envelope The mint bound only user_id/server_id into the envelope, which gave admission no way to reload the caller's key and enforce its current restrictions. Seal the hashed authorizing key instead (a one-way digest, not a usable credential), so admission reloads the live UserAPIKeyAuth by it and the key's team/org/tool permissions and revocation apply per request. Extract the token endpoint's key resolution into a shared _resolve_active_litellm_key so the per-user token store (user_id) and the bridge mint (key hash) derive from one active-key-gated path, and fail the mint closed with invalid_request when no active key accompanies the request. --- .../mcp_server/discoverable_endpoints.py | 83 +++++++++++++------ .../mcp_server/test_discoverable_endpoints.py | 78 +++++++++++++++-- 2 files changed, 127 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 41ed49a7508..eeadeb290f3 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -351,44 +351,73 @@ def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> Optional[str]: return key_obj.user_id -async def _extract_user_id_from_request(request: Request) -> Optional[str]: - """Resolve the LiteLLM ``user_id`` at the OAuth token endpoint so a per-user token is stored - under the same identity the egress later reads it by (``user_api_key_auth.user_id``). +async def _resolve_active_litellm_key(request: Request) -> Optional[Tuple[str, "UserAPIKeyAuth"]]: + """Resolve the presented litellm key to ``(its hash, the live active key record)``, or ``None`` + when the key is absent, unresolvable, or blocked/expired. - Resolves authoritatively via ``get_key_object`` (cache first, then DB) instead of a raw cache - peek. On a multi-replica gateway the token-exchange request can land on a worker whose in-memory - cache never saw the key, and a cross-replica Redis hit deserializes to a plain ``dict`` rather - than a ``UserAPIKeyAuth``; the previous code read only ``Authorization`` and did - ``getattr(cached, "user_id")`` with no ``model_type`` rehydration and no DB fallback, so it - silently returned ``None`` and the token was never persisted, which makes the egress 401 on every - reconnect. The resolved key is validated (``_active_key_user_id``) before its identity is trusted, - so a blocked or expired key cannot write. Returns ``None`` when no key is present, the key cannot - be resolved, or it is blocked/expired. + Single resolution path the OAuth token endpoint reuses. Resolves authoritatively via + ``get_key_object`` (cache first, then DB) instead of a raw cache peek. On a multi-replica gateway + the token-exchange request can land on a worker whose in-memory cache never saw the key, and a + cross-replica Redis hit deserializes to a plain ``dict`` rather than a ``UserAPIKeyAuth``; the + previous code read only ``Authorization`` and did ``getattr(cached, "user_id")`` with no + ``model_type`` rehydration and no DB fallback, so it silently returned ``None``. The resolved key + is validated (``_active_key_user_id``) before it is trusted, so a blocked or expired key resolves + to ``None``. The returned hash is the value ``get_key_object`` and the cache/DB layer key the + record by. Callers derive the ``user_id`` (per-user token store) or seal the hash (dcr_bridge + envelope) from the result. """ token = _litellm_key_from_request(request) if not token: return None try: - from litellm.proxy._types import hash_token # noqa: PLC0415 - from litellm.proxy.auth.auth_checks import get_key_object # noqa: PLC0415 - from litellm.proxy.proxy_server import ( # noqa: PLC0415 + from litellm.proxy._types import ( # noqa: PLC0415 # inline import avoids a module-load circular import + hash_token, + ) + from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import + get_key_object, + ) + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import prisma_client, user_api_key_cache, ) + key_hash = hash_token(token) key_obj = await get_key_object( - hashed_token=hash_token(token), + hashed_token=key_hash, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, ) - return _active_key_user_id(key_obj) - except Exception as exc: + except Exception as exc: # noqa: BLE001 # fail closed to None on any key-resolution error verbose_logger.debug( - "_extract_user_id_from_request: could not resolve a LiteLLM user_id for the presented " - "key (%s); per-user token will not be stored server-side.", + "_resolve_active_litellm_key: could not resolve the presented key (%s)", type(exc).__name__, ) return None + if _active_key_user_id(key_obj) is None: + return None + return key_hash, key_obj + + +async def _extract_user_id_from_request(request: Request) -> Optional[str]: + """The LiteLLM ``user_id`` for the token request, so a per-user token is stored under the same + identity the egress later reads it by (``user_api_key_auth.user_id``). ``None`` when no active + key is present. See :func:`_resolve_active_litellm_key` for the resolution and active-key gate. + """ + resolved = await _resolve_active_litellm_key(request) + return _active_key_user_id(resolved[1]) if resolved else None + + +async def _extract_active_key_hash_from_request(request: Request) -> Optional[str]: + """The hash of the litellm key that authorized the token request, when it maps to an active key. + + A DCR-bridge envelope seals this hash so admission can reload the live ``UserAPIKeyAuth`` record + and enforce the key's current team/org/tool restrictions and revocation, rather than trusting a + frozen identity. The hash is a one-way digest, not a usable credential (the edge rejects a bare + hash presented as a bearer). ``None`` when no active key is present, so no envelope is minted for + a missing, unresolvable, or revoked key. + """ + resolved = await _resolve_active_litellm_key(request) + return resolved[0] if resolved else None async def _store_per_user_token_server_side( @@ -688,10 +717,12 @@ async def _mint_bridge_delegate_token_response( ) -> JSONResponse: """Return the client-held envelope bearer for a DCR-bridge ``oauth_delegate`` token exchange. - The envelope binds the caller's litellm identity (resolved from the token request) to the + The envelope binds the authorizing litellm key (its hash, resolved from the token request) to the upstream grant, so the client holds one bearer that later admits it and forwards the upstream - token, with nothing stored server-side. Fails closed with an OAuth ``invalid_request`` when no - litellm identity accompanies the token request rather than minting an identity-less credential. + token, with nothing stored server-side. Admission reloads the live key by that hash, so the key's + current restrictions and revocation gate the request. Fails closed with an OAuth + ``invalid_request`` when no active litellm key accompanies the token request rather than minting + an unbound credential. """ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import build_bridge_token_response, @@ -708,8 +739,8 @@ async def _mint_bridge_delegate_token_response( if not master_key: raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set") - user_id = await _extract_user_id_from_request(request) - if not user_id: + key_hash = await _extract_active_key_hash_from_request(request) + if not key_hash: raise HTTPException( status_code=400, detail={ @@ -727,7 +758,7 @@ async def _mint_bridge_delegate_token_response( now = datetime.now(timezone.utc) keys = envelope_keys_from_master_key(master_key) - identity = EnvelopeIdentity(user_id=user_id, server_id=mcp_server.server_id) + identity = EnvelopeIdentity(server_id=mcp_server.server_id, key_hash=key_hash) sealed = build_bridge_token_response(identity, grant, keys, now) if not isinstance(sealed, SealedEnvelope): raise HTTPException(status_code=500, detail="Failed to mint the gateway-bound credential") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index a2e8d693fab..e69d94d9615 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4365,7 +4365,7 @@ async def test_register_bridge_relay_never_persists(): _BRIDGE_MASTER_KEY = "sk-bridge-producer-master-key-0123456789abcdef" -async def _exchange_for_bridge_server(server, upstream_body, user_id): +async def _exchange_for_bridge_server(server, upstream_body, key_hash): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( exchange_token_with_server, ) @@ -4382,8 +4382,8 @@ async def _exchange_for_bridge_server(server, upstream_body, user_id): return_value=fake_http_client, ), patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_user_id_from_request", - new=AsyncMock(return_value=user_id), + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request", + new=AsyncMock(return_value=key_hash), ), patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), ): @@ -4416,7 +4416,7 @@ async def test_oauth_delegate_bridge_token_exchange_mints_envelope_not_raw_token server = _bridge_server(auth_type=MCPAuth.oauth_delegate) upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} - response = await _exchange_for_bridge_server(server, upstream, user_id="user-77") + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") body = json.loads(response.body) token = body["access_token"] @@ -4429,7 +4429,7 @@ async def test_oauth_delegate_bridge_token_exchange_mints_envelope_not_raw_token keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY) opened = resolve_bridge_envelope(token, keys, datetime.now(timezone.utc), server.server_id) assert isinstance(opened, BridgeEnvelopeAdmitted) - assert opened.identity.user_id == "user-77" + assert opened.identity.key_hash == "hashed-litellm-key-77" assert opened.upstream_authorization.get_secret_value() == "Bearer UPSTREAM-SECRET-TOKEN" @@ -4443,7 +4443,7 @@ async def test_oauth_delegate_bridge_token_exchange_fails_closed_without_litellm upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} with pytest.raises(HTTPException) as exc: - await _exchange_for_bridge_server(server, upstream, user_id=None) + await _exchange_for_bridge_server(server, upstream, key_hash=None) assert exc.value.status_code == 400 assert exc.value.detail["error"] == "invalid_request" @@ -4457,7 +4457,7 @@ async def test_true_passthrough_bridge_token_exchange_returns_raw_upstream_token server = _bridge_server(auth_type=MCPAuth.true_passthrough) upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} - response = await _exchange_for_bridge_server(server, upstream, user_id="user-77") + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") body = json.loads(response.body) assert body["access_token"] == "UPSTREAM-SECRET-TOKEN" @@ -4472,7 +4472,7 @@ async def test_non_bridge_oauth_delegate_token_exchange_returns_raw_upstream_tok server = _bridge_server(auth_type=MCPAuth.oauth_delegate, dcr_bridge=None) upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} - response = await _exchange_for_bridge_server(server, upstream, user_id="user-77") + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") body = json.loads(response.body) assert body["access_token"] == "UPSTREAM-SECRET-TOKEN" @@ -4822,6 +4822,68 @@ async def test_extract_user_id_rejects_expired_key(proxy_globals): assert await _extract_user_id_from_request(request) is None +@pytest.mark.asyncio +async def test_extract_active_key_hash_returns_hash_for_active_key(proxy_globals): + """The dcr_bridge mint seals the hash of the authorizing key so admission can reload the live + record. For an active key the resolver returns exactly hash_token(key), the same value + get_key_object and the whole cache/DB layer key the record by, so the sealed reference resolves + back to this key at admission.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _extract_active_key_hash_from_request, + ) + from litellm.proxy._types import UserAPIKeyAuth, hash_token + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + key = "sk-alice-key" + cache = UserApiKeyCache() + await cache.async_set_cache( + hash_token(key), + UserAPIKeyAuth(token=hash_token(key), user_id="alice"), + model_type=UserAPIKeyAuth, + ) + proxy_globals.user_api_key_cache = cache + proxy_globals.prisma_client = object() + + request = _token_request({"x-litellm-api-key": f"Bearer {key}"}) + assert await _extract_active_key_hash_from_request(request) == hash_token(key) + + +@pytest.mark.asyncio +async def test_extract_active_key_hash_rejects_blocked_key(proxy_globals): + """A blocked key must not yield a hash, so no gateway-bound envelope is minted for a revoked key; + the mint fails closed with invalid_request instead.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _extract_active_key_hash_from_request, + ) + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + class _FakePrisma: + async def get_data(self, token, table_name, parent_otel_span=None, proxy_logging_obj=None): + return UserAPIKeyAuth(token=token, user_id="blocked-user", blocked=True) + + proxy_globals.user_api_key_cache = UserApiKeyCache() + proxy_globals.prisma_client = _FakePrisma() + + request = _token_request({"x-litellm-api-key": "sk-blocked-key"}) + assert await _extract_active_key_hash_from_request(request) is None + + +@pytest.mark.asyncio +async def test_extract_active_key_hash_none_without_litellm_key(proxy_globals): + """No LiteLLM key on the request yields no hash without consulting the resolver.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _extract_active_key_hash_from_request, + ) + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + proxy_globals.user_api_key_cache = UserApiKeyCache() + proxy_globals.prisma_client = object() + + request = _token_request({"content-type": "application/json"}) + assert await _extract_active_key_hash_from_request(request) is None + + @pytest.mark.asyncio async def test_token_endpoint_uses_client_secret_basic_when_configured(): """LIT-4091: a server with token_endpoint_auth_method=client_secret_basic must send the From 7df848aa6c6efe2fd32ea9174cd2e5bc51ed479f Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 14:41:44 -0700 Subject: [PATCH 038/123] fix(mcp): return 502 not KeyError when a bridge upstream response lacks access_token The eager access_token = token_response["access_token"] extraction ran before the dcr_bridge branch, so a missing upstream access_token raised an unhandled KeyError and _bridge_grant_from_token_response's nil guard (which maps to a clean 502) was dead code. Move the extraction onto the non-bridge result path so the bridge branch reaches its 502 guard. --- .../mcp_server/discoverable_endpoints.py | 3 +-- .../mcp_server/test_discoverable_endpoints.py | 17 +++++++++++++++++ 2 files changed, 18 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index eeadeb290f3..ef42fef312d 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -865,7 +865,6 @@ async def exchange_token_with_server( ) raise token_response = response.json() - access_token = token_response["access_token"] # Validate token response against server-configured rules before any storage. # This rejects tokens from wrong Slack workspaces, Atlassian orgs, etc. @@ -912,7 +911,7 @@ async def exchange_token_with_server( return await _mint_bridge_delegate_token_response(request, mcp_server, token_response) result = { - "access_token": access_token, + "access_token": token_response["access_token"], "token_type": token_response.get("token_type", "Bearer"), } diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index e69d94d9615..5826d1f28b6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4449,6 +4449,23 @@ async def test_oauth_delegate_bridge_token_exchange_fails_closed_without_litellm assert exc.value.detail["error"] == "invalid_request" +@pytest.mark.asyncio +async def test_oauth_delegate_bridge_token_exchange_missing_access_token_is_502_not_keyerror(): + """When the upstream token response has no access_token, a dcr_bridge oauth_delegate exchange + returns a clean 502 rather than raising a KeyError. The eager access_token extraction used to run + before the bridge branch, so a missing token raised KeyError and _bridge_grant_from_token_response's + nil guard (which maps to 502) was dead code; the extraction now lives on the non-bridge path only.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"token_type": "Bearer", "expires_in": 3600} + + with pytest.raises(HTTPException) as exc: + await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + + assert exc.value.status_code == 502 + + @pytest.mark.asyncio async def test_true_passthrough_bridge_token_exchange_returns_raw_upstream_token(): """Only oauth_delegate mints. A true_passthrough dcr_bridge server relays the raw upstream token From 362c78e30864f5826a3dc826dc5428c84a7d4f5e Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 15:05:30 -0700 Subject: [PATCH 039/123] fix(mcp): let a keyless-user active key mint a bridge envelope _resolve_active_litellm_key gated on _active_key_user_id, which returns None both for blocked/expired keys AND for valid keys with no user_id, so a team-scoped or service-account key was wrongly rejected with invalid_request at bridge token exchange. Split the active-state gate (_key_is_active: blocked/expiry only) from the user_id extraction; the mint seals the key hash, not the user, and admission already handles a keyless-user key. The per-user token store still gets no user for such a key, as there is none to key a stored credential by. --- .../mcp_server/discoverable_endpoints.py | 44 ++++++++++++------- .../mcp_server/test_discoverable_endpoints.py | 29 ++++++++++++ 2 files changed, 57 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ef42fef312d..df1615aff99 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -329,26 +329,37 @@ def _litellm_key_from_request(request: Request) -> Optional[str]: return None -def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> Optional[str]: - """The key's ``user_id``, or ``None`` if the key is blocked or expired. +def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool: + """``True`` when the presented key is neither blocked nor past its expiry. - The OAuth token endpoint is unauthenticated, so the presented key is validated here before its - identity is trusted to key a stored credential; a revoked or expired key must not be able to - write or overwrite the per-user OAuth token. ``get_key_object`` resolves a row without these - checks (the main ``user_api_key_auth`` pipeline enforces them downstream, which this endpoint - bypasses), so they are applied here. Deleted keys are already rejected upstream, where - ``get_key_object`` raises on a row that no longer exists. + The OAuth token endpoint is unauthenticated, so the presented key is validated here before it is + trusted; a revoked or expired key must not mint a bridge envelope or write a stored credential. + ``get_key_object`` resolves a row without these checks (the main ``user_api_key_auth`` pipeline + enforces them downstream, which this endpoint bypasses), so they are applied here. Deleted keys + are already rejected upstream, where ``get_key_object`` raises on a row that no longer exists. + + This is an active-state gate only; it deliberately does not require a ``user_id``. A valid + team-scoped or service-account key has no ``user_id`` yet is a legitimate credential, so gating + on ``user_id`` presence would wrongly reject it. Callers that need the user (the per-user token + store) derive it separately via :func:`_active_key_user_id`. """ if key_obj.blocked is True: - return None + return False expires = key_obj.expires if expires is not None: expiry = expires if isinstance(expires, datetime) else datetime.fromisoformat(expires) if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None: expiry = expiry.replace(tzinfo=timezone.utc) if expiry < datetime.now(timezone.utc): - return None - return key_obj.user_id + return False + return True + + +def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> Optional[str]: + """The active key's ``user_id``, or ``None`` when the key is blocked/expired or simply has no + ``user_id`` (a team-scoped or service-account key). Used only by the per-user token store, which + needs a user to key the stored credential; the bridge mint uses the key hash and does not.""" + return key_obj.user_id if _key_is_active(key_obj) else None async def _resolve_active_litellm_key(request: Request) -> Optional[Tuple[str, "UserAPIKeyAuth"]]: @@ -361,10 +372,11 @@ async def _resolve_active_litellm_key(request: Request) -> Optional[Tuple[str, " cross-replica Redis hit deserializes to a plain ``dict`` rather than a ``UserAPIKeyAuth``; the previous code read only ``Authorization`` and did ``getattr(cached, "user_id")`` with no ``model_type`` rehydration and no DB fallback, so it silently returned ``None``. The resolved key - is validated (``_active_key_user_id``) before it is trusted, so a blocked or expired key resolves - to ``None``. The returned hash is the value ``get_key_object`` and the cache/DB layer key the - record by. Callers derive the ``user_id`` (per-user token store) or seal the hash (dcr_bridge - envelope) from the result. + is validated (``_key_is_active``) before it is trusted, so a blocked or expired key resolves to + ``None``, while a valid team-scoped or service-account key (no ``user_id``) still resolves so it + can mint a bridge envelope. The returned hash is the value ``get_key_object`` and the cache/DB + layer key the record by. Callers derive the ``user_id`` (per-user token store) or seal the hash + (dcr_bridge envelope) from the result. """ token = _litellm_key_from_request(request) if not token: @@ -393,7 +405,7 @@ async def _resolve_active_litellm_key(request: Request) -> Optional[Tuple[str, " type(exc).__name__, ) return None - if _active_key_user_id(key_obj) is None: + if not _key_is_active(key_obj): return None return key_hash, key_obj diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 5826d1f28b6..2bf5e49ec79 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4865,6 +4865,35 @@ async def test_extract_active_key_hash_returns_hash_for_active_key(proxy_globals assert await _extract_active_key_hash_from_request(request) == hash_token(key) +@pytest.mark.asyncio +async def test_extract_active_key_hash_returns_hash_for_active_key_without_user_id(proxy_globals): + """A valid team-scoped or service-account key has no user_id but is a legitimate credential, so it + must still resolve to a hash and be able to mint a bridge envelope. Gating the resolver on user_id + presence wrongly rejected these keys with invalid_request; the active-state gate now checks only + blocked and expiry, and the key hash (not the user) is what the mint seals. The per-user token + store still gets no user for such a key, since there is none to key a stored credential by.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _extract_active_key_hash_from_request, + _extract_user_id_from_request, + ) + from litellm.proxy._types import UserAPIKeyAuth, hash_token + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + key = "sk-team-scoped-key" + cache = UserApiKeyCache() + await cache.async_set_cache( + hash_token(key), + UserAPIKeyAuth(token=hash_token(key), user_id=None, team_id="team-x"), + model_type=UserAPIKeyAuth, + ) + proxy_globals.user_api_key_cache = cache + proxy_globals.prisma_client = object() + + request = _token_request({"x-litellm-api-key": f"Bearer {key}"}) + assert await _extract_active_key_hash_from_request(request) == hash_token(key) + assert await _extract_user_id_from_request(request) is None + + @pytest.mark.asyncio async def test_extract_active_key_hash_rejects_blocked_key(proxy_globals): """A blocked key must not yield a hash, so no gateway-bound envelope is minted for a revoked key; From ceff2d1f3c99f2e314e497e86c302372ba26e64b Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 15:47:27 -0700 Subject: [PATCH 040/123] style(mcp): use X | None annotations on the touched key-resolution helpers The keyless-user fix moved these signatures, so their pre-existing Optional[...] annotations counted against the diff and tripped the UP045 strict-budget gate. Modernize the four touched return annotations to the X | None form the gate wants; runtime behavior is unchanged. --- .../_experimental/mcp_server/discoverable_endpoints.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index df1615aff99..c5d7c3b40fd 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -355,14 +355,14 @@ def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool: return True -def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> Optional[str]: +def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None: """The active key's ``user_id``, or ``None`` when the key is blocked/expired or simply has no ``user_id`` (a team-scoped or service-account key). Used only by the per-user token store, which needs a user to key the stored credential; the bridge mint uses the key hash and does not.""" return key_obj.user_id if _key_is_active(key_obj) else None -async def _resolve_active_litellm_key(request: Request) -> Optional[Tuple[str, "UserAPIKeyAuth"]]: +async def _resolve_active_litellm_key(request: Request) -> Tuple[str, "UserAPIKeyAuth"] | None: """Resolve the presented litellm key to ``(its hash, the live active key record)``, or ``None`` when the key is absent, unresolvable, or blocked/expired. @@ -410,7 +410,7 @@ async def _resolve_active_litellm_key(request: Request) -> Optional[Tuple[str, " return key_hash, key_obj -async def _extract_user_id_from_request(request: Request) -> Optional[str]: +async def _extract_user_id_from_request(request: Request) -> str | None: """The LiteLLM ``user_id`` for the token request, so a per-user token is stored under the same identity the egress later reads it by (``user_api_key_auth.user_id``). ``None`` when no active key is present. See :func:`_resolve_active_litellm_key` for the resolution and active-key gate. @@ -419,7 +419,7 @@ async def _extract_user_id_from_request(request: Request) -> Optional[str]: return _active_key_user_id(resolved[1]) if resolved else None -async def _extract_active_key_hash_from_request(request: Request) -> Optional[str]: +async def _extract_active_key_hash_from_request(request: Request) -> str | None: """The hash of the litellm key that authorized the token request, when it maps to an active key. A DCR-bridge envelope seals this hash so admission can reload the live ``UserAPIKeyAuth`` record From 2f349f6cd18194a67b7c5985e989de5009bec4c9 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 16:16:52 -0700 Subject: [PATCH 041/123] fix(mcp): coerce numeric expires_in and make the active-key check total Two correctness gaps in the bridge mint. _bridge_grant_from_token_response only accepted an int expires_in, dropping a float (3600.0) or numeric-string ('3600') lifetime to None so the envelope fell back to its 1h cap and could outlive a shorter-lived upstream token; coerce it to a positive int (bool excluded). And _key_is_active called datetime.fromisoformat on the str|datetime expires outside the resolver's try, so a malformed stored expiry raised an unhandled 500 instead of the fail-closed invalid_request; it now fails closed (inactive) on an unparseable expiry. Regression tests cover int/float/string/bool coercion, the short-float TTL, and the malformed-expiry fail-closed path. --- .../mcp_server/discoverable_endpoints.py | 39 ++++++++++-- .../mcp_server/test_discoverable_endpoints.py | 60 +++++++++++++++++++ 2 files changed, 95 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index c5d7c3b40fd..00586a6afb3 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -342,12 +342,23 @@ def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool: team-scoped or service-account key has no ``user_id`` yet is a legitimate credential, so gating on ``user_id`` presence would wrongly reject it. Callers that need the user (the per-user token store) derive it separately via :func:`_active_key_user_id`. + + Total by design: ``expires`` is typed ``str | datetime``, and an unparseable string would make + ``datetime.fromisoformat`` raise. Since the callers run this outside their key-resolution + ``try``, an uncaught parse error would surface as a 500 instead of the endpoint's fail-closed + behavior, so a malformed expiry is treated as inactive (return ``False``) rather than raising. """ if key_obj.blocked is True: return False expires = key_obj.expires if expires is not None: - expiry = expires if isinstance(expires, datetime) else datetime.fromisoformat(expires) + if isinstance(expires, datetime): + expiry = expires + else: + try: + expiry = datetime.fromisoformat(expires) + except (ValueError, TypeError): + return False if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None: expiry = expiry.replace(tzinfo=timezone.utc) if expiry < datetime.now(timezone.utc): @@ -698,10 +709,31 @@ async def authorize_with_server( return response +def _coerce_positive_expires_in(value: object) -> int | None: + """Coerce an upstream ``expires_in`` to a positive int, or ``None`` when it is absent or not a + usable number. IdPs return it as an int, a float (``3600.0``), or a numeric string (``"3600"``); + accepting only ``int`` would drop the float/string cases to ``None`` and fall back to the + envelope's 1h cap, which can outlive a shorter-lived upstream token and forward a stale bearer. + ``bool`` is excluded (it is an ``int`` subclass but never a real lifetime).""" + if isinstance(value, bool): + return None + if isinstance(value, (int, float)): + seconds = int(value) + return seconds if seconds > 0 else None + if isinstance(value, str): + try: + seconds = int(float(value.strip())) + except (ValueError, TypeError): + return None + return seconds if seconds > 0 else None + return None + + def _bridge_grant_from_token_response(token_response: object) -> Optional["UpstreamTokenGrant"]: """Validate an upstream OAuth token response into a typed grant, or None when it lacks a usable access token. Each field is isinstance-checked so nothing untyped from ``response.json()`` flows - into the grant.""" + into the grant; ``expires_in`` is numerically coerced so a float/string lifetime is honored + rather than dropped to the envelope's default cap.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import UpstreamTokenGrant, ) @@ -714,13 +746,12 @@ def _bridge_grant_from_token_response(token_response: object) -> Optional["Upstr token_type = token_response.get("token_type") refresh = token_response.get("refresh_token") scope = token_response.get("scope") - expires_in = token_response.get("expires_in") return UpstreamTokenGrant( access_token=SecretStr(access), token_type=token_type if isinstance(token_type, str) and token_type else "Bearer", refresh_token=SecretStr(refresh) if isinstance(refresh, str) and refresh else None, scope=scope if isinstance(scope, str) and scope else None, - expires_in=expires_in if isinstance(expires_in, int) and expires_in > 0 else None, + expires_in=_coerce_positive_expires_in(token_response.get("expires_in")), ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 2bf5e49ec79..0dc8d8116df 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4449,6 +4449,43 @@ async def test_oauth_delegate_bridge_token_exchange_fails_closed_without_litellm assert exc.value.detail["error"] == "invalid_request" +def test_bridge_grant_coerces_numeric_expires_in(): + """expires_in from an IdP may be an int, a float (3600.0), or a numeric string ("3600"); coerce + it to a positive int so the envelope TTL honors the real lifetime instead of dropping a non-int + value and defaulting to the 1h cap (which can outlive a shorter-lived upstream token). bool and + non-numeric values become None.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _bridge_grant_from_token_response, + ) + + def ei(v): + return _bridge_grant_from_token_response({"access_token": "x", "expires_in": v}).expires_in + + assert ei(300) == 300 + assert ei(300.0) == 300 + assert ei("300") == 300 + assert ei(" 300 ") == 300 + assert ei(True) is None + assert ei("nope") is None + assert ei(0) is None + assert ei(-5) is None + assert ei(None) is None + + +@pytest.mark.asyncio +async def test_bridge_token_exchange_honors_short_float_expires_in_ttl(): + """A short float expires_in from the upstream caps the envelope TTL, so the client-held envelope + does not outlive the upstream token. Before coercion a float was dropped and the envelope + defaulted to the 1h cap (3600), which would forward a stale bearer after the upstream token + expired.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "UP", "token_type": "Bearer", "expires_in": 120.0} + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + assert json.loads(response.body)["expires_in"] <= 120 + + @pytest.mark.asyncio async def test_oauth_delegate_bridge_token_exchange_missing_access_token_is_502_not_keyerror(): """When the upstream token response has no access_token, a dcr_bridge oauth_delegate exchange @@ -4915,6 +4952,29 @@ async def test_extract_active_key_hash_rejects_blocked_key(proxy_globals): assert await _extract_active_key_hash_from_request(request) is None +@pytest.mark.asyncio +async def test_extract_active_key_hash_fails_closed_on_malformed_expiry(proxy_globals): + """A key whose stored expires string does not parse must fail closed to no-hash (the mint then + returns invalid_request), not surface an unhandled 500. The active-state check runs outside the + resolver's try, so it must be total over a bad expires rather than letting datetime.fromisoformat + raise. Before the fix this raised a ValueError instead of returning None.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _extract_active_key_hash_from_request, + ) + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + class _FakePrisma: + async def get_data(self, token, table_name, parent_otel_span=None, proxy_logging_obj=None): + return UserAPIKeyAuth(token=token, user_id="u", expires="not-a-parseable-date") + + proxy_globals.user_api_key_cache = UserApiKeyCache() + proxy_globals.prisma_client = _FakePrisma() + + request = _token_request({"x-litellm-api-key": "sk-bad-expiry-key"}) + assert await _extract_active_key_hash_from_request(request) is None + + @pytest.mark.asyncio async def test_extract_active_key_hash_none_without_litellm_key(proxy_globals): """No LiteLLM key on the request yields no hash without consulting the resolver.""" From 7a63e516252a5fc78a6150da4b3fc6cd03bab1c6 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 16:36:08 -0700 Subject: [PATCH 042/123] fix(mcp): harden the bridge token mint (multi-lens review pass) Findings from a full adversarial review of the mint path across security, correctness, error-handling, concurrency, and OAuth-protocol dimensions. - expires_in coercion is now total: int(float(...)) can raise OverflowError on Infinity / a giant numeric string, which escaped the ValueError/TypeError catch and 500'd the token endpoint. Unified to catch OverflowError too. - Resolve the litellm identity BEFORE exchanging the single-use upstream code, so a missing or transiently-unresolvable identity fails closed with invalid_request without burning the code (the mint re-resolves via a cache hit). - The no-identity failure is now an RFC 6749 5.2-shaped invalid_request (JSONResponse, top-level error, no-store) instead of a detail-wrapped HTTPException, matching the BYOK OAuth endpoint. - EnvelopeTooLarge (upstream token too big to seal) surfaces a 502, not a 500. - The upstream refresh_token is no longer sealed into the envelope: the edge never consumes it, so it was dead weight embedding a long-lived upstream credential in the client bearer and enlarging the envelope; refresh is a follow-up (a dedicated refresh-envelope). Security review found no exploitable defect (forgery, cross-server/user replay, leakage, confused-deputy all closed). Regression tests cover the OverflowError, the code-not-burned path, the RFC-shaped error, the 502, and the dropped refresh. --- .../mcp_server/discoverable_endpoints.py | 76 ++++++++++++------- .../mcp_server/test_discoverable_endpoints.py | 71 +++++++++++++++-- 2 files changed, 116 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 00586a6afb3..4f72c2cb8c2 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -373,7 +373,7 @@ def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None: return key_obj.user_id if _key_is_active(key_obj) else None -async def _resolve_active_litellm_key(request: Request) -> Tuple[str, "UserAPIKeyAuth"] | None: +async def _resolve_active_litellm_key(request: Request) -> tuple[str, "UserAPIKeyAuth"] | None: """Resolve the presented litellm key to ``(its hash, the live active key record)``, or ``None`` when the key is absent, unresolvable, or blocked/expired. @@ -714,19 +714,17 @@ def _coerce_positive_expires_in(value: object) -> int | None: usable number. IdPs return it as an int, a float (``3600.0``), or a numeric string (``"3600"``); accepting only ``int`` would drop the float/string cases to ``None`` and fall back to the envelope's 1h cap, which can outlive a shorter-lived upstream token and forward a stale bearer. - ``bool`` is excluded (it is an ``int`` subclass but never a real lifetime).""" - if isinstance(value, bool): + ``bool`` is excluded (it is an ``int`` subclass but never a real lifetime). Total over hostile + input: a non-numeric string, ``NaN``, ``Infinity``, or an over-large value all resolve to + ``None`` rather than raising (``int(float(...))`` can raise ``ValueError`` or ``OverflowError``), + so a malformed upstream ``expires_in`` never surfaces as a 500 from the token endpoint.""" + if isinstance(value, bool) or not isinstance(value, (int, float, str)): return None - if isinstance(value, (int, float)): - seconds = int(value) - return seconds if seconds > 0 else None - if isinstance(value, str): - try: - seconds = int(float(value.strip())) - except (ValueError, TypeError): - return None - return seconds if seconds > 0 else None - return None + try: + seconds = int(float(value)) + except (ValueError, TypeError, OverflowError): + return None + return seconds if seconds > 0 else None def _bridge_grant_from_token_response(token_response: object) -> Optional["UpstreamTokenGrant"]: @@ -744,17 +742,38 @@ def _bridge_grant_from_token_response(token_response: object) -> Optional["Upstr if not isinstance(access, str) or not access: return None token_type = token_response.get("token_type") - refresh = token_response.get("refresh_token") scope = token_response.get("scope") return UpstreamTokenGrant( access_token=SecretStr(access), token_type=token_type if isinstance(token_type, str) and token_type else "Bearer", - refresh_token=SecretStr(refresh) if isinstance(refresh, str) and refresh else None, + # The upstream refresh_token is deliberately NOT sealed: the edge never consumes it (it forwards + # only token_type + access_token), so it would be dead weight embedding a long-lived upstream + # credential in the client-held bearer, and it enlarges the envelope. Refresh support is a + # follow-up (a dedicated refresh-envelope); the client re-runs authorization_code at the cap. + refresh_token=None, scope=scope if isinstance(scope, str) and scope else None, expires_in=_coerce_positive_expires_in(token_response.get("expires_in")), ) +def _bridge_invalid_request_response() -> JSONResponse: + """RFC 6749 §5.2-shaped ``invalid_request`` for a bridge token exchange that carries no resolvable + litellm identity. Returned (not raised) so the OAuth error members sit at the top level rather than + wrapped in FastAPI's ``detail``, with the no-store token-endpoint headers, matching the BYOK OAuth + endpoint and what a strict DCR client parses per RFC 6749 §5.2.""" + return JSONResponse( + status_code=400, + content={ + "error": "invalid_request", + "error_description": ( + "this server issues a gateway-bound credential; send a litellm credential " + "(x-litellm-api-key or Authorization) on the token request" + ), + }, + headers=TOKEN_NO_CACHE_HEADERS, + ) + + async def _mint_bridge_delegate_token_response( request: Request, mcp_server: MCPServer, token_response: object ) -> JSONResponse: @@ -784,16 +803,7 @@ async def _mint_bridge_delegate_token_response( key_hash = await _extract_active_key_hash_from_request(request) if not key_hash: - raise HTTPException( - status_code=400, - detail={ - "error": "invalid_request", - "error_description": ( - "this server issues a gateway-bound credential; send a litellm credential " - "(x-litellm-api-key or Authorization) on the token request" - ), - }, - ) + return _bridge_invalid_request_response() grant = _bridge_grant_from_token_response(token_response) if grant is None: @@ -804,7 +814,11 @@ async def _mint_bridge_delegate_token_response( identity = EnvelopeIdentity(server_id=mcp_server.server_id, key_hash=key_hash) sealed = build_bridge_token_response(identity, grant, keys, now) if not isinstance(sealed, SealedEnvelope): - raise HTTPException(status_code=500, detail="Failed to mint the gateway-bound credential") + # build_bridge_token_response returns EnvelopeTooLarge as a value when the upstream token is + # too large to seal; that is an upstream-payload condition, so surface a 502, not a 500. + raise HTTPException( + status_code=502, detail="Upstream token is too large to seal into a gateway-bound credential" + ) expires_in = max(1, int((sealed.expires_at - now).total_seconds())) body = {"access_token": sealed.token.get_secret_value(), "token_type": "Bearer", "expires_in": expires_in} @@ -884,6 +898,16 @@ async def exchange_token_with_server( if code_verifier: token_data["code_verifier"] = code_verifier + # For a bridge oauth_delegate mint, resolve the litellm identity BEFORE exchanging the + # single-use upstream code. A missing or transiently-unresolvable identity then fails closed + # with invalid_request without consuming the code, so the client can retry the same code + # instead of being forced back through the full interactive authorize. The mint below + # re-resolves authoritatively; get_key_object is cache-first, so that second call is a cache + # hit and this adds no extra database round-trip. + if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: + if not await _extract_active_key_hash_from_request(request): + return _bridge_invalid_request_response() + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) response = await async_client.post( mcp_server.token_url, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 0dc8d8116df..8aad8b24e88 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4365,7 +4365,7 @@ async def test_register_bridge_relay_never_persists(): _BRIDGE_MASTER_KEY = "sk-bridge-producer-master-key-0123456789abcdef" -async def _exchange_for_bridge_server(server, upstream_body, key_hash): +async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_client_out=None): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( exchange_token_with_server, ) @@ -4375,6 +4375,8 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash): fake_http_response.raise_for_status = MagicMock() fake_http_client = MagicMock() fake_http_client.post = AsyncMock(return_value=fake_http_response) + if fake_client_out is not None: + fake_client_out["client"] = fake_http_client with ( patch( @@ -4436,17 +4438,69 @@ async def test_oauth_delegate_bridge_token_exchange_mints_envelope_not_raw_token @pytest.mark.asyncio async def test_oauth_delegate_bridge_token_exchange_fails_closed_without_litellm_identity(): """Without a resolvable litellm identity on the token request, the exchange must not mint an - identity-less envelope; it returns an OAuth invalid_request so the client sends a credential.""" + identity-less envelope. It returns an RFC 6749 §5.2-shaped invalid_request (error at the top + level, not wrapped in detail) BEFORE exchanging the upstream code, so the single-use code is not + burned and the client can retry.""" from litellm.types.mcp import MCPAuth server = _bridge_server(auth_type=MCPAuth.oauth_delegate) upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} + captured: dict = {} + response = await _exchange_for_bridge_server(server, upstream, key_hash=None, fake_client_out=captured) + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_request" + # identity resolution failed first, so the upstream single-use code was never exchanged (not burned) + captured["client"].post.assert_not_called() + + +@pytest.mark.asyncio +async def test_bridge_envelope_too_large_upstream_token_is_502(): + """An upstream token too large to seal into the envelope is an upstream-payload condition, so the + mint surfaces a 502 rather than a 500 (build_bridge_token_response returns EnvelopeTooLarge as a + value, and the caller maps it to a truthful status).""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "x" * 40000, "token_type": "Bearer", "expires_in": 3600} with pytest.raises(HTTPException) as exc: - await _exchange_for_bridge_server(server, upstream, key_hash=None) + await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + assert exc.value.status_code == 502 - assert exc.value.status_code == 400 - assert exc.value.detail["error"] == "invalid_request" + +@pytest.mark.asyncio +async def test_bridge_envelope_does_not_seal_upstream_refresh_token(): + """The upstream refresh_token is never sealed into the client-held envelope: the edge never + consumes it and a long-lived upstream credential should not live in the client bearer. The opened + envelope's grant carries no refresh token even when the upstream returned one, and neither does + the response body.""" + from datetime import datetime, timezone + + from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( + envelope_keys_from_master_key, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( + OpenedEnvelope, + open_envelope, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = { + "access_token": "UP", + "token_type": "Bearer", + "expires_in": 3600, + "refresh_token": "UPSTREAM-REFRESH", + } + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + + body = json.loads(response.body) + assert "refresh_token" not in body + assert "UPSTREAM-REFRESH" not in body["access_token"] + keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY) + opened = open_envelope(body["access_token"], keys, datetime.now(timezone.utc)) + assert isinstance(opened, OpenedEnvelope) + assert opened.grant.refresh_token is None def test_bridge_grant_coerces_numeric_expires_in(): @@ -4470,6 +4524,13 @@ def test_bridge_grant_coerces_numeric_expires_in(): assert ei(0) is None assert ei(-5) is None assert ei(None) is None + # hostile numerics must not raise (int(float(...)) can OverflowError) -> None + assert ei("inf") is None + assert ei("1e999") is None + assert ei("-inf") is None + assert ei("nan") is None + assert ei(float("inf")) is None + assert ei(10**400) is None @pytest.mark.asyncio From e16ad044c3773ac958b1cffff0ad6d15bb5e0296 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 16:58:14 -0700 Subject: [PATCH 043/123] fix(mcp): close the burn-before-check gate for both grants and validate master_key first Follow-up to the pre-exchange identity gate, which I had only added to the authorization_code branch and which left the master_key check inside the mint (after the upstream exchange) - so the very burn-then-fail pattern it was meant to prevent still applied to refresh_token grants and to a misconfigured gateway. - Hoist a single pre-exchange gate above the upstream call that covers BOTH grant types: it fails closed (invalid_request) on an unresolvable litellm identity and 500s on an unset master_key BEFORE the single-use code or refresh token is exchanged/rotated, so a bad key or a misconfigured gateway never burns the upstream credential. - Report expires_in from the envelope JWT's own second-truncated exp (rounding the elapsed portion up) instead of the raw expires_at - now delta, so the client is never told the bearer is valid past the ~1s point admission already expires it. Regression tests assert the upstream exchange is never called on the no-identity refresh grant and the master_key-unset path, and that the reported expires_in does not overstate the JWT exp. --- .../mcp_server/discoverable_endpoints.py | 28 ++++-- .../mcp_server/test_discoverable_endpoints.py | 97 +++++++++++++++++++ 2 files changed, 115 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 4f72c2cb8c2..5bf4de09bba 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1,6 +1,7 @@ import asyncio import html as _html import json +import math import secrets import time from datetime import datetime, timezone @@ -820,7 +821,10 @@ async def _mint_bridge_delegate_token_response( status_code=502, detail="Upstream token is too large to seal into a gateway-bound credential" ) - expires_in = max(1, int((sealed.expires_at - now).total_seconds())) + # The JWT exp is int(expires_at.timestamp()) (second-truncated), and admission expires the envelope + # against that exp. Report expires_in from the same truncated exp, rounding the elapsed portion up, + # so the client is never told the bearer lives past the point admission already rejects it. + expires_in = max(1, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp())) body = {"access_token": sealed.token.get_secret_value(), "token_type": "Bearer", "expires_in": expires_in} return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS) @@ -898,15 +902,19 @@ async def exchange_token_with_server( if code_verifier: token_data["code_verifier"] = code_verifier - # For a bridge oauth_delegate mint, resolve the litellm identity BEFORE exchanging the - # single-use upstream code. A missing or transiently-unresolvable identity then fails closed - # with invalid_request without consuming the code, so the client can retry the same code - # instead of being forced back through the full interactive authorize. The mint below - # re-resolves authoritatively; get_key_object is cache-first, so that second call is a cache - # hit and this adds no extra database round-trip. - if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: - if not await _extract_active_key_hash_from_request(request): - return _bridge_invalid_request_response() + # A bridge oauth_delegate mint must fail closed BEFORE the upstream exchange consumes or rotates the + # single-use code (or refresh token): confirm the gateway can mint at all (master_key set) and that + # the request carries a resolvable litellm identity. Applies to both grant types, so an invalid key + # or a misconfigured gateway never burns the upstream credential. The mint below re-checks + # authoritatively; get_key_object is cache-first, so the identity re-resolution is a cache hit and + # adds no extra database round-trip. + if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: + from litellm.proxy.proxy_server import master_key as _bridge_master_key # noqa: PLC0415 + + if not _bridge_master_key: + raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set") + if not await _extract_active_key_hash_from_request(request): + return _bridge_invalid_request_response() async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) response = await async_client.post( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 8aad8b24e88..4a4ff398915 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4503,6 +4503,103 @@ async def test_bridge_envelope_does_not_seal_upstream_refresh_token(): assert opened.grant.refresh_token is None +@pytest.mark.asyncio +async def test_bridge_refresh_grant_fails_closed_before_upstream_when_no_identity(): + """The pre-exchange identity gate covers the refresh_token grant, not just authorization_code: an + unresolvable litellm identity fails closed with invalid_request BEFORE the upstream refresh is + exchanged, so the client's refresh token is not rotated/consumed on a rejected request.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + fake_http_client = MagicMock() + fake_http_client.post = AsyncMock() + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=fake_http_client, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request", + new=AsyncMock(return_value=None), + ), + patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), + ): + response = await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=server, + grant_type="refresh_token", + code=None, + redirect_uri=None, + client_id="dcr-client-123", + client_secret=None, + code_verifier=None, + refresh_token="client-refresh-token", + ) + + assert response.status_code == 400 + assert json.loads(response.body)["error"] == "invalid_request" + fake_http_client.post.assert_not_called() + + +@pytest.mark.asyncio +async def test_bridge_mint_fails_closed_before_upstream_when_master_key_unset(): + """master_key is validated BEFORE the upstream exchange, so a misconfigured gateway 500s without + consuming the single-use code, avoiding the burn-then-fail the pre-exchange gate exists to prevent.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + fake_http_client = MagicMock() + fake_http_client.post = AsyncMock() + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=fake_http_client, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request", + new=AsyncMock(return_value="hashed-litellm-key-77"), + ), + patch("litellm.proxy.proxy_server.master_key", None), + ): + with pytest.raises(HTTPException) as exc: + await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="https://claude.ai/api/mcp/auth_callback", + client_id="dcr-client-123", + client_secret=None, + code_verifier="verifier", + ) + + assert exc.value.status_code == 500 + fake_http_client.post.assert_not_called() + + +@pytest.mark.asyncio +async def test_bridge_reported_expires_in_does_not_overstate_jwt_exp(): + """The reported expires_in is derived from the envelope JWT's second-truncated exp (rounding the + elapsed portion up), so the client is never told the bearer lives past the point admission expires + it. Regression for the sub-second overstatement of the raw (expires_at - now) delta.""" + import time + + import jwt as _jwt + + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "UP", "token_type": "Bearer", "expires_in": 300} + before = int(time.time()) + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + body = json.loads(response.body) + claims = _jwt.decode(body["access_token"].removeprefix("llm_env_"), options={"verify_signature": False}) + # projecting the reported lifetime from a time no later than the mint must not exceed the JWT exp + assert before + body["expires_in"] <= claims["exp"] + + def test_bridge_grant_coerces_numeric_expires_in(): """expires_in from an IdP may be an int, a float (3600.0), or a numeric string ("3600"); coerce it to a positive int so the envelope TTL honors the real lifetime instead of dropping a non-int From 2f0ddc82f7f38b282dff5e299436cd99285e1ba7 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 17:22:01 -0700 Subject: [PATCH 044/123] refactor(mcp): make the bridge delegate mint a phased failures-as-values pipeline The dcr_bridge oauth_delegate token mint validated its preconditions in two places: a pre-exchange guard inside exchange_token_with_server (master_key set, resolvable litellm identity) and an authoritative re-check inside the post-exchange _mint_bridge_delegate_token_response. Keeping the two in step by hand is what kept producing the same class of finding: a precondition guarded on one grant branch but not the other, master_key checked after the exchange on one path, identity resolved twice, and each failure raising an ad-hoc HTTPException with its own status and body shape. Model the mint as three phases whose failures are values. _prepare_bridge_mint runs before the exchange, checks every precondition once (master_key, then identity), and returns either a frozen _BridgeMintReady carrying the resolved key hash and the master-key-derived envelope keys, or a _BridgeMintError literal. Because every precondition lives in prepare, and prepare runs before the upstream POST, no failure can burn the single-use code or rotate a refresh token, for either grant type, by construction rather than by a guard we have to remember to keep in sync. _finish_bridge_mint runs after the exchange and has no preconditions left that can fail; its only failure values are properties of the upstream response itself (no usable access_token, or a token too large to seal). One mapper, _bridge_mint_error_response, turns each _BridgeMintError into an RFC 6749 section 5.2-shaped body with a status truthful about where the failure is (400 for the caller, 500 for gateway config, 502 for the upstream), with an exhaustive match plus assert_never so a new failure mode cannot be added without a matching status. Behavior is unchanged for the client. Every failure that previously raised now returns the same status as an OAuth error body, which is the correct token-endpoint contract; the three tests that asserted a raised HTTPException now assert the returned response. _exchange_for_bridge_server additionally asserts the identity resolver is awaited exactly once for a bridge server and never for a non-bridge one. --- .../mcp_server/discoverable_endpoints.py | 167 +++++++++++------- .../mcp_server/test_discoverable_endpoints.py | 63 ++++--- 2 files changed, 140 insertions(+), 90 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 5bf4de09bba..1bdbc8ea8a4 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -4,6 +4,7 @@ import json import math import secrets import time +from dataclasses import dataclass from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse @@ -12,6 +13,7 @@ import httpx from fastapi import APIRouter, Form, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response from pydantic import BaseModel, SecretStr, ValidationError +from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( @@ -39,6 +41,7 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( + EnvelopeKeys, UpstreamTokenGrant, ) from litellm.proxy._types import LiteLLM_MCPServerTable, UserAPIKeyAuth @@ -757,73 +760,112 @@ def _bridge_grant_from_token_response(token_response: object) -> Optional["Upstr ) -def _bridge_invalid_request_response() -> JSONResponse: - """RFC 6749 §5.2-shaped ``invalid_request`` for a bridge token exchange that carries no resolvable - litellm identity. Returned (not raised) so the OAuth error members sit at the top level rather than - wrapped in FastAPI's ``detail``, with the no-store token-endpoint headers, matching the BYOK OAuth - endpoint and what a strict DCR client parses per RFC 6749 §5.2.""" - return JSONResponse( - status_code=400, - content={ - "error": "invalid_request", - "error_description": ( +# --------------------------------------------------------------------------- +# DCR-bridge oauth_delegate mint: a three-phase pipeline whose failures are values. +# +# prepare (before the upstream exchange) -> validate every precondition and resolve identity+keys +# exchange (the single-use upstream code is consumed here, in exchange_token_with_server) +# finish (after the exchange) -> seal the upstream grant into the client-held envelope +# +# Every precondition lives in ``prepare``, which runs BEFORE the exchange, so no failure can burn the +# single-use code or rotate a refresh token, for either grant type -- that whole class of bug is gone +# by construction rather than guarded case by case. Failures are values mapped to an OAuth-shaped +# response in one place (``_bridge_mint_error_response``), so status codes and the RFC 6749 §5.2 body +# shape are uniform. Adding a failure mode is a new literal plus a match arm the type checker forces. +# --------------------------------------------------------------------------- + +_BridgeMintError = Literal["not_configured", "no_identity", "no_upstream_token", "too_large"] + + +@dataclass(frozen=True, slots=True) +class _BridgeMintReady: + """Everything the seal needs, resolved once before the exchange: the authorizing key hash and the + master-key-derived envelope keys. Passing this forward means identity resolution and key derivation + happen exactly once, and ``_finish_bridge_mint`` has no preconditions left that could fail.""" + + key_hash: str + keys: "EnvelopeKeys" + + +def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: + """Map a bridge-mint failure value to its token-endpoint response. One place, RFC 6749 §5.2 shape + (top-level ``error``, no-store) for every case, with a status truthful about where the failure is: + the caller's request (400), the gateway config (500), or the upstream (502).""" + if error == "no_identity": + status, code, desc = ( + 400, + "invalid_request", + ( "this server issues a gateway-bound credential; send a litellm credential " "(x-litellm-api-key or Authorization) on the token request" ), - }, - headers=TOKEN_NO_CACHE_HEADERS, + ) + elif error == "not_configured": + status, code, desc = ( + 500, + "server_error", + ("the gateway is not configured to mint a gateway-bound credential (master_key is not set)"), + ) + elif error == "no_upstream_token": + status, code, desc = 502, "server_error", "the upstream token response has no usable access_token" + elif error == "too_large": + status, code, desc = ( + 502, + "server_error", + ("the upstream token is too large to seal into a gateway-bound credential"), + ) + else: + assert_never(error) + return JSONResponse( + status_code=status, content={"error": code, "error_description": desc}, headers=TOKEN_NO_CACHE_HEADERS ) -async def _mint_bridge_delegate_token_response( - request: Request, mcp_server: MCPServer, token_response: object -) -> JSONResponse: - """Return the client-held envelope bearer for a DCR-bridge ``oauth_delegate`` token exchange. +async def _prepare_bridge_mint(request: Request, mcp_server: MCPServer) -> "_BridgeMintReady | _BridgeMintError": + """Phase 1, BEFORE the upstream exchange: validate that the gateway can mint (master_key set) and + that the request carries a resolvable litellm identity, and derive the envelope keys. Returns a + ready context or a failure value. Running before the exchange is what makes a missing master_key or + an unresolvable identity fail closed without consuming the single-use code / rotating a refresh + token, for both grant types.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import + envelope_keys_from_master_key, + ) + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import + master_key, + ) - The envelope binds the authorizing litellm key (its hash, resolved from the token request) to the - upstream grant, so the client holds one bearer that later admits it and forwards the upstream - token, with nothing stored server-side. Admission reloads the live key by that hash, so the key's - current restrictions and revocation gate the request. Fails closed with an OAuth - ``invalid_request`` when no active litellm key accompanies the token request rather than minting - an unbound credential. - """ + if not master_key: + return "not_configured" + key_hash = await _extract_active_key_hash_from_request(request) + if not key_hash: + return "no_identity" + return _BridgeMintReady(key_hash=key_hash, keys=envelope_keys_from_master_key(master_key)) + + +def _finish_bridge_mint( + ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime +) -> "JSONResponse | _BridgeMintError": + """Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held envelope using + the pre-resolved identity and keys, so the client holds one bearer that admits it and forwards the + upstream token with nothing stored server-side. The only failures here are properties of the + upstream response (no usable token, or a token too large to seal), returned as values.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import build_bridge_token_response, - envelope_keys_from_master_key, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import EnvelopeIdentity, SealedEnvelope, ) - from litellm.proxy.proxy_server import ( - master_key, # noqa: PLC0415 # inline import avoids a module-load circular import - ) - - if not master_key: - raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set") - - key_hash = await _extract_active_key_hash_from_request(request) - if not key_hash: - return _bridge_invalid_request_response() grant = _bridge_grant_from_token_response(token_response) if grant is None: - raise HTTPException(status_code=502, detail="Upstream token response has no usable access_token") - - now = datetime.now(timezone.utc) - keys = envelope_keys_from_master_key(master_key) - identity = EnvelopeIdentity(server_id=mcp_server.server_id, key_hash=key_hash) - sealed = build_bridge_token_response(identity, grant, keys, now) + return "no_upstream_token" + identity = EnvelopeIdentity(server_id=mcp_server.server_id, key_hash=ready.key_hash) + sealed = build_bridge_token_response(identity, grant, ready.keys, now) if not isinstance(sealed, SealedEnvelope): - # build_bridge_token_response returns EnvelopeTooLarge as a value when the upstream token is - # too large to seal; that is an upstream-payload condition, so surface a 502, not a 500. - raise HTTPException( - status_code=502, detail="Upstream token is too large to seal into a gateway-bound credential" - ) - - # The JWT exp is int(expires_at.timestamp()) (second-truncated), and admission expires the envelope - # against that exp. Report expires_in from the same truncated exp, rounding the elapsed portion up, - # so the client is never told the bearer lives past the point admission already rejects it. + return "too_large" + # Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the + # client is never told the bearer lives past the point admission (which uses that exp) rejects it. expires_in = max(1, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp())) body = {"access_token": sealed.token.get_secret_value(), "token_type": "Bearer", "expires_in": expires_in} return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS) @@ -902,19 +944,15 @@ async def exchange_token_with_server( if code_verifier: token_data["code_verifier"] = code_verifier - # A bridge oauth_delegate mint must fail closed BEFORE the upstream exchange consumes or rotates the - # single-use code (or refresh token): confirm the gateway can mint at all (master_key set) and that - # the request carries a resolvable litellm identity. Applies to both grant types, so an invalid key - # or a misconfigured gateway never burns the upstream credential. The mint below re-checks - # authoritatively; get_key_object is cache-first, so the identity re-resolution is a cache hit and - # adds no extra database round-trip. + # Phase 1: for a bridge oauth_delegate mint, validate all preconditions and resolve identity+keys + # BEFORE the exchange below consumes the single-use upstream code, and carry the ready context to + # phase 3. A failure here returns without ever touching the upstream credential. + bridge_mint_ready: _BridgeMintReady | None = None if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: - from litellm.proxy.proxy_server import master_key as _bridge_master_key # noqa: PLC0415 - - if not _bridge_master_key: - raise HTTPException(status_code=500, detail="Server misconfigured: master_key is not set") - if not await _extract_active_key_hash_from_request(request): - return _bridge_invalid_request_response() + prepared = await _prepare_bridge_mint(request, mcp_server) + if not isinstance(prepared, _BridgeMintReady): + return _bridge_mint_error_response(prepared) + bridge_mint_ready = prepared async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) response = await async_client.post( @@ -982,8 +1020,11 @@ async def exchange_token_with_server( # A DCR-bridge oauth_delegate server hands the client a gateway-bound envelope (identity plus the # upstream token) instead of the raw upstream token, so the one bearer both admits the caller and # forwards the upstream credential. Only this mode mints; every other server returns the raw token. - if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: - return await _mint_bridge_delegate_token_response(request, mcp_server, token_response) + if bridge_mint_ready is not None: + # Phase 3: seal the upstream grant into the client-held envelope; failures map through the same + # OAuth-shaped response as the phase-1 preconditions. + minted = _finish_bridge_mint(bridge_mint_ready, mcp_server, token_response, datetime.now(timezone.utc)) + return minted if isinstance(minted, JSONResponse) else _bridge_mint_error_response(minted) result = { "access_token": token_response["access_token"], diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4a4ff398915..4bbaf7b0b76 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4375,6 +4375,7 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_clie fake_http_response.raise_for_status = MagicMock() fake_http_client = MagicMock() fake_http_client.post = AsyncMock(return_value=fake_http_response) + key_resolver = AsyncMock(return_value=key_hash) if fake_client_out is not None: fake_client_out["client"] = fake_http_client @@ -4385,11 +4386,11 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_clie ), patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request", - new=AsyncMock(return_value=key_hash), + new=key_resolver, ), patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), ): - return await exchange_token_with_server( + response = await exchange_token_with_server( request=_bridge_mock_request(), mcp_server=server, grant_type="authorization_code", @@ -4399,6 +4400,11 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_clie client_secret=None, code_verifier="verifier", ) + if server.is_oauth_delegate and server.is_dcr_bridge: + key_resolver.assert_awaited_once() + else: + key_resolver.assert_not_awaited() + return response @pytest.mark.asyncio @@ -4457,15 +4463,16 @@ async def test_oauth_delegate_bridge_token_exchange_fails_closed_without_litellm @pytest.mark.asyncio async def test_bridge_envelope_too_large_upstream_token_is_502(): """An upstream token too large to seal into the envelope is an upstream-payload condition, so the - mint surfaces a 502 rather than a 500 (build_bridge_token_response returns EnvelopeTooLarge as a - value, and the caller maps it to a truthful status).""" + mint surfaces a 502 (as an RFC 6749 §5.2 error body, not a raised HTTPException) rather than a 500: + build_bridge_token_response returns EnvelopeTooLarge as a value, _finish_bridge_mint returns the + "too_large" failure, and _bridge_mint_error_response maps it to a truthful status.""" from litellm.types.mcp import MCPAuth server = _bridge_server(auth_type=MCPAuth.oauth_delegate) upstream = {"access_token": "x" * 40000, "token_type": "Bearer", "expires_in": 3600} - with pytest.raises(HTTPException) as exc: - await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") - assert exc.value.status_code == 502 + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + assert response.status_code == 502 + assert json.loads(response.body)["error"] == "server_error" @pytest.mark.asyncio @@ -4544,8 +4551,10 @@ async def test_bridge_refresh_grant_fails_closed_before_upstream_when_no_identit @pytest.mark.asyncio async def test_bridge_mint_fails_closed_before_upstream_when_master_key_unset(): - """master_key is validated BEFORE the upstream exchange, so a misconfigured gateway 500s without - consuming the single-use code, avoiding the burn-then-fail the pre-exchange gate exists to prevent.""" + """master_key is validated BEFORE the upstream exchange (in _prepare_bridge_mint), so a + misconfigured gateway returns a 500 server_error without consuming the single-use code, avoiding + the burn-then-fail the pre-exchange phase exists to prevent. The failure is returned as an RFC 6749 + error body, not raised.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server from litellm.types.mcp import MCPAuth @@ -4563,19 +4572,19 @@ async def test_bridge_mint_fails_closed_before_upstream_when_master_key_unset(): ), patch("litellm.proxy.proxy_server.master_key", None), ): - with pytest.raises(HTTPException) as exc: - await exchange_token_with_server( - request=_bridge_mock_request(), - mcp_server=server, - grant_type="authorization_code", - code="auth-code", - redirect_uri="https://claude.ai/api/mcp/auth_callback", - client_id="dcr-client-123", - client_secret=None, - code_verifier="verifier", - ) + response = await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="https://claude.ai/api/mcp/auth_callback", + client_id="dcr-client-123", + client_secret=None, + code_verifier="verifier", + ) - assert exc.value.status_code == 500 + assert response.status_code == 500 + assert json.loads(response.body)["error"] == "server_error" fake_http_client.post.assert_not_called() @@ -4647,18 +4656,18 @@ async def test_bridge_token_exchange_honors_short_float_expires_in_ttl(): @pytest.mark.asyncio async def test_oauth_delegate_bridge_token_exchange_missing_access_token_is_502_not_keyerror(): """When the upstream token response has no access_token, a dcr_bridge oauth_delegate exchange - returns a clean 502 rather than raising a KeyError. The eager access_token extraction used to run - before the bridge branch, so a missing token raised KeyError and _bridge_grant_from_token_response's - nil guard (which maps to 502) was dead code; the extraction now lives on the non-bridge path only.""" + returns a clean 502 error body rather than raising a KeyError. _finish_bridge_mint asks + _bridge_grant_from_token_response for a typed grant, gets None, and returns the "no_upstream_token" + failure, which maps to 502; nothing indexes token_response["access_token"] on the bridge path.""" from litellm.types.mcp import MCPAuth server = _bridge_server(auth_type=MCPAuth.oauth_delegate) upstream = {"token_type": "Bearer", "expires_in": 3600} - with pytest.raises(HTTPException) as exc: - await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") - assert exc.value.status_code == 502 + assert response.status_code == 502 + assert json.loads(response.body)["error"] == "server_error" @pytest.mark.asyncio From 4ba7221b7a1ae717d55ff548247aa8416dbe8232 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 17:47:54 -0700 Subject: [PATCH 045/123] fix(mcp): let the bridge envelope report expires_in 0 at the jwt exp boundary _finish_bridge_mint floored the reported expires_in at 1. Admission expires the envelope against the JWT's second-truncated exp, so when the mint lands in the same second that exp falls on (a sub-second upstream lifetime, for instance), the true remaining life is 0 and reporting 1 tells the client the bearer lives one second past the point admission already rejects it. Floor at 0 instead so the reported lifetime never overstates the exp; the value still cannot go negative. The regression pins the boundary directly: minting at now=100.25 with a 1s upstream token seals exp=101, and the reported expires_in is max(0, 101 - ceil(100.25)) = 0. Under the old floor of 1 it reads 1, so the test fails on that mutation. Also drops the unused mcp_server parameter from _prepare_bridge_mint; identity and key derivation there never referenced the server. --- .../mcp_server/discoverable_endpoints.py | 6 ++-- .../mcp_server/test_discoverable_endpoints.py | 29 +++++++++++++++++++ 2 files changed, 32 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 1bdbc8ea8a4..abb375b5b6a 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -821,7 +821,7 @@ def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: ) -async def _prepare_bridge_mint(request: Request, mcp_server: MCPServer) -> "_BridgeMintReady | _BridgeMintError": +async def _prepare_bridge_mint(request: Request) -> "_BridgeMintReady | _BridgeMintError": """Phase 1, BEFORE the upstream exchange: validate that the gateway can mint (master_key set) and that the request carries a resolvable litellm identity, and derive the envelope keys. Returns a ready context or a failure value. Running before the exchange is what makes a missing master_key or @@ -866,7 +866,7 @@ def _finish_bridge_mint( return "too_large" # Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the # client is never told the bearer lives past the point admission (which uses that exp) rejects it. - expires_in = max(1, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp())) + expires_in = max(0, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp())) body = {"access_token": sealed.token.get_secret_value(), "token_type": "Bearer", "expires_in": expires_in} return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS) @@ -949,7 +949,7 @@ async def exchange_token_with_server( # phase 3. A failure here returns without ever touching the upstream credential. bridge_mint_ready: _BridgeMintReady | None = None if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: - prepared = await _prepare_bridge_mint(request, mcp_server) + prepared = await _prepare_bridge_mint(request) if not isinstance(prepared, _BridgeMintReady): return _bridge_mint_error_response(prepared) bridge_mint_ready = prepared diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4bbaf7b0b76..4aa23247249 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4609,6 +4609,35 @@ async def test_bridge_reported_expires_in_does_not_overstate_jwt_exp(): assert before + body["expires_in"] <= claims["exp"] +def test_bridge_reported_expires_in_can_be_zero_at_jwt_exp_boundary(): + from datetime import datetime, timezone + + from fastapi.responses import JSONResponse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _BridgeMintReady, + _finish_bridge_mint, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( + envelope_keys_from_master_key, + ) + from litellm.types.mcp import MCPAuth + + ready = _BridgeMintReady( + key_hash="hashed-litellm-key-77", + keys=envelope_keys_from_master_key(_BRIDGE_MASTER_KEY), + ) + response = _finish_bridge_mint( + ready=ready, + mcp_server=_bridge_server(auth_type=MCPAuth.oauth_delegate), + token_response={"access_token": "UP", "expires_in": 1}, + now=datetime.fromtimestamp(100.25, tz=timezone.utc), + ) + + assert isinstance(response, JSONResponse) + assert json.loads(response.body)["expires_in"] == 0 + + def test_bridge_grant_coerces_numeric_expires_in(): """expires_in from an IdP may be an int, a float (3600.0), or a numeric string ("3600"); coerce it to a positive int so the envelope TTL honors the real lifetime instead of dropping a non-int From a07aba05798941c75e369789831e701801355daa Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 18:28:32 -0700 Subject: [PATCH 046/123] refactor(mcp): make bridge-mint resolvers return tagged unions so status is truthful by construction MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three findings landed together, all one defect: a resolution step crushed several distinct outcomes into a single None or a silent default, so the mint's error mapper could not tell them apart and assigned the wrong status. Identity resolution mapped a database outage to the same None as a missing credential, which the mint reported as 400 invalid_request, blaming the caller for a gateway outage while admission statuses the same outage 503/500. Lifetime coercion mapped an explicit non-positive expires_in to the same None as an absent one, so an upstream token the IdP reports as already dead was sealed into an hour-long envelope. And the refresh_token grant was run through the upstream exchange (which can rotate the client's upstream refresh credential) and its result then discarded, even though a bridge server seals no refresh_token and the client never holds one to present. Rather than add a mapping branch per finding, the fix changes the return types so a wrong status is not representable. Each resolution step now returns a precise tagged value instead of None: identity resolution returns a _ResolvedKey or one of no_active_key / unavailable / unresolvable, classified the same way admission's _reload_admitted_key classifies the same conditions; upstream-lifetime classification returns a positive number of seconds, "unspecified" (absent or unparseable, which the envelope caps), or "expired" (a parseable non-positive value, an already-dead token); and upstream-grant validation returns a typed grant or one of no_access_token / expired_lifetime. Thin exhaustive mappers (match plus assert_never) lift each vocabulary into one bridge-mint taxonomy of eight named failures, and a single _bridge_mint_error_response gives each its truthful RFC 6749 §5.2 status: 400 for the caller's missing credential or an unsupported grant, 503 for a transient auth-DB outage, 500 for a gateway that cannot resolve identity or is not configured, and 502 for an upstream response with no usable token, an already-expired lifetime, or a token too large to seal. Adding a failure mode now requires a new literal and a match arm the type checker forces, so the class of wrong-status bug cannot recur silently. The refresh_token grant is rejected in _prepare_bridge_mint before the exchange with unsupported_grant_type, so it can never rotate or consume the client's upstream refresh credential; renewal is re-running authorization_code, as the sealed refresh_token=None already intends. An absent or unparseable expires_in still mints a capped envelope (the by-design behaviour for an upstream that omits the field); only an explicitly-dead lifetime is rejected. Tests cover the resolver's three failure classes (including a real connection-error outage and a missing prisma_client), the mint statuses for each (503 before the upstream exchange, 500, 502 on an expired upstream lifetime, and a capped mint on an unknown one), and the refresh-grant rejection before any exchange. The three findings are mutation-checked: reverting each fix turns its regression test red. --- .../mcp_server/discoverable_endpoints.py | 334 ++++++++++++------ .../mcp_server/test_discoverable_endpoints.py | 259 +++++++++++--- 2 files changed, 424 insertions(+), 169 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index abb375b5b6a..b8d10335182 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -377,74 +377,92 @@ def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None: return key_obj.user_id if _key_is_active(key_obj) else None -async def _resolve_active_litellm_key(request: Request) -> tuple[str, "UserAPIKeyAuth"] | None: - """Resolve the presented litellm key to ``(its hash, the live active key record)``, or ``None`` - when the key is absent, unresolvable, or blocked/expired. +@dataclass(frozen=True, slots=True) +class _ResolvedKey: + """An active litellm key resolved from the token request: its hash (the value ``get_key_object`` + and the cache/DB layer key the record by) and the live record.""" - Single resolution path the OAuth token endpoint reuses. Resolves authoritatively via - ``get_key_object`` (cache first, then DB) instead of a raw cache peek. On a multi-replica gateway - the token-exchange request can land on a worker whose in-memory cache never saw the key, and a - cross-replica Redis hit deserializes to a plain ``dict`` rather than a ``UserAPIKeyAuth``; the - previous code read only ``Authorization`` and did ``getattr(cached, "user_id")`` with no - ``model_type`` rehydration and no DB fallback, so it silently returned ``None``. The resolved key - is validated (``_key_is_active``) before it is trusted, so a blocked or expired key resolves to - ``None``, while a valid team-scoped or service-account key (no ``user_id``) still resolves so it - can mint a bridge envelope. The returned hash is the value ``get_key_object`` and the cache/DB - layer key the record by. Callers derive the ``user_id`` (per-user token store) or seal the hash - (dcr_bridge envelope) from the result. - """ + key_hash: str + key: "UserAPIKeyAuth" + + +_KeyResolutionFailure = Literal["no_active_key", "unavailable", "unresolvable"] +"""Why a token request yielded no active litellm key, kept distinct so a caller statuses each truthfully +instead of blaming the client for a gateway problem: +- ``no_active_key``: none was presented, or the presented key is unknown / blocked / expired (the + caller's request is at fault) +- ``unavailable``: the auth database was transiently unreachable while resolving (retryable) +- ``unresolvable``: the gateway cannot resolve identity right now (no DB connection, or an unexpected + error) -- a gateway fault, not the caller's +The classification mirrors admission's ``_reload_admitted_key`` so the mint (ingress) and admission +(egress) never disagree on the status of the same outage.""" + + +async def _resolve_active_litellm_key(request: Request) -> "_ResolvedKey | _KeyResolutionFailure": + """Resolve the presented litellm key to an active key record, or say precisely why not. + + Single resolution path the OAuth token endpoint reuses, resolving authoritatively via + ``get_key_object`` (cache first, then DB). The failure is a value, not a bare ``None``, so a caller + can tell "the client sent no usable credential" (a request error) apart from "the gateway could not + check" (an infrastructure error) and status each truthfully; collapsing both to ``None`` is what let + a DB outage read as a 400. A resolved key is still gated by ``_key_is_active``, so a blocked or + expired key is ``no_active_key`` while a valid team-scoped or service-account key (no ``user_id``) + resolves. Classification mirrors admission's ``_reload_admitted_key``: no DB connection is a gateway + fault, a ``ProxyException`` / ``HTTPException`` from ``get_key_object`` is an unknown or invalid key, + a database-service-unavailable error is a retryable outage, and anything else is an unexpected + gateway fault.""" token = _litellm_key_from_request(request) if not token: - return None - try: - from litellm.proxy._types import ( # noqa: PLC0415 # inline import avoids a module-load circular import - hash_token, - ) - from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import - get_key_object, - ) - from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import - prisma_client, - user_api_key_cache, - ) + return "no_active_key" + from litellm.proxy._types import ( # noqa: PLC0415 # inline import avoids a module-load circular import + ProxyException, + hash_token, + ) + from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import + get_key_object, + ) + from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import + PrismaDBExceptionHandler, + ) + from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import + prisma_client, + user_api_key_cache, + ) - key_hash = hash_token(token) + if prisma_client is None: + return "unresolvable" + key_hash = hash_token(token) + try: key_obj = await get_key_object( hashed_token=key_hash, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, ) - except Exception as exc: # noqa: BLE001 # fail closed to None on any key-resolution error + except (ProxyException, HTTPException): + return "no_active_key" + except Exception as exc: # noqa: BLE001 # classify: a DB outage is retryable, anything else is an opaque gateway fault + if PrismaDBExceptionHandler.is_database_service_unavailable_error(exc): + return "unavailable" verbose_logger.debug( - "_resolve_active_litellm_key: could not resolve the presented key (%s)", + "_resolve_active_litellm_key: unexpected key-resolution error (%s)", type(exc).__name__, ) - return None + return "unresolvable" if not _key_is_active(key_obj): - return None - return key_hash, key_obj + return "no_active_key" + return _ResolvedKey(key_hash=key_hash, key=key_obj) async def _extract_user_id_from_request(request: Request) -> str | None: - """The LiteLLM ``user_id`` for the token request, so a per-user token is stored under the same - identity the egress later reads it by (``user_api_key_auth.user_id``). ``None`` when no active - key is present. See :func:`_resolve_active_litellm_key` for the resolution and active-key gate. - """ + """The litellm ``user_id`` for the token request, so a per-user token is stored under the same + identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome + (including a transient DB outage) collapses to ``None`` here and the caller simply skips the store; + the bridge mint, which must status those outcomes differently, consumes + :func:`_resolve_active_litellm_key` directly.""" resolved = await _resolve_active_litellm_key(request) - return _active_key_user_id(resolved[1]) if resolved else None - - -async def _extract_active_key_hash_from_request(request: Request) -> str | None: - """The hash of the litellm key that authorized the token request, when it maps to an active key. - - A DCR-bridge envelope seals this hash so admission can reload the live ``UserAPIKeyAuth`` record - and enforce the key's current team/org/tool restrictions and revocation, rather than trusting a - frozen identity. The hash is a one-way digest, not a usable credential (the edge rejects a bare - hash presented as a bearer). ``None`` when no active key is present, so no envelope is minted for - a missing, unresolvable, or revoked key. - """ - resolved = await _resolve_active_litellm_key(request) - return resolved[0] if resolved else None + if not isinstance(resolved, _ResolvedKey): + return None + return _active_key_user_id(resolved.key) async def _store_per_user_token_server_side( @@ -713,38 +731,51 @@ async def authorize_with_server( return response -def _coerce_positive_expires_in(value: object) -> int | None: - """Coerce an upstream ``expires_in`` to a positive int, or ``None`` when it is absent or not a - usable number. IdPs return it as an int, a float (``3600.0``), or a numeric string (``"3600"``); - accepting only ``int`` would drop the float/string cases to ``None`` and fall back to the - envelope's 1h cap, which can outlive a shorter-lived upstream token and forward a stale bearer. - ``bool`` is excluded (it is an ``int`` subclass but never a real lifetime). Total over hostile - input: a non-numeric string, ``NaN``, ``Infinity``, or an over-large value all resolve to - ``None`` rather than raising (``int(float(...))`` can raise ``ValueError`` or ``OverflowError``), - so a malformed upstream ``expires_in`` never surfaces as a 500 from the token endpoint.""" - if isinstance(value, bool) or not isinstance(value, (int, float, str)): - return None +_UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"] +"""Why an upstream token response cannot back a bridge envelope: +- ``no_access_token``: the response carries no usable ``access_token`` +- ``expired_lifetime``: the response reports a parseable, non-positive ``expires_in``, i.e. an upstream + token that is already dead, so sealing it would forward a bearer the edge cannot use +An absent or unparseable ``expires_in`` is NOT a rejection; the lifetime is merely unknown and the +envelope caps it, the by-design behaviour for an upstream that omits the field.""" + + +def _classify_upstream_lifetime(raw_expires_in: object) -> "int | Literal['unspecified', 'expired']": + """Classify an upstream ``expires_in`` into a positive number of seconds, ``"unspecified"`` (absent + or unparseable, so the envelope caps it), or ``"expired"`` (a parseable non-positive value the + upstream reports as already elapsed). Telling "we do not know the lifetime" apart from "the upstream + says it is already dead" is what stops an explicitly-expired token from silently receiving the + envelope's 1h cap. ``bool`` is excluded (an ``int`` subclass but never a real lifetime), and + ``int(float(...))`` can raise on ``NaN`` / ``Infinity`` / oversized input, which reads as + unparseable rather than surfacing as a 500.""" + if raw_expires_in is None or isinstance(raw_expires_in, bool) or not isinstance(raw_expires_in, (int, float, str)): + return "unspecified" try: - seconds = int(float(value)) + seconds = int(float(raw_expires_in)) except (ValueError, TypeError, OverflowError): - return None - return seconds if seconds > 0 else None + return "unspecified" + return seconds if seconds > 0 else "expired" -def _bridge_grant_from_token_response(token_response: object) -> Optional["UpstreamTokenGrant"]: - """Validate an upstream OAuth token response into a typed grant, or None when it lacks a usable - access token. Each field is isinstance-checked so nothing untyped from ``response.json()`` flows - into the grant; ``expires_in`` is numerically coerced so a float/string lifetime is honored - rather than dropped to the envelope's default cap.""" +def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenGrant | _UpstreamGrantRejection": + """Validate an upstream OAuth token response into a typed grant, or say why it cannot back an + envelope. Each field is isinstance-checked so nothing untyped from ``response.json()`` reaches the + grant. ``expires_in`` is read three ways (see :func:`_classify_upstream_lifetime`): an unknown + lifetime leaves the grant ``expires_in`` ``None`` for the envelope to cap, a positive value is + honoured, and an explicit already-elapsed value is a rejection rather than a silent fall-through to + the cap.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import UpstreamTokenGrant, ) if not isinstance(token_response, dict): - return None + return "no_access_token" access = token_response.get("access_token") if not isinstance(access, str) or not access: - return None + return "no_access_token" + lifetime = _classify_upstream_lifetime(token_response.get("expires_in")) + if lifetime == "expired": + return "expired_lifetime" token_type = token_response.get("token_type") scope = token_response.get("scope") return UpstreamTokenGrant( @@ -756,7 +787,7 @@ def _bridge_grant_from_token_response(token_response: object) -> Optional["Upstr # follow-up (a dedicated refresh-envelope); the client re-runs authorization_code at the cap. refresh_token=None, scope=scope if isinstance(scope, str) and scope else None, - expires_in=_coerce_positive_expires_in(token_response.get("expires_in")), + expires_in=lifetime if isinstance(lifetime, int) else None, ) @@ -774,7 +805,16 @@ def _bridge_grant_from_token_response(token_response: object) -> Optional["Upstr # shape are uniform. Adding a failure mode is a new literal plus a match arm the type checker forces. # --------------------------------------------------------------------------- -_BridgeMintError = Literal["not_configured", "no_identity", "no_upstream_token", "too_large"] +_BridgeMintError = Literal[ + "no_identity", + "unsupported_grant", + "identity_unavailable", + "identity_unresolvable", + "not_configured", + "no_upstream_token", + "upstream_token_expired", + "too_large", +] @dataclass(frozen=True, slots=True) @@ -788,45 +828,105 @@ class _BridgeMintReady: def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: - """Map a bridge-mint failure value to its token-endpoint response. One place, RFC 6749 §5.2 shape - (top-level ``error``, no-store) for every case, with a status truthful about where the failure is: - the caller's request (400), the gateway config (500), or the upstream (502).""" - if error == "no_identity": - status, code, desc = ( - 400, - "invalid_request", - ( + """Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape + (top-level ``error``, no-store headers) for every case, with a status truthful about where the + failure is. The caller's request is 400, a transient gateway outage is 503, a gateway + misconfiguration is 500, and an upstream problem is 502. The identity-resolution statuses match how + admission statuses the same conditions on the egress side, so mint and admit never disagree under + one outage.""" + match error: + case "no_identity": + status, code, desc = ( + 400, + "invalid_request", "this server issues a gateway-bound credential; send a litellm credential " - "(x-litellm-api-key or Authorization) on the token request" - ), - ) - elif error == "not_configured": - status, code, desc = ( - 500, - "server_error", - ("the gateway is not configured to mint a gateway-bound credential (master_key is not set)"), - ) - elif error == "no_upstream_token": - status, code, desc = 502, "server_error", "the upstream token response has no usable access_token" - elif error == "too_large": - status, code, desc = ( - 502, - "server_error", - ("the upstream token is too large to seal into a gateway-bound credential"), - ) - else: - assert_never(error) + "(x-litellm-api-key or Authorization) on the token request", + ) + case "unsupported_grant": + status, code, desc = ( + 400, + "unsupported_grant_type", + "this server issues a gateway-bound credential and supports only the authorization_code " + "grant; re-run authorization_code to renew rather than refresh_token", + ) + case "identity_unavailable": + status, code, desc = ( + 503, + "temporarily_unavailable", + "the authentication database is temporarily unreachable; retry shortly", + ) + case "identity_unresolvable": + status, code, desc = ( + 500, + "server_error", + "the gateway could not resolve the litellm identity for this request", + ) + case "not_configured": + status, code, desc = ( + 500, + "server_error", + "the gateway is not configured to mint a gateway-bound credential (master_key is not set)", + ) + case "no_upstream_token": + status, code, desc = ( + 502, + "server_error", + "the upstream token response has no usable access_token", + ) + case "upstream_token_expired": + status, code, desc = ( + 502, + "server_error", + "the upstream token response reports an already-expired lifetime", + ) + case "too_large": + status, code, desc = ( + 502, + "server_error", + "the upstream token is too large to seal into a gateway-bound credential", + ) + case _: + assert_never(error) return JSONResponse( status_code=status, content={"error": code, "error_description": desc}, headers=TOKEN_NO_CACHE_HEADERS ) -async def _prepare_bridge_mint(request: Request) -> "_BridgeMintReady | _BridgeMintError": - """Phase 1, BEFORE the upstream exchange: validate that the gateway can mint (master_key set) and - that the request carries a resolvable litellm identity, and derive the envelope keys. Returns a - ready context or a failure value. Running before the exchange is what makes a missing master_key or - an unresolvable identity fail closed without consuming the single-use code / rotating a refresh - token, for both grant types.""" +def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError: + """Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays + truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that + cannot resolve identity is 500.""" + match failure: + case "no_active_key": + return "no_identity" + case "unavailable": + return "identity_unavailable" + case "unresolvable": + return "identity_unresolvable" + case _: + assert_never(failure) + + +def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _BridgeMintError: + """Lift an upstream-response rejection into the mint taxonomy; both are upstream faults (502).""" + match rejection: + case "no_access_token": + return "no_upstream_token" + case "expired_lifetime": + return "upstream_token_expired" + case _: + assert_never(rejection) + + +async def _prepare_bridge_mint(request: Request, grant_type: str) -> "_BridgeMintReady | _BridgeMintError": + """Phase 1, BEFORE the upstream exchange: reject a grant this mint does not support, confirm the + gateway can mint (master_key set), resolve the litellm identity, and derive the envelope keys. + Returns a ready context or a precise failure value. Running before the exchange is what makes every + failure here fail closed without consuming the single-use code or rotating a refresh token. A bridge + server issues only envelopes and seals no upstream refresh_token, so the client holds none to + present: the refresh_token grant is rejected up front rather than exchanged (which could rotate the + upstream credential) and its result then discarded. Identity-resolution failures keep their origin + so the mapper statuses each truthfully.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import envelope_keys_from_master_key, ) @@ -834,12 +934,14 @@ async def _prepare_bridge_mint(request: Request) -> "_BridgeMintReady | _BridgeM master_key, ) + if grant_type != "authorization_code": + return "unsupported_grant" if not master_key: return "not_configured" - key_hash = await _extract_active_key_hash_from_request(request) - if not key_hash: - return "no_identity" - return _BridgeMintReady(key_hash=key_hash, keys=envelope_keys_from_master_key(master_key)) + resolved = await _resolve_active_litellm_key(request) + if not isinstance(resolved, _ResolvedKey): + return _key_resolution_failure_to_mint_error(resolved) + return _BridgeMintReady(key_hash=resolved.key_hash, keys=envelope_keys_from_master_key(master_key)) def _finish_bridge_mint( @@ -848,18 +950,20 @@ def _finish_bridge_mint( """Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held envelope using the pre-resolved identity and keys, so the client holds one bearer that admits it and forwards the upstream token with nothing stored server-side. The only failures here are properties of the - upstream response (no usable token, or a token too large to seal), returned as values.""" + upstream response (no usable token, an already-expired lifetime, or a token too large to seal), + returned as values.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import build_bridge_token_response, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import EnvelopeIdentity, SealedEnvelope, + UpstreamTokenGrant, ) grant = _bridge_grant_from_token_response(token_response) - if grant is None: - return "no_upstream_token" + if not isinstance(grant, UpstreamTokenGrant): + return _upstream_rejection_to_mint_error(grant) identity = EnvelopeIdentity(server_id=mcp_server.server_id, key_hash=ready.key_hash) sealed = build_bridge_token_response(identity, grant, ready.keys, now) if not isinstance(sealed, SealedEnvelope): @@ -949,7 +1053,7 @@ async def exchange_token_with_server( # phase 3. A failure here returns without ever touching the upstream credential. bridge_mint_ready: _BridgeMintReady | None = None if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: - prepared = await _prepare_bridge_mint(request) + prepared = await _prepare_bridge_mint(request, grant_type) if not isinstance(prepared, _BridgeMintReady): return _bridge_mint_error_response(prepared) bridge_mint_ready = prepared diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 4aa23247249..f611d2f6a24 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4367,6 +4367,7 @@ _BRIDGE_MASTER_KEY = "sk-bridge-producer-master-key-0123456789abcdef" async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_client_out=None): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _ResolvedKey, exchange_token_with_server, ) @@ -4375,7 +4376,10 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_clie fake_http_response.raise_for_status = MagicMock() fake_http_client = MagicMock() fake_http_client.post = AsyncMock(return_value=fake_http_response) - key_resolver = AsyncMock(return_value=key_hash) + # The mint consumes _resolve_active_litellm_key's tagged result: an active key resolves to a + # _ResolvedKey carrying its hash; a request with no usable credential resolves to "no_active_key". + resolution = _ResolvedKey(key_hash=key_hash, key=MagicMock()) if key_hash is not None else "no_active_key" + key_resolver = AsyncMock(return_value=resolution) if fake_client_out is not None: fake_client_out["client"] = fake_http_client @@ -4385,7 +4389,7 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_clie return_value=fake_http_client, ), patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request", + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_active_litellm_key", new=key_resolver, ), patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), @@ -4511,10 +4515,12 @@ async def test_bridge_envelope_does_not_seal_upstream_refresh_token(): @pytest.mark.asyncio -async def test_bridge_refresh_grant_fails_closed_before_upstream_when_no_identity(): - """The pre-exchange identity gate covers the refresh_token grant, not just authorization_code: an - unresolvable litellm identity fails closed with invalid_request BEFORE the upstream refresh is - exchanged, so the client's refresh token is not rotated/consumed on a rejected request.""" +async def test_bridge_refresh_grant_is_rejected_before_upstream(): + """A bridge oauth_delegate server issues only envelopes and seals no upstream refresh_token, so the + client never holds one to present. _prepare_bridge_mint rejects the refresh_token grant up front + with unsupported_grant_type, BEFORE any upstream exchange, so a stray refresh request can never + rotate or consume the client's upstream refresh credential; renewal is re-running + authorization_code. This is checked before identity resolution, so it holds even with a valid key.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server from litellm.types.mcp import MCPAuth @@ -4526,10 +4532,6 @@ async def test_bridge_refresh_grant_fails_closed_before_upstream_when_no_identit "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", return_value=fake_http_client, ), - patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request", - new=AsyncMock(return_value=None), - ), patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), ): response = await exchange_token_with_server( @@ -4545,7 +4547,7 @@ async def test_bridge_refresh_grant_fails_closed_before_upstream_when_no_identit ) assert response.status_code == 400 - assert json.loads(response.body)["error"] == "invalid_request" + assert json.loads(response.body)["error"] == "unsupported_grant_type" fake_http_client.post.assert_not_called() @@ -4567,8 +4569,8 @@ async def test_bridge_mint_fails_closed_before_upstream_when_master_key_unset(): return_value=fake_http_client, ), patch( - "litellm.proxy._experimental.mcp_server.discoverable_endpoints._extract_active_key_hash_from_request", - new=AsyncMock(return_value="hashed-litellm-key-77"), + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_active_litellm_key", + new=AsyncMock(return_value="no_active_key"), ), patch("litellm.proxy.proxy_server.master_key", None), ): @@ -4588,6 +4590,93 @@ async def test_bridge_mint_fails_closed_before_upstream_when_master_key_unset(): fake_http_client.post.assert_not_called() +async def _prepare_only_bridge_exchange(resolver_result): + """Drive exchange_token_with_server for a bridge oauth_delegate authorization_code request with the + identity resolver stubbed to a given tagged result, returning (response, post_mock) so a test can + assert the mapped status and that the single-use code was never exchanged.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + fake_http_client = MagicMock() + fake_http_client.post = AsyncMock() + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=fake_http_client, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_active_litellm_key", + new=AsyncMock(return_value=resolver_result), + ), + patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), + ): + response = await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="https://claude.ai/api/mcp/auth_callback", + client_id="dcr-client-123", + client_secret=None, + code_verifier="verifier", + ) + return response, fake_http_client.post + + +@pytest.mark.asyncio +async def test_bridge_mint_db_outage_is_503_before_upstream(): + """A DB outage while resolving identity is a retryable gateway failure, so the mint returns 503 + temporarily_unavailable WITHOUT consuming the single-use code, matching how admission statuses the + same outage on the egress side. Collapsing every resolution failure to None used to blame the + client with 400 invalid_request for an infrastructure problem.""" + response, post = await _prepare_only_bridge_exchange("unavailable") + assert response.status_code == 503 + assert json.loads(response.body)["error"] == "temporarily_unavailable" + post.assert_not_called() + + +@pytest.mark.asyncio +async def test_bridge_mint_unresolvable_identity_is_500_before_upstream(): + """An unresolvable identity (no DB connection, or an unexpected resolution error) is a gateway + fault, so the mint returns 500 server_error before the exchange, a status distinct from both the + caller's 400 and the transient 503, matching admission's 500-vs-503 split for the same conditions.""" + response, post = await _prepare_only_bridge_exchange("unresolvable") + assert response.status_code == 500 + assert json.loads(response.body)["error"] == "server_error" + post.assert_not_called() + + +@pytest.mark.asyncio +async def test_bridge_mint_upstream_expired_lifetime_is_502(): + """An upstream token response reporting an already-elapsed lifetime (a parseable non-positive + expires_in) is rejected with 502 rather than sealed into an hour-long envelope around a dead + bearer. Regression for expires_in<=0 silently falling through to the 1h cap.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "UP", "token_type": "Bearer", "expires_in": 0} + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + assert response.status_code == 502 + assert json.loads(response.body)["error"] == "server_error" + + +@pytest.mark.asyncio +async def test_bridge_mint_unknown_lifetime_is_capped_not_rejected(): + """An absent or unparseable expires_in leaves the lifetime unknown, which the envelope caps (never + inventing a longer life than the upstream stated); it is NOT rejected. Only an explicitly-dead + lifetime fails, so a metadata glitch on an otherwise-valid token still mints a bounded envelope.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "UP", "token_type": "Bearer", "expires_in": "not-a-number"} + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + assert response.status_code == 200 + body = json.loads(response.body) + assert body["access_token"].startswith("llm_env_") + assert 0 < body["expires_in"] <= 3600 + + @pytest.mark.asyncio async def test_bridge_reported_expires_in_does_not_overstate_jwt_exp(): """The reported expires_in is derived from the envelope JWT's second-truncated exp (rounding the @@ -4638,34 +4727,51 @@ def test_bridge_reported_expires_in_can_be_zero_at_jwt_exp_boundary(): assert json.loads(response.body)["expires_in"] == 0 -def test_bridge_grant_coerces_numeric_expires_in(): - """expires_in from an IdP may be an int, a float (3600.0), or a numeric string ("3600"); coerce - it to a positive int so the envelope TTL honors the real lifetime instead of dropping a non-int - value and defaulting to the 1h cap (which can outlive a shorter-lived upstream token). bool and - non-numeric values become None.""" - from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _bridge_grant_from_token_response, - ) +def test_classify_upstream_lifetime(): + """expires_in from an IdP may be an int, a float (3600.0), or a numeric string ("3600"); each + coerces to a positive number of seconds. Absent or unparseable input (bool, non-numeric, NaN/inf, + oversized) is "unspecified" so the envelope caps it, while a parseable non-positive value is + "expired": the upstream reporting an already-dead token, which the mint must reject rather than + silently give the 1h cap.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _classify_upstream_lifetime - def ei(v): - return _bridge_grant_from_token_response({"access_token": "x", "expires_in": v}).expires_in + assert _classify_upstream_lifetime(300) == 300 + assert _classify_upstream_lifetime(300.0) == 300 + assert _classify_upstream_lifetime("300") == 300 + assert _classify_upstream_lifetime(" 300 ") == 300 + # explicit, parseable, non-positive -> the upstream says the token is already dead + assert _classify_upstream_lifetime(0) == "expired" + assert _classify_upstream_lifetime(-5) == "expired" + # unknown lifetime -> cap (never invent a longer life than the upstream stated) + assert _classify_upstream_lifetime(None) == "unspecified" + assert _classify_upstream_lifetime(True) == "unspecified" + assert _classify_upstream_lifetime("nope") == "unspecified" + # hostile numerics must not raise (int(float(...)) can OverflowError) -> unspecified + assert _classify_upstream_lifetime("inf") == "unspecified" + assert _classify_upstream_lifetime("1e999") == "unspecified" + assert _classify_upstream_lifetime("-inf") == "unspecified" + assert _classify_upstream_lifetime("nan") == "unspecified" + assert _classify_upstream_lifetime(float("inf")) == "unspecified" + assert _classify_upstream_lifetime(10**400) == "unspecified" - assert ei(300) == 300 - assert ei(300.0) == 300 - assert ei("300") == 300 - assert ei(" 300 ") == 300 - assert ei(True) is None - assert ei("nope") is None - assert ei(0) is None - assert ei(-5) is None - assert ei(None) is None - # hostile numerics must not raise (int(float(...)) can OverflowError) -> None - assert ei("inf") is None - assert ei("1e999") is None - assert ei("-inf") is None - assert ei("nan") is None - assert ei(float("inf")) is None - assert ei(10**400) is None + +def test_bridge_grant_honors_and_rejects_upstream_lifetime(): + """The grant validator honors a positive lifetime, leaves an unknown one None for the envelope to + cap, and rejects an explicitly-expired one with "expired_lifetime" so a dead upstream token is + never sealed into an hour-long envelope.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _bridge_grant_from_token_response + + def grant(v): + return _bridge_grant_from_token_response({"access_token": "x", "expires_in": v}) + + assert grant(300).expires_in == 300 + assert grant(120.0).expires_in == 120 + # unknown lifetime backs a grant whose expires_in the envelope caps; it is not a rejection + assert grant("nope").expires_in is None + assert _bridge_grant_from_token_response({"access_token": "x"}).expires_in is None + # an explicitly already-dead lifetime is rejected, not silently capped at 1h + assert grant(0) == "expired_lifetime" + assert grant(-5) == "expired_lifetime" @pytest.mark.asyncio @@ -5073,13 +5179,14 @@ async def test_extract_user_id_rejects_expired_key(proxy_globals): @pytest.mark.asyncio -async def test_extract_active_key_hash_returns_hash_for_active_key(proxy_globals): +async def test_resolve_active_litellm_key_returns_resolved_key_for_active_key(proxy_globals): """The dcr_bridge mint seals the hash of the authorizing key so admission can reload the live record. For an active key the resolver returns exactly hash_token(key), the same value get_key_object and the whole cache/DB layer key the record by, so the sealed reference resolves back to this key at admission.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _extract_active_key_hash_from_request, + _resolve_active_litellm_key, + _ResolvedKey, ) from litellm.proxy._types import UserAPIKeyAuth, hash_token from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -5095,19 +5202,22 @@ async def test_extract_active_key_hash_returns_hash_for_active_key(proxy_globals proxy_globals.prisma_client = object() request = _token_request({"x-litellm-api-key": f"Bearer {key}"}) - assert await _extract_active_key_hash_from_request(request) == hash_token(key) + resolved = await _resolve_active_litellm_key(request) + assert isinstance(resolved, _ResolvedKey) + assert resolved.key_hash == hash_token(key) @pytest.mark.asyncio -async def test_extract_active_key_hash_returns_hash_for_active_key_without_user_id(proxy_globals): +async def test_resolve_active_litellm_key_resolves_key_without_user_id(proxy_globals): """A valid team-scoped or service-account key has no user_id but is a legitimate credential, so it must still resolve to a hash and be able to mint a bridge envelope. Gating the resolver on user_id presence wrongly rejected these keys with invalid_request; the active-state gate now checks only blocked and expiry, and the key hash (not the user) is what the mint seals. The per-user token store still gets no user for such a key, since there is none to key a stored credential by.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _extract_active_key_hash_from_request, _extract_user_id_from_request, + _resolve_active_litellm_key, + _ResolvedKey, ) from litellm.proxy._types import UserAPIKeyAuth, hash_token from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -5123,16 +5233,18 @@ async def test_extract_active_key_hash_returns_hash_for_active_key_without_user_ proxy_globals.prisma_client = object() request = _token_request({"x-litellm-api-key": f"Bearer {key}"}) - assert await _extract_active_key_hash_from_request(request) == hash_token(key) + resolved = await _resolve_active_litellm_key(request) + assert isinstance(resolved, _ResolvedKey) + assert resolved.key_hash == hash_token(key) assert await _extract_user_id_from_request(request) is None @pytest.mark.asyncio -async def test_extract_active_key_hash_rejects_blocked_key(proxy_globals): +async def test_resolve_active_litellm_key_rejects_blocked_key(proxy_globals): """A blocked key must not yield a hash, so no gateway-bound envelope is minted for a revoked key; the mint fails closed with invalid_request instead.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _extract_active_key_hash_from_request, + _resolve_active_litellm_key, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -5145,17 +5257,17 @@ async def test_extract_active_key_hash_rejects_blocked_key(proxy_globals): proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": "sk-blocked-key"}) - assert await _extract_active_key_hash_from_request(request) is None + assert await _resolve_active_litellm_key(request) == "no_active_key" @pytest.mark.asyncio -async def test_extract_active_key_hash_fails_closed_on_malformed_expiry(proxy_globals): +async def test_resolve_active_litellm_key_fails_closed_on_malformed_expiry(proxy_globals): """A key whose stored expires string does not parse must fail closed to no-hash (the mint then returns invalid_request), not surface an unhandled 500. The active-state check runs outside the resolver's try, so it must be total over a bad expires rather than letting datetime.fromisoformat raise. Before the fix this raised a ValueError instead of returning None.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _extract_active_key_hash_from_request, + _resolve_active_litellm_key, ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -5168,14 +5280,14 @@ async def test_extract_active_key_hash_fails_closed_on_malformed_expiry(proxy_gl proxy_globals.prisma_client = _FakePrisma() request = _token_request({"x-litellm-api-key": "sk-bad-expiry-key"}) - assert await _extract_active_key_hash_from_request(request) is None + assert await _resolve_active_litellm_key(request) == "no_active_key" @pytest.mark.asyncio -async def test_extract_active_key_hash_none_without_litellm_key(proxy_globals): +async def test_resolve_active_litellm_key_no_active_key_without_litellm_key(proxy_globals): """No LiteLLM key on the request yields no hash without consulting the resolver.""" from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - _extract_active_key_hash_from_request, + _resolve_active_litellm_key, ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -5183,7 +5295,46 @@ async def test_extract_active_key_hash_none_without_litellm_key(proxy_globals): proxy_globals.prisma_client = object() request = _token_request({"content-type": "application/json"}) - assert await _extract_active_key_hash_from_request(request) is None + assert await _resolve_active_litellm_key(request) == "no_active_key" + + +@pytest.mark.asyncio +async def test_resolve_active_litellm_key_db_outage_is_unavailable(proxy_globals): + """A database outage while resolving the presented key is a retryable infrastructure failure, not + the caller's fault, so the resolver reports "unavailable" (the mint statuses it 503) rather than + collapsing it to the same value as a missing credential. is_database_service_unavailable_error + classifies a connection error (an OSError) as an outage, matching admission's egress-side handling.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _resolve_active_litellm_key, + ) + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + class _OutagePrisma: + async def get_data(self, token, table_name, parent_otel_span=None, proxy_logging_obj=None): + raise ConnectionError("connection refused") + + proxy_globals.user_api_key_cache = UserApiKeyCache() + proxy_globals.prisma_client = _OutagePrisma() + + request = _token_request({"x-litellm-api-key": "sk-during-outage"}) + assert await _resolve_active_litellm_key(request) == "unavailable" + + +@pytest.mark.asyncio +async def test_resolve_active_litellm_key_no_database_is_unresolvable(proxy_globals): + """With no database connection configured the gateway cannot verify the presented key at all, so + the resolver reports "unresolvable" (the mint statuses it 500) instead of blaming the caller. + Mirrors admission, which 500s a missing prisma_client on the egress side.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _resolve_active_litellm_key, + ) + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + proxy_globals.user_api_key_cache = UserApiKeyCache() + proxy_globals.prisma_client = None + + request = _token_request({"x-litellm-api-key": "sk-no-db"}) + assert await _resolve_active_litellm_key(request) == "unresolvable" @pytest.mark.asyncio From 55ff3a242c269228dc1d54d0ffaffa40fb66ca5d Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 18:39:00 -0700 Subject: [PATCH 047/123] fix(mcp): treat a positive sub-second upstream lifetime as alive, not expired _classify_upstream_lifetime decided "expired" from int(float(expires_in)), which truncates toward zero, so a positive fractional lifetime in (0, 1) became 0 and was misread as already elapsed. That rejected the mint with 502 in _finish_bridge_mint after the single-use upstream code had already been consumed, even though the upstream reported a positive remaining lifetime. Decide expired on the parsed numeric value rather than its truncated int, so only a genuinely non-positive value is expired. The envelope works in whole seconds and cannot represent a sub-second lifetime, so a positive value that truncates to 0 clamps up to the 1s floor instead of being rejected. Values >= 1 still truncate toward zero so the envelope never claims more life than the upstream stated, and NaN / Infinity / oversized input still read as unparseable ("unspecified"). Regression covers the classifier (0.5 and 0.001 clamp to 1, 1.9 truncates to 1, -0.5 stays expired) and the mint (a 0.5s upstream lifetime mints a 200 envelope rather than a 502); reverting to the truncate-then-check reddens both. --- .../mcp_server/discoverable_endpoints.py | 21 ++++++++++------- .../mcp_server/test_discoverable_endpoints.py | 23 +++++++++++++++++++ 2 files changed, 36 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index b8d10335182..8d1713a5911 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -742,19 +742,24 @@ envelope caps it, the by-design behaviour for an upstream that omits the field." def _classify_upstream_lifetime(raw_expires_in: object) -> "int | Literal['unspecified', 'expired']": """Classify an upstream ``expires_in`` into a positive number of seconds, ``"unspecified"`` (absent - or unparseable, so the envelope caps it), or ``"expired"`` (a parseable non-positive value the - upstream reports as already elapsed). Telling "we do not know the lifetime" apart from "the upstream - says it is already dead" is what stops an explicitly-expired token from silently receiving the - envelope's 1h cap. ``bool`` is excluded (an ``int`` subclass but never a real lifetime), and - ``int(float(...))`` can raise on ``NaN`` / ``Infinity`` / oversized input, which reads as - unparseable rather than surfacing as a 500.""" + or unparseable, so the envelope caps it), or ``"expired"`` (a non-positive value the upstream reports + as already elapsed). Telling "we do not know the lifetime" apart from "the upstream says it is + already dead" is what stops an explicitly-expired token from silently receiving the envelope's 1h + cap. The expired decision is made on the parsed numeric value, not on ``int(...)`` of it, so a + positive sub-second lifetime in ``(0, 1)`` is not truncated to ``0`` and misread as elapsed; the + envelope works in whole seconds, so such a lifetime clamps up to its 1s floor. ``bool`` is excluded + (an ``int`` subclass but never a real lifetime), and the conversions can raise on ``NaN`` / + ``Infinity`` / oversized input, which reads as unparseable rather than surfacing as a 500.""" if raw_expires_in is None or isinstance(raw_expires_in, bool) or not isinstance(raw_expires_in, (int, float, str)): return "unspecified" try: - seconds = int(float(raw_expires_in)) + numeric = float(raw_expires_in) + seconds = int(numeric) except (ValueError, TypeError, OverflowError): return "unspecified" - return seconds if seconds > 0 else "expired" + if numeric <= 0: + return "expired" + return max(1, seconds) def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenGrant | _UpstreamGrantRejection": diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f611d2f6a24..68466e624ec 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4661,6 +4661,22 @@ async def test_bridge_mint_upstream_expired_lifetime_is_502(): assert json.loads(response.body)["error"] == "server_error" +@pytest.mark.asyncio +async def test_bridge_mint_positive_sub_second_lifetime_mints_not_502(): + """A positive fractional expires_in in (0, 1) is a live token, not an elapsed one, so it mints a + (1s-floored) envelope rather than being truncated to 0 and rejected with 502 after the single-use + code was already consumed. Regression for classifying a sub-second remaining lifetime as expired.""" + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + upstream = {"access_token": "UP", "token_type": "Bearer", "expires_in": 0.5} + response = await _exchange_for_bridge_server(server, upstream, key_hash="hashed-litellm-key-77") + assert response.status_code == 200 + body = json.loads(response.body) + assert body["access_token"].startswith("llm_env_") + assert body["expires_in"] >= 0 + + @pytest.mark.asyncio async def test_bridge_mint_unknown_lifetime_is_capped_not_rejected(): """An absent or unparseable expires_in leaves the lifetime unknown, which the envelope caps (never @@ -4742,6 +4758,13 @@ def test_classify_upstream_lifetime(): # explicit, parseable, non-positive -> the upstream says the token is already dead assert _classify_upstream_lifetime(0) == "expired" assert _classify_upstream_lifetime(-5) == "expired" + assert _classify_upstream_lifetime(-0.5) == "expired" + # a positive sub-second lifetime is alive, not elapsed; it clamps up to the envelope's 1s floor + # rather than truncating to 0 and being misread as expired + assert _classify_upstream_lifetime(0.5) == 1 + assert _classify_upstream_lifetime(0.001) == 1 + # a positive value >= 1 truncates toward zero (never overstating the stated lifetime) + assert _classify_upstream_lifetime(1.9) == 1 # unknown lifetime -> cap (never invent a longer life than the upstream stated) assert _classify_upstream_lifetime(None) == "unspecified" assert _classify_upstream_lifetime(True) == "unspecified" From 0e90f61e48a84ae8fd7610333032a40b22e6c928 Mon Sep 17 00:00:00 2001 From: Abhimanyu Kapur <38531241+akapur99@users.noreply.github.com> Date: Sat, 11 Jul 2026 19:21:51 -0700 Subject: [PATCH 048/123] fix(auto_router): inline error for missing LLM classifier model Selecting the LLM classifier without picking a model only surfaced a toast on submit; the classifier model select now gets the same red outline and helper text as the tier and embedding selects once a submit attempt has failed. --- .../add_model/ComplexityRouterConfig.test.tsx | 22 +++++++++++++++++++ .../add_model/ComplexityRouterConfig.tsx | 9 ++++++++ 2 files changed, 31 insertions(+) diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index 0613b0c02ae..a34a8709918 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -226,6 +226,28 @@ describe("ComplexityRouterConfig", () => { expect(screen.queryByText("This tier is required")).not.toBeInTheDocument(); }); + it("shows an inline error on the classifier model select when llm is selected without a model", () => { + const llmValue: ComplexityRouterConfigValue = { + ...defaultValue, + classifier_type: "llm", + classifier_llm_config: { model: "", timeout_ms: 3000 }, + }; + renderWithProviders(); + fireEvent.click(screen.getByText("Advanced: Classification Method")); + expect(screen.getByText("A classifier model is required")).toBeInTheDocument(); + }); + + it("does not show the classifier model error once a classifier model is set", () => { + const llmValue: ComplexityRouterConfigValue = { + ...defaultValue, + classifier_type: "llm", + classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, + }; + renderWithProviders(); + fireEvent.click(screen.getByText("Advanced: Classification Method")); + expect(screen.queryByText("A classifier model is required")).not.toBeInTheDocument(); + }); + it("shows a validation error only under unfilled tiers when showValidationErrors is true", () => { renderWithProviders( = ({ label: model.model_group, })); + const classifierModelMissing = + showValidationErrors && value.classifier_type === "llm" && !value.classifier_llm_config?.model; + const handleTierChange = (tier: keyof ComplexityTiers, model: string) => { onChange({ ...value, @@ -233,7 +236,13 @@ const ComplexityRouterConfig: React.FC = ({ showSearch style={{ width: "100%" }} options={modelOptions} + status={classifierModelMissing ? "error" : undefined} /> + {classifierModelMissing && ( + + A classifier model is required + + )}
From c9beaf85ff9713371ff5a20154a71429aeac8017 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 19:25:39 -0700 Subject: [PATCH 049/123] build(dev-env): add make bootstrap and unprovisioned-checkout preflight to pre-commit --- CLAUDE.md | 2 ++ Makefile | 15 ++++++++++++++- scripts/pre_commit_lint.sh | 23 +++++++++++++++++++++-- 3 files changed, 37 insertions(+), 3 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 78da2c65d96..0c679e92113 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -39,6 +39,8 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a Python max line length is 120, not 88 +On a fresh worktree or clone, run `make bootstrap` before anything else. It provisions everything tests, `make pre-commit`, and a local proxy need: the uv env with proxy extras, the Prisma client, and the dashboard's node_modules; on worktrees it also copies `.env` from the main checkout (it never overwrites an existing `.env`) + Run tests before you commit. Also, run `make pre-commit` right before each commit, which generates types (as needed) and formats/lints your code. Any errors found must be fixed. It only runs when there are staged frontend and/or backend changes and calculates violations, generates types, etc. based on the worktree, so stage what you need or stash/delete unwanted files in litellm/ or ui/ (where backend and frontend lint run, respectively) before running it. If it fails because dashboard api types are stale, it already regenerated them for you. You just need to stage the schema.d.ts, re-run `make pre-commit` to confirm it passes, and commit When you fix violations gated by `ruff-strict-budget.json`, `type-discipline-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered limits so the ceilings ratchet down instead of leaving stale headroom. It measures the working tree, so it must contain exactly the fixes you're committing diff --git a/Makefile b/Makefile index 965a3254616..d035f1703bd 100644 --- a/Makefile +++ b/Makefile @@ -9,11 +9,12 @@ lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \ install-dev install-proxy-dev install-test-deps install-hooks \ install-helm-unittest check-circular-imports check-import-safety pre-commit \ - lint-install lint-fetch-base + lint-install lint-fetch-base bootstrap # Default target help: @echo "Available commands:" + @echo " make bootstrap - Provision a fresh clone/worktree: Python env with proxy extras, Prisma client, dashboard node_modules; worktrees also copy .env from the main checkout" @echo " make install-dev - Install development dependencies" @echo " make install-proxy-dev - Install proxy development dependencies" @echo " make install-dev-ci - Install dev dependencies (CI-compatible, pins OpenAI)" @@ -71,6 +72,18 @@ info: install-dev: $(UV) sync --inexact --frozen +bootstrap: + $(UV) sync --inexact --frozen --extra proxy --group proxy-dev --group e2e-dev + $(UV_RUN) python scripts/prisma_generate_if_needed.py + cd ui/litellm-dashboard && npm ci --no-audit --no-fund + @main_root=$$(git worktree list --porcelain | head -1 | sed 's/^worktree //'); \ + if [ "$$main_root" != "$$(git rev-parse --show-toplevel)" ] && [ -f "$$main_root/.env" ] && [ ! -f .env ]; then \ + cp "$$main_root/.env" .env && echo "bootstrap: copied .env from $$main_root"; \ + else \ + echo "bootstrap: .env left untouched"; \ + fi + @echo "bootstrap: done" + install-proxy-dev: $(UV) sync --frozen --group proxy-dev --extra proxy diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index d7d560ce947..cce0cb61c1e 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -89,6 +89,11 @@ EOF status=0 +bootstrap_hint() { + echo " This checkout looks unprovisioned (fresh worktree or clone)." >&2 + echo " Fix: make bootstrap" >&2 +} + if [ -n "$litellm_py_files" ]; then echo "pre-commit: linting Python (make lint)" make lint || { echo "✗ Python lint failed. Fix the reds above, then re-run make pre-commit." >&2; status=1; } @@ -109,7 +114,13 @@ fi if [ -n "$ui_prettier_files" ] || [ -n "$ui_eslint_files" ]; then echo "pre-commit: linting dashboard (prettier + eslint + lint budgets)" - lint_dashboard || { echo "✗ Dashboard lint failed. See above; format with: (cd ui/litellm-dashboard && npm run format)." >&2; status=1; } + if [ ! -d ui/litellm-dashboard/node_modules ]; then + echo "✗ ui/litellm-dashboard/node_modules is missing; dashboard lint cannot run." >&2 + bootstrap_hint + status=1 + else + lint_dashboard || { echo "✗ Dashboard lint failed. See above; format with: (cd ui/litellm-dashboard && npm run format)." >&2; status=1; } + fi fi if [ -n "$spec_files" ]; then @@ -118,7 +129,15 @@ if [ -n "$spec_files" ]; then # and an up-to-date Prisma client; check-ui-api-types.yml installs those and runs # prisma generate before gen:api, so mirror that here or a stale client can mask # drift that CI will still flag. - if ! uv run --no-sync python scripts/prisma_generate_if_needed.py; then + if [ ! -d ui/litellm-dashboard/node_modules ]; then + echo "✗ ui/litellm-dashboard/node_modules is missing; the gen:api sync check cannot run." >&2 + bootstrap_hint + status=1 + elif ! uv run --no-sync python -c "import orjson, prisma" 2>/dev/null; then + echo "✗ The Python env lacks the proxy deps (orjson/prisma) that gen:api needs." >&2 + bootstrap_hint + status=1 + elif ! uv run --no-sync python scripts/prisma_generate_if_needed.py; then echo "✗ Could not regenerate Prisma client (prisma generate failed)." >&2 status=1 elif ( cd ui/litellm-dashboard && LITELLM_PYTHON="uv run --no-sync python" npm run gen:api ); then From f717e3b2f0d4914ee5311f058a2514d676306988 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 11 Jul 2026 19:27:11 -0700 Subject: [PATCH 050/123] feat(router): random-pick multi-model complexity tiers (#32967) * feat(router): random-pick multi-model complexity tiers Tier pools already make sense without adaptive; stop pinning lists to index 0 and shuffle within the classified tier instead. Co-authored-by: Cursor * fix(ci): format complexity router config Co-authored-by: Cursor * fix(ci): use PEP 585 types for tier pools Co-authored-by: Cursor --------- Co-authored-by: Cursor --- .../complexity_router/complexity_router.py | 23 ++++++++------ .../complexity_router/config.py | 25 +++++++++++++--- .../router_strategy/test_complexity_router.py | 30 +++++++++++++++++++ 3 files changed, 65 insertions(+), 13 deletions(-) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 11719b8a18f..2138a0112a0 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -14,6 +14,7 @@ Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter """ import asyncio +import random import re from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast @@ -437,22 +438,26 @@ class ComplexityRouter(CustomLogger): """ tier_key = tier.value if isinstance(tier, ComplexityTier) else tier - # Check config tiers mapping - model = self.config.tiers.get(tier_key) - if model: - return model + if tier_key in self.config.tiers: + return self._pick_from_tier_value(self.config.tiers[tier_key], tier_key) - # Fallback to default model if configured if self.config.default_model: return self.config.default_model - # Last resort: return MEDIUM tier model or error - medium_model = self.config.tiers.get(ComplexityTier.MEDIUM.value) - if medium_model: - return medium_model + medium_key = ComplexityTier.MEDIUM.value + if medium_key in self.config.tiers: + return self._pick_from_tier_value(self.config.tiers[medium_key], medium_key) raise ValueError(f"No model configured for tier {tier_key} and no default_model set") + @staticmethod + def _pick_from_tier_value(model: str | list[str], tier_key: str) -> str: + if isinstance(model, str): + return model + if not model: + raise ValueError(f"Empty model pool for tier {tier_key}") + return random.choice(model) + def _lexical_tier_override(self, user_message: str) -> Optional[ComplexityTier]: """When keyword_tier_rules match literally, the most-severe matched tier wins. diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 125de6f7489..8c8e5acb51f 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -8,7 +8,7 @@ All values are configurable via proxy config.yaml. from enum import Enum from typing import Dict, List, Literal, Optional -from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator class ComplexityTier(str, Enum): @@ -244,10 +244,12 @@ class ClassifierLLMConfig(BaseModel): class ComplexityRouterConfig(BaseModel): """Configuration for the ComplexityRouter.""" - # Tier to model mapping - tiers: Dict[str, str] = Field( + # string = pin; list = random pick from the tier pool + tiers: dict[str, str | list[str]] = Field( default_factory=lambda: DEFAULT_TIER_MODELS.copy(), - description="Mapping of complexity tiers to model names", + description=( + "Mapping of complexity tiers to a model or model pool. A list is randomly picked from for that tier" + ), ) # Tier boundaries (normalized scores) @@ -335,6 +337,21 @@ class ComplexityRouterConfig(BaseModel): model_config = ConfigDict(extra="allow") # Allow additional fields + @field_validator("tiers", mode="before") + @classmethod + def _coerce_tier_values(cls, value: object) -> object: + if not isinstance(value, dict): + return value + coerced: dict[str, object] = {} + for key, item in value.items(): + if isinstance(item, str): + coerced[key] = item + elif isinstance(item, (list, tuple)): + coerced[key] = list(item) + else: + coerced[key] = item + return coerced + @model_validator(mode="after") def _validate_llm_classifier_config(self) -> "ComplexityRouterConfig": if self.classifier_type == "llm" and self.classifier_llm_config is None: diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index f47c19b2baa..e1133620a57 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -315,6 +315,36 @@ class TestModelSelection: model = router.get_model_for_tier(ComplexityTier.SIMPLE) assert model == "fallback-model" + def test_get_model_for_tier_list_random_choice(self, mock_router_instance): + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + "tiers": {"SIMPLE": ["cheap", "premium"], "MEDIUM": "mid"}, + "default_model": "mid", + }, + ) + pool = ["cheap", "premium"] + with patch( + "litellm.router_strategy.complexity_router.complexity_router.random.choice", + return_value="premium", + ) as choice: + assert router.get_model_for_tier(ComplexityTier.SIMPLE) == "premium" + choice.assert_called_once_with(pool) + assert router.get_model_for_tier(ComplexityTier.MEDIUM) == "mid" + + def test_get_model_for_tier_empty_pool_raises(self, mock_router_instance): + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + "tiers": {"SIMPLE": []}, + "default_model": "mid", + }, + ) + with pytest.raises(ValueError, match="Empty model pool for tier SIMPLE"): + router.get_model_for_tier(ComplexityTier.SIMPLE) + class TestPreRoutingHook: """Test the async_pre_routing_hook method.""" From f61fd2fb6d5dfd07850cb0ae1a486a6b9adcb18b Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Sat, 11 Jul 2026 19:46:02 -0700 Subject: [PATCH 051/123] fix(xecguard): sanitize scan result before recording it for logging (#32935) --- .../guardrail_hooks/xecguard/xecguard.py | 11 +++++++- .../guardrail_hooks/test_xecguard.py | 27 +++++++++++++++++++ 2 files changed, 37 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index f4a6f0aeb3b..7fe942bcb38 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -44,6 +44,8 @@ from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, ) +from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys +from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -64,6 +66,13 @@ if TYPE_CHECKING: ) +def _sanitize_scan_result_for_logging(scan_result: dict) -> dict: + without_secrets = {key: value for key, value in scan_result.items() if key != "secret_fields"} + redacted = redact_nested_match_and_regex_keys(without_secrets) + masked = mask_credentials_in_payload(redacted if isinstance(redacted, dict) else without_secrets) + return masked if isinstance(masked, dict) else without_secrets + + _DEFAULT_API_BASE = "https://api-xecguard.cycraft.ai" _SCAN_ENDPOINT = "/xecguard/v1/scan" _GROUNDING_ENDPOINT = "/xecguard/v1/grounding" @@ -253,7 +262,7 @@ class XecGuardGuardrail(CustomGuardrail): slg = StandardLoggingGuardrailInformation( guardrail_name=self.guardrail_name or "xecguard", guardrail_mode=GuardrailEventHooks.logging_only, - guardrail_response=scan_result, + guardrail_response=_sanitize_scan_result_for_logging(scan_result), guardrail_status=guardrail_status, start_time=start_time.timestamp(), end_time=end_time.timestamp(), diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py index f64967abbb0..6e601df897b 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_xecguard.py @@ -1675,6 +1675,33 @@ class TestXecGuardLoggingHook: assert info_list[1]["guardrail_name"] == "test-xecguard" assert info_list[1]["guardrail_response"]["trace_id"] == "lg-4" + @pytest.mark.asyncio + async def test_async_logging_hook_sanitizes_scan_result( + self, xecguard_guardrail, mock_request_data + ): + resp = _make_response( + { + "decision": "SAFE", + "trace_id": "lg-5", + "secret_fields": {"authorization": "Bearer xgs_raw"}, + "detections": [{"match": "raw matched span", "policy": "pii"}], + "api_key": "xgs_super_secret_value", + } + ) + with patch.object(xecguard_guardrail.async_handler, "post", return_value=resp): + kwargs = {**mock_request_data, "standard_logging_object": {}} + await xecguard_guardrail.async_logging_hook( + kwargs=kwargs, + result=_build_model_response("some answer"), + call_type="acompletion", + ) + info = kwargs["standard_logging_object"]["guardrail_information"][0] + guardrail_response = info["guardrail_response"] + assert "secret_fields" not in guardrail_response + assert guardrail_response["detections"][0]["match"] == "[REDACTED]" + assert guardrail_response["api_key"] != "xgs_super_secret_value" + assert guardrail_response["trace_id"] == "lg-5" + @pytest.mark.asyncio async def test_async_logging_hook_without_response_records_info( self, xecguard_guardrail, mock_request_data From 1bff68c0ced29c926212e33e260d9ccdcb76a88a Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 20:23:41 -0700 Subject: [PATCH 052/123] chore: keep it brief --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index 0c679e92113..9f708716c6d 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -39,7 +39,7 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a Python max line length is 120, not 88 -On a fresh worktree or clone, run `make bootstrap` before anything else. It provisions everything tests, `make pre-commit`, and a local proxy need: the uv env with proxy extras, the Prisma client, and the dashboard's node_modules; on worktrees it also copies `.env` from the main checkout (it never overwrites an existing `.env`) +On a fresh worktree or clone, run `make bootstrap` before anything else. It provisions everything tests, `make pre-commit`, and a local proxy need Run tests before you commit. Also, run `make pre-commit` right before each commit, which generates types (as needed) and formats/lints your code. Any errors found must be fixed. It only runs when there are staged frontend and/or backend changes and calculates violations, generates types, etc. based on the worktree, so stage what you need or stash/delete unwanted files in litellm/ or ui/ (where backend and frontend lint run, respectively) before running it. If it fails because dashboard api types are stale, it already regenerated them for you. You just need to stage the schema.d.ts, re-run `make pre-commit` to confirm it passes, and commit From 732c382644f714f3e7f5be1d454d351e9cc17a5b Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 20:25:53 -0700 Subject: [PATCH 053/123] chore: keep it brief --- Makefile | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Makefile b/Makefile index d035f1703bd..8b657dcb465 100644 --- a/Makefile +++ b/Makefile @@ -14,7 +14,7 @@ # Default target help: @echo "Available commands:" - @echo " make bootstrap - Provision a fresh clone/worktree: Python env with proxy extras, Prisma client, dashboard node_modules; worktrees also copy .env from the main checkout" + @echo " make bootstrap - Provision a fresh clone/worktree" @echo " make install-dev - Install development dependencies" @echo " make install-proxy-dev - Install proxy development dependencies" @echo " make install-dev-ci - Install dev dependencies (CI-compatible, pins OpenAI)" From 6401908f65b8bd2ed18f33a38fab843ee5fea184 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 20:29:16 -0700 Subject: [PATCH 054/123] docs(readme): point developer-mode setup at make bootstrap --- README.md | 13 ++++--------- 1 file changed, 4 insertions(+), 9 deletions(-) diff --git a/README.md b/README.md index 90d3e944fcc..0e6038a9e4b 100644 --- a/README.md +++ b/README.md @@ -552,17 +552,12 @@ The Terraform modules live at [`terraform/litellm/aws/`](./terraform/litellm/aws 2. Run dependent services `docker-compose up db prometheus` #### Backend -1. (In root) create virtual environment `python -m venv .venv` -2. Activate virtual environment `source .venv/bin/activate` -3. Install dependencies `uv sync --all-extras --group proxy-dev` -4. `uv run prisma generate` -5. `prisma generate` -6. Start proxy backend `python litellm/proxy/proxy_cli.py` +1. (In root) provision the checkout with `make bootstrap` (installs the Python env with proxy extras, generates the Prisma client, and installs the dashboard's node_modules) +2. Start proxy backend `uv run python litellm/proxy/proxy_cli.py` #### Frontend -1. Navigate to `ui/litellm-dashboard` -2. Install dependencies `npm install` -3. Run `npm run dev` to start the dashboard +1. Navigate to `ui/litellm-dashboard` (dependencies were already installed by `make bootstrap`) +2. Run `npm run dev` to start the dashboard ### Verify Docker Image Signatures From a523895a573f0ed7867e15b8f1fa6e975d690a25 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 11 Jul 2026 20:32:34 -0700 Subject: [PATCH 055/123] chore: keep it concise --- README.md | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/README.md b/README.md index 0e6038a9e4b..32b0160dbaa 100644 --- a/README.md +++ b/README.md @@ -552,12 +552,12 @@ The Terraform modules live at [`terraform/litellm/aws/`](./terraform/litellm/aws 2. Run dependent services `docker-compose up db prometheus` #### Backend -1. (In root) provision the checkout with `make bootstrap` (installs the Python env with proxy extras, generates the Prisma client, and installs the dashboard's node_modules) -2. Start proxy backend `uv run python litellm/proxy/proxy_cli.py` +1. Run `make bootstrap` +2. Start proxy backend: `uv run python litellm/proxy/proxy_cli.py` #### Frontend -1. Navigate to `ui/litellm-dashboard` (dependencies were already installed by `make bootstrap`) -2. Run `npm run dev` to start the dashboard +1. Navigate to `ui/litellm-dashboard` (dependencies were already installed w/ `make bootstrap`) +2. Start dashboard: `npm run dev` ### Verify Docker Image Signatures From 85f9bdd4129588cdf47c746978fc4b658faec747 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 11 Jul 2026 21:38:18 -0700 Subject: [PATCH 056/123] feat(router): add Router(plugins=[...]) routing-plugin pipeline (#32972) * feat(router): add Router(plugins=[...]) routing-plugin pipeline Runs a sequence of user-supplied plugins before the routing decision is made. Each plugin reads/mutates a RoutingContext (messages, candidate models, metadata, signals); the narrowed candidate list is enforced when picking a deployment, raising rather than silently falling back if a plugin narrows to zero candidates. Prototype for the routing-plugin pipeline discussed in #32168. * fix(router): use ruff-modern typing, add raw/structured messages to RoutingContext - Use dict/list/X|None instead of Dict/List/Optional in new code, staying within the ruff strict-rule budget ratchet - Extract the guardrail-translation message normalization ComplexityRouter already had into a shared resolve_structured_messages() helper (litellm_core_utils/prompt_templates/factory.py), reused by ComplexityRouter and the new routing-plugin pipeline instead of duplicating it - RoutingContext now exposes both raw_messages (as received) and structured_messages (normalized across chat completions / Anthropic messages / Responses API), mirroring CustomGuardrail.apply_guardrail's pattern, per review feedback on #32972 - Add direct unit tests for _run_routing_plugins and _filter_by_routing_plugin_candidates (router_code_coverage gate requires every router.py function be called by name somewhere in tests/) * fix(test): rename to test_router_routing_plugins.py router_code_coverage.py's AST scanner only inspects test files whose filename contains the substring "router" -- test_routing_plugins.py doesn't match (routing != router), so it silently skipped this file and flagged _run_routing_plugins/_filter_by_routing_plugin_candidates as untested despite the direct unit tests added for them. * fix(router): fail closed when plugins are configured but the resolved routing path can't run them Router.completion() (and other sync entry points) resolves deployments via the synchronous get_available_deployment(), which never runs async_pre_routing_hook and therefore never runs the routing-plugin pipeline. async_get_available_deployment() itself falls back to that same synchronous method for routing strategies without an async-native selector (e.g. legacy "usage-based-routing" v1). Both paths would let a policy plugin (e.g. a deny-all rule) be silently bypassed. Raise instead of silently proceeding when self.routing_plugins is configured and the sync path is reached, since applying the pipeline to every selector path is a larger change out of scope for this PR. Per review: https://github.com/BerriAI/litellm/pull/32972/changes/BASE..bdfb583c2c6f8df10004fb249e11629d41ce71fa#r3565373303 --- .../prompt_templates/factory.py | 53 +++++ litellm/router.py | 101 ++++++++ .../complexity_router/complexity_router.py | 38 +-- litellm/types/router.py | 31 ++- .../test_router_routing_plugins.py | 220 ++++++++++++++++++ 5 files changed, 407 insertions(+), 36 deletions(-) create mode 100644 tests/test_litellm/router_strategy/test_router_routing_plugins.py diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 8bb0e12905e..f7ff4d6b16f 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -5494,3 +5494,56 @@ def has_tool_with_name(tools: Any, tool_name: str) -> bool: elif tool.get("name") == tool_name: return True return False + + +def resolve_structured_messages( + messages: list[dict[str, Any]] | None, + request_kwargs: dict[str, Any], +) -> list[dict[str, Any]] | None: + """ + Normalize a request's messages to OpenAI-spec chat-completions shape, + regardless of which API surface produced them (chat completions, + Anthropic /v1/messages, Responses API ``input``, etc). + + Returns ``messages`` unchanged if already present. Otherwise dispatches + through the guardrail translation handlers (the same per-surface + conversion logic guardrails use) to convert e.g. Responses API ``input`` + into a message list. Returns ``None`` if no messages could be resolved. + """ + if messages: + return messages + + from litellm.litellm_core_utils.api_route_to_call_types import ( + get_call_types_for_route, + ) + from litellm.llms import load_guardrail_translation_mappings + from litellm.types.utils import CallTypes + + mappings = load_guardrail_translation_mappings() + call_type: CallTypes | None = None + + # 1. Try route-based inference from proxy metadata + route = request_kwargs.get("litellm_metadata", {}).get("user_api_key_request_route") + if route: + call_types_list = get_call_types_for_route(route) + if call_types_list: + for ct in call_types_list: + if ct in mappings: + call_type = ct + break + + # 2. Fallback: try each mapped handler until one produces messages + handlers_to_try: list[Any] = [] + if call_type is not None and call_type in mappings: + handlers_to_try.append(mappings[call_type]()) + else: + handlers_to_try.extend(handler_cls() for handler_cls in mappings.values()) + + for handler in handlers_to_try: + structured = handler.get_structured_messages(request_kwargs) + if structured: + return [ + msg if isinstance(msg, dict) else msg.model_dump() # type: ignore + for msg in structured + ] + return None diff --git a/litellm/router.py b/litellm/router.py index 6e773a06c7f..245a50545e7 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -179,7 +179,9 @@ from litellm.types.router import ( RouterModelGroupAliasItem, RouterRateLimitError, RouterRateLimitErrorBasic, + RoutingContext, RoutingGroup, + RoutingPlugin, RoutingStrategy, SearchToolTypedDict, ) @@ -299,6 +301,7 @@ class Router: enable_pre_call_checks: bool = False, enable_tag_filtering: bool = False, tag_filtering_match_any: bool = True, + plugins: list[RoutingPlugin] | None = None, retry_after: int = 0, # min time to wait before retrying a failed request retry_policy: Optional[Union[RetryPolicy, dict]] = None, # set custom retries for different exceptions model_group_retry_policy: Dict[str, RetryPolicy] = {}, # set custom retry policies based on model group @@ -477,6 +480,7 @@ class Router: self.complexity_routers: Dict[str, "ComplexityRouter"] = {} self.adaptive_routers: Dict[str, "AdaptiveRouter"] = {} self.quality_routers: Dict[str, "QualityRouter"] = {} + self.routing_plugins: list[RoutingPlugin] = list(plugins) if plugins else [] # Initialize model_group_alias early since it's used in set_model_list self.model_group_alias: Dict[str, Union[str, RouterModelGroupAliasItem]] = ( @@ -10321,6 +10325,12 @@ class Router: metadata_variable_name=self._get_metadata_variable_name_from_kwargs(request_kwargs), ) + # narrow to whatever `self.routing_plugins` left in candidate_models + healthy_deployments = self._filter_by_routing_plugin_candidates( + healthy_deployments=healthy_deployments, + request_kwargs=request_kwargs, + ) + ## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2) _target_order = (request_kwargs or {}).pop("_target_order", None) healthy_deployments = litellm.utils._get_order_filtered_deployments( @@ -10596,6 +10606,76 @@ class Router: ) raise e + async def _run_routing_plugins( + self, + model: str, + request_kwargs: dict, + messages: list[dict[str, Any]] | None, + ) -> RoutingContext: + """ + Build a RoutingContext for `model`, run it through `self.routing_plugins` + in order, then stash the narrowed candidate list and accumulated signals + onto `request_kwargs["metadata"]` so `_filter_by_routing_plugin_candidates` + (called later, during healthy-deployment filtering) can consume them. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + resolve_structured_messages, + ) + + deployments = self.get_model_list(model_name=model) or [] + candidate_models = [ + d["litellm_params"]["model"] for d in deployments if d.get("litellm_params", {}).get("model") + ] + + metadata_key = self._get_metadata_variable_name_from_kwargs(request_kwargs) + metadata = request_kwargs.setdefault(metadata_key, {}) + + context = RoutingContext( + raw_messages=messages or [], + structured_messages=resolve_structured_messages(messages=messages, request_kwargs=request_kwargs) or [], + candidate_models=candidate_models, + metadata=metadata, + ) + + for plugin in self.routing_plugins: + context = await plugin.run(context) + + metadata["routing_plugin_signals"] = context.signals + if len(context.candidate_models) < len(candidate_models): + metadata["_routing_plugin_candidate_models"] = context.candidate_models + + return context + + def _filter_by_routing_plugin_candidates( + self, + healthy_deployments: Union[list[dict], dict], + request_kwargs: dict, + ) -> Union[list[dict], dict]: + """ + Narrow `healthy_deployments` to whatever `self.routing_plugins` left in + `context.candidate_models`. Raises rather than silently falling back to + the unfiltered pool -- a plugin narrowing to nothing is a policy decision + (e.g. no model this tenant's budget allows), not something to bypass. + """ + if not self.routing_plugins or not isinstance(healthy_deployments, list): + return healthy_deployments + + metadata_key = self._get_metadata_variable_name_from_kwargs(request_kwargs) + candidate_models = (request_kwargs.get(metadata_key) or {}).get("_routing_plugin_candidate_models") + # `is None` (not falsy-check): a plugin narrowing to an empty list must + # still hit the "no deployments left" raise below, not be treated the + # same as "no plugin ever set this key". + if candidate_models is None: + return healthy_deployments + + candidate_set = set(candidate_models) + filtered = [d for d in healthy_deployments if d.get("litellm_params", {}).get("model") in candidate_set] + + if not filtered: + raise ValueError(f"No deployments left after routing-plugin filtering. candidate_models={candidate_models}") + + return filtered + async def async_pre_routing_hook( self, model: str, @@ -10609,6 +10689,15 @@ class Router: Used for the litellm auto-router to modify the request before the routing decision is made. """ + ######################################################### + # Run the routing-plugin pipeline, if any plugins are configured. + # Plugins narrow the candidate deployment pool (consumed later by + # `_filter_by_routing_plugin_candidates`) and may attach signals for + # downstream strategies (auto-router, complexity-router, ...) to read. + ######################################################### + if self.routing_plugins: + await self._run_routing_plugins(model=model, request_kwargs=request_kwargs, messages=messages) + ######################################################### # Check if any auto-router should be used ######################################################### @@ -10671,6 +10760,18 @@ class Router: """ Returns the deployment based on routing strategy """ + if self.routing_plugins: + raise ValueError( + "Router(plugins=[...]) is configured, but this call resolved to the synchronous " + "deployment-selection path, which never runs the routing-plugin pipeline. This " + "happens for sync Router methods (e.g. Router.completion()) and for async calls " + "with a routing_strategy that has no async-native selector (e.g. legacy " + "'usage-based-routing', v1). Silently skipping " + "configured plugins would let a policy plugin (e.g. a deny-all rule) be bypassed. " + "Use an async Router method with a supported routing_strategy (simple-shuffle, " + "usage-based-routing-v2, cost-based-routing, latency-based-routing, least-busy), " + "or remove `plugins` from the Router config." + ) # users need to explicitly call a specific deployment, by setting `specific_deployment = True` as completion()/embedding() kwarg # When this was no explicit we had several issues with fallbacks timing out diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 2138a0112a0..74644f01be8 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -601,43 +601,11 @@ class ComplexityRouter(CustomLogger): Uses the guardrail translation handler dispatch to convert Responses API ``input`` (or other non-chat-completions formats) into OpenAI-spec messages. """ - if messages: - return messages - - from litellm.litellm_core_utils.api_route_to_call_types import ( - get_call_types_for_route, + from litellm.litellm_core_utils.prompt_templates.factory import ( + resolve_structured_messages, ) - from litellm.llms import load_guardrail_translation_mappings - from litellm.types.utils import CallTypes - mappings = load_guardrail_translation_mappings() - call_type: Optional[CallTypes] = None - - # 1. Try route-based inference from proxy metadata - route = request_kwargs.get("litellm_metadata", {}).get("user_api_key_request_route") - if route: - call_types_list = get_call_types_for_route(route) - if call_types_list: - for ct in call_types_list: - if ct in mappings: - call_type = ct - break - - # 2. Fallback: try each mapped handler until one produces messages - handlers_to_try: List[Any] = [] - if call_type is not None and call_type in mappings: - handlers_to_try.append(mappings[call_type]()) - else: - handlers_to_try.extend(handler_cls() for handler_cls in mappings.values()) - - for handler in handlers_to_try: - structured = handler.get_structured_messages(request_kwargs) - if structured: - return [ - msg if isinstance(msg, dict) else msg.model_dump() # type: ignore - for msg in structured - ] - return None + return resolve_structured_messages(messages=messages, request_kwargs=request_kwargs) @staticmethod def _extract_user_message_and_system_prompt( diff --git a/litellm/types/router.py b/litellm/types/router.py index 4bac9358392..3bedd97c20c 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -9,7 +9,7 @@ from typing import Any, Dict, List, Literal, Optional, Tuple, Union, get_type_hi import httpx from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator -from typing_extensions import Required, TypedDict +from typing_extensions import Protocol, Required, TypedDict from litellm._uuid import uuid @@ -829,6 +829,35 @@ class PreRoutingHookResponse(BaseModel): messages: Optional[List[Dict[str, Any]]] +class RoutingContext(BaseModel): + """ + Passed through a Router's `plugins` pipeline before the routing decision is made. + + Each plugin reads and mutates this object; the next plugin sees the previous + plugin's changes. `candidate_models` narrows as the pipeline runs -- Router + only selects a deployment whose `litellm_params.model` survives the pipeline. + + `raw_messages` and `structured_messages` mirror the pattern + `CustomGuardrail.apply_guardrail` uses: the message shape differs by API + surface (chat completions, Anthropic /v1/messages, Responses API `input`, + ...), so plugins that need a stable, provider-agnostic shape should read + `structured_messages` (normalized to OpenAI chat-completions format); + plugins that need the exact original payload can read `raw_messages`. + """ + + raw_messages: list[dict[str, Any]] + structured_messages: list[dict[str, Any]] + candidate_models: list[str] + metadata: dict[str, Any] = Field(default_factory=dict) + signals: dict[str, Any] = Field(default_factory=dict) + + +class RoutingPlugin(Protocol): + """Interface a custom routing plugin must implement to run in `Router(plugins=[...])`.""" + + async def run(self, context: RoutingContext) -> RoutingContext: ... + + class RequestType(str, enum.Enum): """Fixed v0 taxonomy. User-extensible types come in v1.""" diff --git a/tests/test_litellm/router_strategy/test_router_routing_plugins.py b/tests/test_litellm/router_strategy/test_router_routing_plugins.py new file mode 100644 index 00000000000..e9c12d009e2 --- /dev/null +++ b/tests/test_litellm/router_strategy/test_router_routing_plugins.py @@ -0,0 +1,220 @@ +""" +Tests for Router(plugins=[...]) -- a pipeline of routing plugins that run +before the routing decision is made, narrowing the candidate deployment pool. + +Discussion: https://github.com/BerriAI/litellm/discussions/32168 +""" + +import pytest + +from litellm import Router +from litellm.types.router import RoutingContext + + +class LanguageDetector: + async def run(self, context: RoutingContext) -> RoutingContext: + context.signals["language-detector"] = {"lang": "en"} + return context + + +class DomainClassifier: + async def run(self, context: RoutingContext) -> RoutingContext: + context.signals["domain-classifier"] = {"domain": "coding", "confidence": 0.93} + return context + + +class TenantPolicy: + ALLOWED_PROVIDERS = {"acme-corp": {"openai", "anthropic"}} + + async def run(self, context: RoutingContext) -> RoutingContext: + tenant = context.metadata.get("tenant", "default") + allowed = self.ALLOWED_PROVIDERS.get(tenant, {"openai", "anthropic", "self-hosted"}) + context.candidate_models = [m for m in context.candidate_models if m.split("/")[0] in allowed] + context.signals["tenant-policy"] = {"tenant": tenant, "allowed_providers": sorted(allowed)} + return context + + +class BudgetPolicy: + COST_CAP_PER_TOKEN = 0.000005 + COST_BY_MODEL = { + "openai/gpt-4o-mini": 0.00000015, + "anthropic/claude-haiku-4-5": 0.000001, + "openai/gpt-5.1": 0.00003, + } + + async def run(self, context: RoutingContext) -> RoutingContext: + context.candidate_models = [ + m for m in context.candidate_models if self.COST_BY_MODEL.get(m, 0) <= self.COST_CAP_PER_TOKEN + ] + context.signals["budget-policy"] = {"daily_limit": 100} + return context + + +class BlockEverything: + async def run(self, context: RoutingContext) -> RoutingContext: + context.candidate_models = [] + return context + + +def _smart_router_model_list(): + return [ + { + "model_name": "smart-router", + "litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "cheap openai"}, + "model_info": {"tags": ["openai"]}, + }, + { + "model_name": "smart-router", + "litellm_params": {"model": "anthropic/claude-haiku-4-5", "mock_response": "anthropic"}, + "model_info": {"tags": ["anthropic"]}, + }, + { + "model_name": "smart-router", + "litellm_params": {"model": "openai/gpt-5.1", "mock_response": "expensive openai"}, + "model_info": {"tags": ["openai"]}, + }, + { + "model_name": "smart-router", + "litellm_params": {"model": "ollama/llama-3-70b", "mock_response": "self hosted"}, + "model_info": {"tags": ["self-hosted"]}, + }, + ] + + +@pytest.mark.asyncio +async def test_routing_plugin_pipeline_matches_jeann2013_e2e_scenario(): + """ + https://github.com/BerriAI/litellm/discussions/32168#discussioncomment-17608820 + + language plugin -> domain classifier -> tenant policy (openai+anthropic only) + -> budget policy (drops over-cap models) -> Router picks the best remaining + candidate. Must never land on the self-hosted or over-budget deployment. + """ + router = Router( + model_list=_smart_router_model_list(), + plugins=[LanguageDetector(), DomainClassifier(), TenantPolicy(), BudgetPolicy()], + ) + + response = await router.acompletion( + model="smart-router", + messages=[{"role": "user", "content": "Write a function to reverse a linked list."}], + metadata={"tenant": "acme-corp"}, + ) + + # response.model is the bare model name (litellm strips the provider/ prefix + # on the response), so compare against bare names rather than litellm_params.model + routed_model = response.model + + assert routed_model in {"gpt-4o-mini", "claude-haiku-4-5"} + assert routed_model not in {"llama-3-70b", "gpt-5.1"} + + +@pytest.mark.asyncio +async def test_routing_plugin_narrowing_to_zero_candidates_raises(): + """A plugin narrowing to nothing is a policy decision -- must raise, not silently + fall back to the unfiltered pool (that would defeat the policy it enforces).""" + router = Router( + model_list=_smart_router_model_list(), + plugins=[BlockEverything()], + ) + + with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"): + await router.acompletion( + model="smart-router", + messages=[{"role": "user", "content": "hi"}], + ) + + +def test_sync_get_available_deployment_rejects_configured_plugins(): + """ + Router.completion() (and any other sync entry point) resolves deployments via + the synchronous get_available_deployment(), which never runs the routing-plugin + pipeline. Silently allowing that would let a deny-all policy plugin be bypassed + just by calling the sync API -- must fail closed instead. + """ + router = Router(model_list=_smart_router_model_list(), plugins=[TenantPolicy()]) + + with pytest.raises(ValueError, match="routing-plugin pipeline"): + router.get_available_deployment(model="smart-router", messages=[{"role": "user", "content": "hi"}]) + + +def test_sync_router_completion_rejects_configured_plugins(): + """End-to-end: Router.completion() (the sync API) must not silently skip plugins either.""" + router = Router(model_list=_smart_router_model_list(), plugins=[TenantPolicy()]) + + with pytest.raises(ValueError, match="routing-plugin pipeline"): + router.completion(model="smart-router", messages=[{"role": "user", "content": "hi"}]) + + +@pytest.mark.asyncio +async def test_async_completion_with_unsupported_strategy_rejects_configured_plugins(): + """ + async_get_available_deployment() itself delegates to the synchronous selector + for routing strategies outside {simple-shuffle, usage-based-routing-v2, + cost-based-routing, latency-based-routing, least-busy} -- e.g. "usage-based-routing" + (v1, not v2) -- which would silently bypass the plugin pipeline on the async path too. + """ + router = Router( + model_list=_smart_router_model_list(), + plugins=[TenantPolicy()], + routing_strategy="usage-based-routing", + ) + + with pytest.raises(ValueError, match="routing-plugin pipeline"): + await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}]) + + +@pytest.mark.asyncio +async def test_router_without_plugins_is_unaffected(): + """Regression guard: a Router with no `plugins` configured behaves exactly as before.""" + router = Router( + model_list=[ + { + "model_name": "smart-router", + "litellm_params": {"model": "openai/gpt-4o-mini", "mock_response": "hi"}, + }, + ], + ) + response = await router.acompletion( + model="smart-router", + messages=[{"role": "user", "content": "hi"}], + ) + assert response.choices[0].message.content == "hi" + + +@pytest.mark.asyncio +async def test_run_routing_plugins_narrows_candidates_and_records_signals(): + """Unit-level check of _run_routing_plugins in isolation, independent of acompletion.""" + router = Router( + model_list=_smart_router_model_list(), + plugins=[LanguageDetector(), DomainClassifier(), TenantPolicy(), BudgetPolicy()], + ) + request_kwargs = {"metadata": {"tenant": "acme-corp"}} + + context = await router._run_routing_plugins( + model="smart-router", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "hi"}], + ) + + assert context.candidate_models == ["openai/gpt-4o-mini", "anthropic/claude-haiku-4-5"] + assert context.signals["domain-classifier"]["domain"] == "coding" + assert request_kwargs["metadata"]["_routing_plugin_candidate_models"] == context.candidate_models + + +def test_filter_by_routing_plugin_candidates_narrows_and_raises_when_empty(): + """Unit-level check of _filter_by_routing_plugin_candidates in isolation.""" + router = Router(model_list=_smart_router_model_list(), plugins=[TenantPolicy()]) + healthy_deployments = router.model_list + + narrowed = router._filter_by_routing_plugin_candidates( + healthy_deployments=healthy_deployments, + request_kwargs={"metadata": {"_routing_plugin_candidate_models": ["openai/gpt-4o-mini"]}}, + ) + assert [d["litellm_params"]["model"] for d in narrowed] == ["openai/gpt-4o-mini"] + + with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"): + router._filter_by_routing_plugin_candidates( + healthy_deployments=healthy_deployments, + request_kwargs={"metadata": {"_routing_plugin_candidate_models": ["nonexistent/model"]}}, + ) From 26ab730bfaac36b6d96af68d5fe5e7eb867af2ca Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 11 Jul 2026 21:56:33 -0700 Subject: [PATCH 057/123] feat(router): soft-floor adaptive mode for complexity router (#32947) * feat(router): soft-floor adaptive mode for complexity router Let complexity_router_config.adaptive=true Thompson-sample across the union of tier pools with a tier-distance penalty, and wire the existing adaptive post-call bandit so mis-tiered requests can still recover. Co-authored-by: Cursor * fix(router): reattach adaptive hooks for hybrid complexity Finalize was wiping every AdaptiveRouterPostCallHook and only re-registering standalone auto_router/adaptive_router deployments, so complexity adaptive=true never received bandit updates. Co-authored-by: Cursor * chore(router): drop unnecessary hybrid docstrings Co-authored-by: Cursor * fix(router): attribute adaptive feedback Credit user reactions to the model that produced the previous response while keeping current-response signals on the serving model Co-authored-by: Cursor * fix(router): tune hybrid cold defaults Use the cost-weighted policy that beat equal-pool complexity in the full bakeoff, and make the committed harness compare identical tier pools Co-authored-by: Cursor * fix(router): preserve hybrid cold quality floor Sample only unobserved models in the classified tier until feedback exists, then apply adaptive scoring without mis-penalizing models shared across tiers Co-authored-by: Cursor * fix(router): bound feedback context cache Cap retained session feedback so unique session IDs cannot exhaust router memory Co-authored-by: Cursor * fix(router): preserve exhaustion signals Include tool-result exhaustion in adaptive feedback and clear strict lint regressions blocking CI Co-authored-by: Cursor * refactor(router): remove stale owner cache Remove obsolete attribution state, tighten the embedded router type, and keep the test diff focused on adaptive behavior Co-authored-by: Cursor * refactor(router): centralize hook cleanup Use the callback manager to discover and remove adaptive hooks across every registered callback list Co-authored-by: Cursor --------- Co-authored-by: Cursor --- litellm/router.py | 29 +- .../router_strategy/adaptive_router/README.md | 15 +- .../adaptive_router/adaptive_router.py | 290 +++++++++++------- .../router_strategy/adaptive_router/hooks.py | 8 - .../adaptive_router/signals.py | 148 ++++++--- .../complexity_router/complexity_router.py | 271 +++++++++++++--- .../complexity_router/config.py | 87 ++++-- ruff-strict-budget.json | 8 +- .../adaptive_router/test_adaptive_router.py | 216 +++++++------ .../test_e2e_adaptive_router.py | 17 - .../adaptive_router/test_hooks.py | 37 ++- .../adaptive_router/test_state_endpoint.py | 25 +- .../router_strategy/test_complexity_router.py | 288 +++++++++++++++++ .../add_model/ComplexityRouterConfig.tsx | 2 +- .../build_complexity_router_config.ts | 3 +- 15 files changed, 1035 insertions(+), 409 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 245a50545e7..6539d3c0c43 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -7556,7 +7556,11 @@ class Router: if default_model is None and complexity_router_config: tiers = complexity_router_config.get("tiers", {}) # Use MEDIUM tier as fallback default - default_model = tiers.get("MEDIUM") or tiers.get("SIMPLE") + medium = tiers.get("MEDIUM") or tiers.get("SIMPLE") + if isinstance(medium, list): + default_model = medium[0] if medium else None + else: + default_model = medium if default_model is None: raise ValueError( @@ -7593,15 +7597,6 @@ class Router: AdaptiveRouterPostCallHook, ) - for _cb_list in ( - litellm.callbacks, - litellm.success_callback, - litellm.failure_callback, - litellm._async_success_callback, - litellm._async_failure_callback, - ): - litellm.logging_callback_manager.remove_callbacks_by_type(_cb_list, AdaptiveRouterPostCallHook) - for entry in self.model_list or []: lp = entry.get("litellm_params") if isinstance(entry, dict) else entry.litellm_params lp_model = (lp.get("model") if isinstance(lp, dict) else lp.model) if lp else None @@ -7619,6 +7614,20 @@ class Router: ) self.init_adaptive_router_deployment(deployment=deployment) + for model_name, complexity_router in self.complexity_routers.items(): + if not complexity_router.config.adaptive or model_name in self.adaptive_routers: + continue + adaptive_router = complexity_router._ensure_adaptive_router() + if adaptive_router is not None: + self.adaptive_routers[model_name] = adaptive_router + + for callback in litellm.logging_callback_manager.get_custom_loggers_for_type(AdaptiveRouterPostCallHook): + litellm.logging_callback_manager.remove_callback_from_all_lists(callback) + for adaptive_router in self.adaptive_routers.values(): + litellm.logging_callback_manager.add_litellm_callback( + AdaptiveRouterPostCallHook(adaptive_router=adaptive_router) + ) + def init_adaptive_router_deployment(self, deployment: Deployment) -> None: """ Build an AdaptiveRouter instance for this deployment and register its diff --git a/litellm/router_strategy/adaptive_router/README.md b/litellm/router_strategy/adaptive_router/README.md index 7f5d7aa21d0..09420a8dd9d 100644 --- a/litellm/router_strategy/adaptive_router/README.md +++ b/litellm/router_strategy/adaptive_router/README.md @@ -56,11 +56,10 @@ Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key - **Per-request decision.** Sample once per eligible model, score with `quality_weight·sample + cost_weight·normalized_cost`, pick the argmax. Routing is stateless per-turn — no sticky lookup. Each call resamples. -- **Owner-cache attribution.** Post-call, the conversation's first picked - model claims an "owner slot" for `OWNER_CACHE_TTL_SECONDS` (24h). Later - turns of the same conversation only fire bandit/state updates if the - same model handled them — mismatches are dropped (no attribution) and - counted in `skipped_updates_total`. Conversation identity is the +- **Previous-response attribution.** Post-call, feedback from the current user + message is attributed to the model that produced the previous response, while + response signals are attributed to the current model. Contexts expire after + 24 hours and the in-memory cache is capped at 1,024 sessions. Conversation identity is the client-supplied `litellm_session_id` if present, otherwise a sha256 over caller identity (api key hash, team, user, end-user) + the first message. - **Per-turn updates.** `satisfaction → +α`. `misalignment, stagnation, @@ -76,12 +75,6 @@ Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key model can still be picked. - **Hard sample cap at 200.** Once `α + β > 200`, deltas are silently dropped. No rescaling — drift is a v1 concern. -- **24h owner-cache TTL.** No explicit eviction below TTL. The in-memory map - can grow if traffic patterns produce many one-shot sessions. -- **Owner-recovery skew.** If model A "owns" a conversation but is then - dethroned in the bandit, later turns served by model B are dropped — so - bandit updates for that conversation flatline until A's TTL expires. - Tracked via `skipped_updates_total`. - **Signals are regex + tool-call only.** No LLM-judge, no embedding similarity, no exemplar storage. Signals are best-effort and biased toward English. - **One AdaptiveRouter per `Router`.** Multiple `adaptive_router/*` deployments diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 69d6a019e68..ec84eb1decf 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -3,25 +3,21 @@ Main adaptive router strategy. See README.md for design overview. One AdaptiveRouter instance per router_name. Holds in-memory caches: - _cells: Beta(alpha, beta) bandit posteriors per (request_type, model) -- _owner_cache: session_key -> (owner_model, expires_at) — the first model - picked for a conversation owns its bandit-update slot - _session_states: (session_key, model) -> SessionState for incremental signal updates Owns the AdaptiveRouterUpdateQueue used by the proxy's flusher to persist state and session snapshots back to Postgres. -Routing is stateless per-turn (Thompson sample fresh on every call). The -owner cache is consulted only at post-call time to decide whether a turn's -signals should fire a bandit update — turns served by a different model than -the conversation's owner are skipped to avoid cross-model misattribution. +Routing is stateless per-turn (Thompson sample fresh on every call). """ from __future__ import annotations import asyncio import time -from dataclasses import asdict -from typing import Any, Dict, List, Optional, Tuple, Union, cast +from collections import OrderedDict +from dataclasses import asdict, dataclass +from typing import Any, Union, cast from litellm._logging import verbose_router_logger from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -38,13 +34,18 @@ from litellm.router_strategy.adaptive_router.config import ( ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY, MIN_QUALITY_TIER_HEADER, MIN_QUALITY_TIER_METADATA_KEY, + MIN_TURNS_FOR_CLEAN_CREDIT, OWNER_CACHE_TTL_SECONDS, ) from litellm.router_strategy.adaptive_router.signals import ( SessionState, SignalDelta, Turn, - apply_turn, + advance_session_state, + apply_signal_delta, + detect_response_signals, + detect_user_feedback, + merge_signal_deltas, ) from litellm.router_strategy.adaptive_router.update_queue import ( AdaptiveRouterUpdateQueue, @@ -53,8 +54,7 @@ from litellm.router_strategy.adaptive_router.update_queue import ( # Sweep session-state cache when it exceeds this many live entries. Expired # entries are dropped in bulk; amortizes to O(1) per insert. _SESSION_STATE_SWEEP_THRESHOLD: int = 1024 -# Same pattern for the owner cache. -_OWNER_CACHE_SWEEP_THRESHOLD: int = 1024 +_FEEDBACK_CONTEXT_MAX_ENTRIES: int = 1024 from litellm.repositories.table_repositories import AdaptiveRouterStateRepository from litellm.types.llms.openai import AllMessageValues from litellm.types.router import ( @@ -70,6 +70,17 @@ def _default_prefs() -> AdaptiveRouterPreferences: return AdaptiveRouterPreferences(quality_tier=2, strengths=[]) +@dataclass(frozen=True, slots=True) +class _FeedbackContext: + model_name: str + request_type: RequestType + user_content: str | None + assistant_content: str | None + turn_count: int + clean_credit_awarded: bool + expires_at: float + + class AdaptiveRouter: """One instance per router_name. Holds in-memory caches + the update queue.""" @@ -77,8 +88,8 @@ class AdaptiveRouter: self, router_name: str, config: AdaptiveRouterConfig, - model_to_prefs: Dict[str, AdaptiveRouterPreferences], - model_to_cost: Dict[str, float], + model_to_prefs: dict[str, AdaptiveRouterPreferences], + model_to_cost: dict[str, float], ) -> None: self.router_name = router_name self.config = config @@ -86,13 +97,14 @@ class AdaptiveRouter: self.model_to_cost = model_to_cost self.queue = AdaptiveRouterUpdateQueue() - self._cells: Dict[Tuple[RequestType, str], BanditCell] = {} - self._owner_cache: Dict[str, Tuple[str, float]] = {} - self._session_states: Dict[Tuple[str, str], SessionState] = {} - # Parallel expiry map for _session_states, same TTL as _owner_cache. - # Evicted opportunistically in `get_or_create_session_state`. - self._session_states_expiry: Dict[Tuple[str, str], float] = {} - self._skipped_updates_total: int = 0 + self._cells: dict[tuple[RequestType, str], BanditCell] = {} + self._session_states: dict[tuple[str, str], SessionState] = {} + self._feedback_contexts: OrderedDict[str, _FeedbackContext] = OrderedDict() + self._session_states_expiry: dict[tuple[str, str], float] = {} + self._feedback_attributed_total: int = 0 + self._feedback_without_context_total: int = 0 + self._cross_model_feedback_total: int = 0 + self._response_signal_updates_total: int = 0 # Set to True once the proxy flusher has loaded persisted priors from # Postgres. Checked to support lazy-load on hot-reloaded routers. self._state_loaded: bool = False @@ -145,11 +157,11 @@ class AdaptiveRouter: async def async_pre_routing_hook( self, model: str, - request_kwargs: Dict[str, Any], - messages: Optional[List[Dict[str, Any]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - ) -> Optional[PreRoutingHookResponse]: + request_kwargs: dict[str, Any], + messages: list[dict[str, Any]] | None = None, + input: Union[str, list] | None = None, + specific_deployment: bool | None = False, + ) -> PreRoutingHookResponse | None: """ Plugin entry point invoked by `Router.async_pre_routing_hook` when the inbound `model` matches this adaptive router's `router_name`. @@ -159,11 +171,9 @@ class AdaptiveRouter: post-call hook can surface it as a response header. Routing is stateless per-turn: every call Thompson-samples fresh, - regardless of any prior pick for the same session. Cross-turn - attribution is enforced post-call via the owner cache (see - `claim_or_check_owner`). + regardless of any prior pick for the same session. """ - user_text = get_last_user_message(cast(List[AllMessageValues], messages or [])) or "" + user_text = get_last_user_message(cast(list[AllMessageValues], messages or [])) or "" request_type = classify_prompt(user_text) min_quality_tier = self._extract_min_quality_tier(request_kwargs) @@ -190,7 +200,7 @@ class AdaptiveRouter: async def pick_model( self, request_type: RequestType, - min_quality_tier: Optional[int] = None, + min_quality_tier: int | None = None, ) -> str: """Thompson-sample across eligible models. Stateless per-turn.""" eligible = self._eligible_models(min_quality_tier) @@ -206,44 +216,7 @@ class AdaptiveRouter: cost_weight=self.config.weights.cost, ) - def claim_or_check_owner(self, session_key: str, current_model: str) -> bool: - """Resolve attribution for a turn under stateless routing. - - Returns True iff this turn should fire a bandit/state update. The - first call for a `session_key` claims ownership for `current_model` - and returns True. Subsequent calls return True only if the owner is - still live AND matches `current_model`. Mismatches (a different - model handled this turn) and expired owners both increment - `_skipped_updates_total` and return False — no attribution. - """ - now = time.time() - existing = self._owner_cache.get(session_key) - if existing is not None and existing[1] > now: - owner_model, _ = existing - if owner_model == current_model: - return True - self._skipped_updates_total += 1 - return False - - # Opportunistic bulk sweep — sessions that never come back would - # otherwise pile up here forever. Same threshold pattern as the - # session-state cache. - if len(self._owner_cache) >= _OWNER_CACHE_SWEEP_THRESHOLD: - self._evict_expired_owner_cache(now) - - # No live owner -> claim for current_model. - self._owner_cache[session_key] = ( - current_model, - now + OWNER_CACHE_TTL_SECONDS, - ) - return True - - def _evict_expired_owner_cache(self, now: float) -> None: - expired = [k for k, (_, exp) in self._owner_cache.items() if exp <= now] - for k in expired: - self._owner_cache.pop(k, None) - - async def get_state_snapshot(self) -> Dict[str, Any]: + async def get_state_snapshot(self) -> dict[str, Any]: """In-memory snapshot for the introspection endpoint. Cheap; no DB hit.""" cells = [] for (rt, model), cell in sorted(self._cells.items(), key=lambda kv: (kv[0][0].value, kv[0][1])): @@ -264,7 +237,7 @@ class AdaptiveRouter: ) queue = await self.queue.queue_size() now = time.time() - owner_cache_live = sum(1 for _, exp in self._owner_cache.values() if exp > now) + feedback_contexts_live = sum(1 for context in self._feedback_contexts.values() if context.expires_at > now) return { "router_name": self.router_name, "available_models": list(self.config.available_models), @@ -274,15 +247,18 @@ class AdaptiveRouter: }, "model_costs": dict(self.model_to_cost), "cells": cells, - "owner_cache_live": owner_cache_live, - "skipped_updates_total": self._skipped_updates_total, + "feedback_contexts_live": feedback_contexts_live, + "feedback_attributed_total": self._feedback_attributed_total, + "feedback_without_context_total": self._feedback_without_context_total, + "cross_model_feedback_total": self._cross_model_feedback_total, + "response_signal_updates_total": self._response_signal_updates_total, "queue": queue, } @staticmethod def _extract_min_quality_tier( - request_kwargs: Dict[str, Any], - ) -> Optional[int]: + request_kwargs: dict[str, Any], + ) -> int | None: """Pull `min_quality_tier` from request headers or metadata. Precedence: headers (`x-litellm-min-quality-tier`) over metadata @@ -310,7 +286,7 @@ class AdaptiveRouter: return None return None - def _eligible_models(self, min_quality_tier: Optional[int]) -> List[str]: + def _eligible_models(self, min_quality_tier: int | None) -> list[str]: if min_quality_tier is None: return list(self.config.available_models) return [ @@ -363,17 +339,131 @@ class AdaptiveRouter: request_type: RequestType, turn: Turn, ) -> SignalDelta: - """Apply one turn, push session snapshot + bandit deltas to the queue.""" - state = self.get_or_create_session_state(session_id, model_name, request_type) - delta = apply_turn(state, turn) - verbose_router_logger.debug("AdaptiveRouter[%s]: record_turn delta=%s", self.router_name, delta) + """Attribute feedback to the previous response and response signals to the current model.""" + async with self._lock: + now = time.time() + while self._feedback_contexts: + oldest_context = next(iter(self._feedback_contexts.values())) + if oldest_context.expires_at > now: + break + self._feedback_contexts.popitem(last=False) + previous = self._feedback_contexts.pop(session_id, None) - # Strip the raw conversation content before persisting. The - # last_user/assistant_content and tool_call_history fields are only - # needed in-memory for the next turn's incremental signal detection; - # writing user prompts and tool payloads to the DB would store PII - # for every adaptive-router conversation. Counts + bookkeeping is - # all the persisted row needs. + effective_request_type = ( + previous.request_type if previous is not None and request_type == RequestType.GENERAL else request_type + ) + current_state = self.get_or_create_session_state( + session_id, + model_name, + effective_request_type, + ) + feedback_delta = detect_user_feedback( + previous.user_content if previous else None, + turn.user_content, + turn.tool_results, + allow_satisfaction=( + previous is not None + and not previous.clean_credit_awarded + and previous.turn_count + 1 >= MIN_TURNS_FOR_CLEAN_CREDIT + ), + ) + previous_assistant = previous.assistant_content if previous else None + response_delta = detect_response_signals( + previous_assistant, + turn.assistant_content, + current_state.tool_call_history, + turn.tool_calls, + turn.tool_results, + turn.response_status, + ) + states_to_persist: dict[str, SessionState] = {model_name: current_state} + bandit_deltas: dict[tuple[RequestType, str], SignalDelta] = {} + + if previous is not None: + feedback_state = self.get_or_create_session_state( + session_id, + previous.model_name, + previous.request_type, + ) + apply_signal_delta(feedback_state, feedback_delta) + if feedback_delta.satisfaction: + feedback_state.clean_credit_awarded = True + states_to_persist[previous.model_name] = feedback_state + if feedback_delta.any_fired(): + self._feedback_attributed_total += 1 + if previous.model_name != model_name: + self._cross_model_feedback_total += 1 + bandit_deltas[(previous.request_type, previous.model_name)] = feedback_delta + else: + if feedback_delta.any_fired(): + self._feedback_without_context_total += 1 + initial_failure = SignalDelta(failure=feedback_delta.failure) + apply_signal_delta(current_state, initial_failure) + bandit_deltas[(effective_request_type, model_name)] = initial_failure + + apply_signal_delta(current_state, response_delta) + if self._compute_bandit_delta(response_delta) != (0.0, 0.0): + self._response_signal_updates_total += 1 + current_key = (effective_request_type, model_name) + bandit_deltas[current_key] = merge_signal_deltas( + bandit_deltas.get(current_key, SignalDelta()), + response_delta, + ) + advance_session_state(current_state, turn) + + next_turn_count = (previous.turn_count if previous else 0) + 1 + clean_credit_awarded = bool((previous and previous.clean_credit_awarded) or feedback_delta.satisfaction) + if len(self._feedback_contexts) >= _FEEDBACK_CONTEXT_MAX_ENTRIES: + self._feedback_contexts.popitem(last=False) + self._feedback_contexts[session_id] = _FeedbackContext( + model_name=model_name, + request_type=effective_request_type, + user_content=turn.user_content, + assistant_content=turn.assistant_content, + turn_count=next_turn_count, + clean_credit_awarded=clean_credit_awarded, + expires_at=now + OWNER_CACHE_TTL_SECONDS, + ) + + for state_model, state in states_to_persist.items(): + await self.queue.add_session_state( + session_id, + self.router_name, + state_model, + self._persistable_session_snapshot(state), + ) + + combined_delta = SignalDelta() + for (attribution_type, target_model), delta in bandit_deltas.items(): + combined_delta = merge_signal_deltas(combined_delta, delta) + d_alpha, d_beta = self._compute_bandit_delta(delta) + if d_alpha == 0 and d_beta == 0: + continue + cell_key = (attribution_type, target_model) + self._cells[cell_key] = apply_delta( + self._cells[cell_key], + d_alpha, + d_beta, + ) + await self.queue.add_state_delta( + self.router_name, + attribution_type.value, + target_model, + d_alpha, + d_beta, + ) + + verbose_router_logger.debug( + "AdaptiveRouter[%s]: feedback_target=%s current_model=%s delta=%s", + self.router_name, + previous.model_name if previous else None, + model_name, + combined_delta, + ) + return combined_delta + + @staticmethod + def _persistable_session_snapshot(state: SessionState) -> dict[str, Any]: snapshot = asdict(state) for sensitive in ( "last_user_content", @@ -382,38 +472,10 @@ class AdaptiveRouter: "pending_tool_calls", ): snapshot.pop(sensitive, None) - await self.queue.add_session_state(session_id, self.router_name, model_name, snapshot) - - d_alpha, d_beta = self._compute_bandit_delta(delta) - verbose_router_logger.debug( - "AdaptiveRouter[%s]: bandit delta alpha=%.2f beta=%.2f", - self.router_name, - d_alpha, - d_beta, - ) - if d_alpha != 0 or d_beta != 0: - # For non-GENERAL turns, attribute to the current-turn classification - # so genuine mid-session topic shifts (e.g. code → math) update the - # correct cell. For GENERAL turns ("thanks!", "ok", "sounds good"), fall - # back to the session's original type so closing pleasantries don't - # misattribute the reward. - attribution_type = ( - request_type if request_type != RequestType.GENERAL else RequestType(state.classified_type) - ) - cell_key = (attribution_type, model_name) - self._cells[cell_key] = apply_delta(self._cells[cell_key], d_alpha, d_beta) - await self.queue.add_state_delta( - self.router_name, - attribution_type.value, - model_name, - d_alpha, - d_beta, - ) - - return delta + return snapshot @staticmethod - def _compute_bandit_delta(delta: SignalDelta) -> Tuple[float, float]: + def _compute_bandit_delta(delta: SignalDelta) -> tuple[float, float]: """ Translate per-turn signal deltas into bandit-cell deltas. diff --git a/litellm/router_strategy/adaptive_router/hooks.py b/litellm/router_strategy/adaptive_router/hooks.py index c3e3f8ca74a..89ae28be227 100644 --- a/litellm/router_strategy/adaptive_router/hooks.py +++ b/litellm/router_strategy/adaptive_router/hooks.py @@ -214,10 +214,6 @@ class AdaptiveRouterPostCallHook(CustomLogger): ) -> None: try: messages = kwargs.get("messages") or [] - if len(messages) < SIGNAL_GATE_MIN_MESSAGES: - # Too few turns for any signal to be meaningful — skip. - return - session_key = _resolve_session_key(kwargs) if not session_key: return @@ -233,10 +229,6 @@ class AdaptiveRouterPostCallHook(CustomLogger): if not current_model: return - if not self.adaptive_router.claim_or_check_owner(session_key, current_model): - # A different model owns this conversation — skip attribution. - return - user_text = _last_user_content(messages) assistant_text, tool_calls = _assistant_content_and_tool_calls(response_obj) tool_results = _recent_tool_results(messages) diff --git a/litellm/router_strategy/adaptive_router/signals.py b/litellm/router_strategy/adaptive_router/signals.py index 2fd1d24fbbe..74fa8936098 100644 --- a/litellm/router_strategy/adaptive_router/signals.py +++ b/litellm/router_strategy/adaptive_router/signals.py @@ -14,7 +14,7 @@ from __future__ import annotations import re from dataclasses import dataclass, field -from typing import Any, Dict, List, Optional, Set +from typing import Any from litellm.router_strategy.adaptive_router.config import ( LOOP_REPEAT_THRESHOLD, @@ -74,26 +74,26 @@ class SessionState: loop_count: int = 0 exhaustion_count: int = 0 - last_user_content: Optional[str] = None - last_assistant_content: Optional[str] = None - tool_call_history: List[str] = field(default_factory=list) - pending_tool_calls: Dict[str, str] = field(default_factory=dict) + last_user_content: str | None = None + last_assistant_content: str | None = None + tool_call_history: list[str] = field(default_factory=list) + pending_tool_calls: dict[str, str] = field(default_factory=dict) turn_count: int = 0 last_processed_turn: int = -1 clean_credit_awarded: bool = False - terminal_status: Optional[int] = None + terminal_status: int | None = None @dataclass class Turn: """One turn of input. Caller assembles this from the request/response.""" - user_content: Optional[str] = None - assistant_content: Optional[str] = None - tool_calls: List[Dict[str, Any]] = field(default_factory=list) - tool_results: List[Dict[str, Any]] = field(default_factory=list) - response_status: Optional[int] = None + user_content: str | None = None + assistant_content: str | None = None + tool_calls: list[dict[str, Any]] = field(default_factory=list) + tool_results: list[dict[str, Any]] = field(default_factory=list) + response_status: int | None = None # ---- Detection helpers ---------------------------------------------------- @@ -101,13 +101,13 @@ class Turn: _TOKEN_RE = re.compile(r"[A-Za-z0-9]+") -def _tokens(text: Optional[str]) -> Set[str]: +def _tokens(text: str | None) -> set[str]: if not text: return set() return {t.lower() for t in _TOKEN_RE.findall(text)} -def _jaccard(a: Set[str], b: Set[str]) -> float: +def _jaccard(a: set[str], b: set[str]) -> float: union = a | b if not union: return 0.0 @@ -130,7 +130,7 @@ _SATISFACTION_PATTERNS = [ ] -def _detect_misalignment(prev_user: Optional[str], curr_user: Optional[str]) -> bool: +def _detect_misalignment(prev_user: str | None, curr_user: str | None) -> bool: """Fires when consecutive user messages share *some* topic (jaccard > 0) but are sufficiently different (jaccard < threshold) — i.e. user is rephrasing, not changing topic, not repeating.""" @@ -140,7 +140,7 @@ def _detect_misalignment(prev_user: Optional[str], curr_user: Optional[str]) -> return 0.0 < j < MISALIGNMENT_JACCARD_THRESHOLD -def _detect_stagnation(prev_asst: Optional[str], curr_asst: Optional[str]) -> bool: +def _detect_stagnation(prev_asst: str | None, curr_asst: str | None) -> bool: """Fires when consecutive assistant messages are near-duplicates.""" if not prev_asst or not curr_asst: return False @@ -148,19 +148,19 @@ def _detect_stagnation(prev_asst: Optional[str], curr_asst: Optional[str]) -> bo return j >= STAGNATION_JACCARD_NEAR_DUP -def _detect_disengagement(curr_user: Optional[str]) -> bool: +def _detect_disengagement(curr_user: str | None) -> bool: if not curr_user: return False return any(p.search(curr_user) for p in _DISENGAGEMENT_PATTERNS) -def _detect_satisfaction(curr_user: Optional[str]) -> bool: +def _detect_satisfaction(curr_user: str | None) -> bool: if not curr_user: return False return any(p.search(curr_user) for p in _SATISFACTION_PATTERNS) -def _detect_failure(tool_results: List[Dict[str, Any]]) -> bool: +def _detect_failure(tool_results: list[dict[str, Any]]) -> bool: """Any tool result explicitly flagged as an error. We do NOT treat empty content as failure — many tools legitimately return @@ -173,7 +173,7 @@ def _detect_failure(tool_results: List[Dict[str, Any]]) -> bool: return False -def _signature(call: Dict[str, Any]) -> str: +def _signature(call: dict[str, Any]) -> str: """Stable signature for loop detection: name + sorted JSON-ish args.""" name = call.get("name") or call.get("function", {}).get("name", "") call_args = call.get("arguments") @@ -184,7 +184,7 @@ def _signature(call: Dict[str, Any]) -> str: return f"{name}({call_args})" -def _detect_loop(history: List[str], new_calls: List[Dict[str, Any]]) -> bool: +def _detect_loop(history: list[str], new_calls: list[dict[str, Any]]) -> bool: """Fires if any new call's signature appears >= LOOP_REPEAT_THRESHOLD-1 times in recent history (so this call would be the Nth).""" if not new_calls: @@ -209,7 +209,7 @@ _EXHAUSTION_KEYWORDS = ( ) -def _detect_exhaustion(status: Optional[int], tool_results: List[Dict[str, Any]]) -> bool: +def _detect_exhaustion(status: int | None, tool_results: list[dict[str, Any]]) -> bool: if status is not None and status in _EXHAUSTION_STATUSES: return True for r in tool_results: @@ -219,39 +219,53 @@ def _detect_exhaustion(status: Optional[int], tool_results: List[Dict[str, Any]] return False -# ---- Public entrypoint ---------------------------------------------------- +def detect_user_feedback( + previous_user_content: str | None, + current_user_content: str | None, + tool_results: list[dict[str, Any]], + allow_satisfaction: bool, +) -> SignalDelta: + return SignalDelta( + misalignment=int(_detect_misalignment(previous_user_content, current_user_content)), + disengagement=int(_detect_disengagement(current_user_content)), + satisfaction=int(allow_satisfaction and _detect_satisfaction(current_user_content)), + failure=int(_detect_failure(tool_results)), + ) -def apply_turn(state: SessionState, turn: Turn) -> SignalDelta: - """ - Detect signals on this turn, mutate state, return the delta. +def detect_response_signals( + previous_assistant_content: str | None, + current_assistant_content: str | None, + tool_call_history: list[str], + tool_calls: list[dict[str, Any]], + tool_results: list[dict[str, Any]], + response_status: int | None, +) -> SignalDelta: + return SignalDelta( + stagnation=int( + _detect_stagnation( + previous_assistant_content, + current_assistant_content, + ) + ), + loop=int(_detect_loop(tool_call_history, tool_calls)), + exhaustion=int(_detect_exhaustion(response_status, tool_results)), + ) - O(1) per turn (no full-history rescan). Only inspects last_*, recent tool history - (which is bounded at TOOL_CALL_HISTORY_MAX), and the new turn payload. - """ - delta = SignalDelta() - if _detect_misalignment(state.last_user_content, turn.user_content): - delta.misalignment = 1 - if _detect_stagnation(state.last_assistant_content, turn.assistant_content): - delta.stagnation = 1 - if _detect_disengagement(turn.user_content): - delta.disengagement = 1 - if _detect_satisfaction(turn.user_content): - # Gate: only award satisfaction credit once per session, and only - # after MIN_TURNS_FOR_CLEAN_CREDIT turns of context. Early "thanks" - # on turn 1-2 is noise, not a validated quality signal. - current_turn_index = state.turn_count + 1 - if not state.clean_credit_awarded and current_turn_index >= MIN_TURNS_FOR_CLEAN_CREDIT: - delta.satisfaction = 1 - state.clean_credit_awarded = True - if _detect_failure(turn.tool_results): - delta.failure = 1 - if _detect_loop(state.tool_call_history, turn.tool_calls): - delta.loop = 1 - if _detect_exhaustion(turn.response_status, turn.tool_results): - delta.exhaustion = 1 +def merge_signal_deltas(*deltas: SignalDelta) -> SignalDelta: + return SignalDelta( + misalignment=sum(delta.misalignment for delta in deltas), + stagnation=sum(delta.stagnation for delta in deltas), + disengagement=sum(delta.disengagement for delta in deltas), + satisfaction=sum(delta.satisfaction for delta in deltas), + failure=sum(delta.failure for delta in deltas), + loop=sum(delta.loop for delta in deltas), + exhaustion=sum(delta.exhaustion for delta in deltas), + ) + +def apply_signal_delta(state: SessionState, delta: SignalDelta) -> None: state.misalignment_count += delta.misalignment state.stagnation_count += delta.stagnation state.disengagement_count += delta.disengagement @@ -260,6 +274,8 @@ def apply_turn(state: SessionState, turn: Turn) -> SignalDelta: state.loop_count += delta.loop state.exhaustion_count += delta.exhaustion + +def advance_session_state(state: SessionState, turn: Turn) -> None: if turn.user_content: state.last_user_content = turn.user_content if turn.assistant_content: @@ -276,4 +292,38 @@ def apply_turn(state: SessionState, turn: Turn) -> SignalDelta: state.turn_count += 1 state.last_processed_turn = state.turn_count + +# ---- Public entrypoint ---------------------------------------------------- + + +def apply_turn(state: SessionState, turn: Turn) -> SignalDelta: + """ + Detect signals on this turn, mutate state, return the delta. + + O(1) per turn (no full-history rescan). Only inspects last_*, recent tool history + (which is bounded at TOOL_CALL_HISTORY_MAX), and the new turn payload. + """ + feedback_delta = detect_user_feedback( + state.last_user_content, + turn.user_content, + turn.tool_results, + allow_satisfaction=(not state.clean_credit_awarded and state.turn_count + 1 >= MIN_TURNS_FOR_CLEAN_CREDIT), + ) + response_delta = detect_response_signals( + state.last_assistant_content, + turn.assistant_content, + state.tool_call_history, + turn.tool_calls, + turn.tool_results, + turn.response_status, + ) + delta = merge_signal_deltas( + feedback_delta, + response_delta, + ) + apply_signal_delta(state, delta) + if delta.satisfaction: + state.clean_credit_awarded = True + advance_session_state(state, turn) + return delta diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 74644f01be8..bebdbba90ef 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -13,10 +13,12 @@ evaluated before either classification strategy and force a tier outright when m Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter """ +from __future__ import annotations + import asyncio import random import re -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast +from typing import TYPE_CHECKING, Any, Literal, Optional, Union, cast from pydantic import BaseModel @@ -38,6 +40,7 @@ if TYPE_CHECKING: from semantic_router.routers import SemanticRouter from litellm.router import Router + from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter from litellm.types.router import PreRoutingHookResponse else: Router = Any @@ -63,7 +66,7 @@ Tiers: {prompt}""" -def _append_custom_keywords(base_keywords: list[str], custom_keywords: Optional[list[str]]) -> list[str]: +def _append_custom_keywords(base_keywords: list[str], custom_keywords: list[str] | None) -> list[str]: if not custom_keywords: return base_keywords base_lowered = frozenset(keyword.lower() for keyword in base_keywords) @@ -95,7 +98,7 @@ def _sanitize_user_api_key_auth(auth: Any) -> Any: return auth -def _classifier_call_metadata(metadata: Optional[dict[str, Any]]) -> Optional[dict[str, Any]]: +def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any] | None: if not metadata: return metadata return { @@ -110,7 +113,7 @@ class DimensionScore: __slots__ = ("name", "score", "signal") - def __init__(self, name: str, score: float, signal: Optional[str] = None): + def __init__(self, name: str, score: float, signal: str | None = None): self.name = name self.score = score self.signal = signal @@ -134,9 +137,9 @@ class ComplexityRouter(CustomLogger): def __init__( self, model_name: str, - litellm_router_instance: "Router", - complexity_router_config: Optional[Dict[str, Any]] = None, - default_model: Optional[str] = None, + litellm_router_instance: Router, + complexity_router_config: dict[str, Any] | None = None, + default_model: str | None = None, ): """ Initialize ComplexityRouter. @@ -173,7 +176,7 @@ class ComplexityRouter(CustomLogger): # embeddings are static, only the prompt is embedded per request). The lock # serializes the one-time build so concurrent cold-start requests don't each # construct the index and fire duplicate embedding calls. - self._semantic_routelayer: Optional[SemanticRouter] = None + self._semantic_routelayer: SemanticRouter | None = None self._semantic_routelayer_lock = asyncio.Lock() # Pre-compile regex patterns for efficiency @@ -185,6 +188,10 @@ class ComplexityRouter(CustomLogger): re.compile(r"[a-z]\)\s", re.IGNORECASE), ] + self.adaptive_router: AdaptiveRouter | None = None + self._model_tiers: dict[str, tuple[ComplexityTier, ...]] = {} + self._adaptive_init_attempted = False + verbose_router_logger.debug(f"ComplexityRouter initialized for {model_name} with tiers: {self.config.tiers}") def _estimate_tokens(self, text: str) -> int: @@ -228,12 +235,12 @@ class ComplexityRouter(CustomLogger): def _score_keyword_match( self, text: str, - keywords: List[str], + keywords: list[str], name: str, signal_label: str, - thresholds: Tuple[int, int], # (low, high) - scores: Tuple[float, float, float], # (none, low, high) - ) -> Tuple[DimensionScore, int]: + thresholds: tuple[int, int], # (low, high) + scores: tuple[float, float, float], # (none, low, high) + ) -> tuple[DimensionScore, int]: """Score based on keyword matches using word boundary matching. Returns: @@ -271,7 +278,7 @@ class ComplexityRouter(CustomLogger): return DimensionScore("questionComplexity", 0.5, f"{count} questions") return DimensionScore("questionComplexity", 0, None) - def classify(self, prompt: str, system_prompt: Optional[str] = None) -> Tuple[ComplexityTier, float, List[str]]: + def classify(self, prompt: str, system_prompt: str | None = None) -> tuple[ComplexityTier, float, list[str]]: """ Classify a prompt by complexity. @@ -330,7 +337,7 @@ class ComplexityRouter(CustomLogger): (0, -1.0, -1.0), ) - dimensions: List[DimensionScore] = [ + dimensions: list[DimensionScore] = [ self._score_token_count(estimated_tokens), code_score, reasoning_score, @@ -372,8 +379,8 @@ class ComplexityRouter(CustomLogger): async def aclassify( self, prompt: str, - system_prompt: Optional[str] = None, - request_kwargs: Optional[dict[str, Any]] = None, + system_prompt: str | None = None, + request_kwargs: dict[str, Any] | None = None, ) -> tuple[ComplexityTier, float, list[str]]: """ Classify a prompt by complexity, using the LLM classifier when configured. @@ -396,8 +403,8 @@ class ComplexityRouter(CustomLogger): async def _classify_with_llm( self, prompt: str, - system_prompt: Optional[str] = None, - request_kwargs: Optional[dict[str, Any]] = None, + system_prompt: str | None = None, + request_kwargs: dict[str, Any] | None = None, ) -> ComplexityTier: """Call the configured classifier model and parse its structured tier response.""" llm_config = self.config.classifier_llm_config @@ -458,7 +465,176 @@ class ComplexityRouter(CustomLogger): raise ValueError(f"Empty model pool for tier {tier_key}") return random.choice(model) - def _lexical_tier_override(self, user_message: str) -> Optional[ComplexityTier]: + def _tier_pools(self) -> dict[str, list[str]]: + return {tier: (models if isinstance(models, list) else [models]) for tier, models in self.config.tiers.items()} + + def _ensure_adaptive_router(self) -> Any | None: + if not self.config.adaptive: + return None + if self.adaptive_router is not None: + return self.adaptive_router + if self._adaptive_init_attempted: + return self.adaptive_router + self._adaptive_init_attempted = True + + from litellm.router_strategy.adaptive_router.adaptive_router import ( + AdaptiveRouter, + ) + from litellm.router_strategy.adaptive_router.config import ( + ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY, + ) + from litellm.types.router import ( + AdaptiveRouterConfig, + AdaptiveRouterPreferences, + ) + + pools = self._tier_pools() + available_models = list(dict.fromkeys(model for models in pools.values() for model in models)) + self._model_tiers = { + model: tuple(ComplexityTier(tier_name) for tier_name, models in pools.items() if model in models) + for model in available_models + } + + model_to_prefs: dict[str, AdaptiveRouterPreferences] = {} + model_to_cost: dict[str, float] = {} + model_list = getattr(self.litellm_router_instance, "model_list", None) or [] + name_to_indices = getattr(self.litellm_router_instance, "model_name_to_deployment_indices", {}) or {} + for name in available_models: + indices = name_to_indices.get(name, []) + if not indices: + model_to_prefs[name] = AdaptiveRouterPreferences(quality_tier=2, strengths=[]) + model_to_cost[name] = 0.0 + continue + deployment = model_list[indices[0]] + mi = deployment.get("model_info") if isinstance(deployment, dict) else deployment.model_info + mi_dict: dict[str, Any] = mi if isinstance(mi, dict) else (mi.model_dump() if mi else {}) + prefs_raw = mi_dict.get("adaptive_router_preferences") + if prefs_raw is not None: + model_to_prefs[name] = AdaptiveRouterPreferences(**prefs_raw) + else: + model_to_prefs[name] = AdaptiveRouterPreferences(quality_tier=2, strengths=[]) + + lp = deployment.get("litellm_params") if isinstance(deployment, dict) else deployment.litellm_params + lp_dict: dict[str, Any] = lp if isinstance(lp, dict) else (lp.model_dump() if lp else {}) + cost = lp_dict.get("input_cost_per_token") + model_to_cost[name] = float(cost) if cost is not None else 0.0 + + self.adaptive_router = AdaptiveRouter( + router_name=self.model_name, + config=AdaptiveRouterConfig( + available_models=available_models, + weights=self.config.adaptive_weights, + ), + model_to_prefs=model_to_prefs, + model_to_cost=model_to_cost, + ) + self._adaptive_chosen_model_key = ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY + return self.adaptive_router + + def _soft_floor_pick( + self, + classified_tier: ComplexityTier, + user_message: str, + request_kwargs: dict[str, Any] | None = None, + ) -> str: + from litellm.router_strategy.adaptive_router.bandit import ( + normalized_cost, + thompson_sample, + ) + from litellm.router_strategy.adaptive_router.classifier import classify_prompt + + adaptive = self._ensure_adaptive_router() + if adaptive is None: + return self.get_model_for_tier(classified_tier) + + request_type = classify_prompt(user_message) + classified_idx = TIER_SEVERITY_ORDER.index(classified_tier) + pools = self._tier_pools() + classified_candidates = tuple(pools.get(classified_tier.value, ())) + cold_start_candidates = tuple( + model for model in classified_candidates if adaptive._cells[(request_type, model)].total_samples == 0 + ) + if cold_start_candidates: + chosen_model = random.choice(cold_start_candidates) + if request_kwargs is not None: + metadata = request_kwargs.setdefault("metadata", {}) + if isinstance(metadata, dict): + metadata["adaptive_router_decision"] = { + "phase": "cold_start", + "classified_tier": classified_tier.value, + "request_type": request_type.value, + "eligible_mode": "classified_tier", + "quality_weight": self.config.adaptive_weights.quality, + "cost_weight": self.config.adaptive_weights.cost, + "tier_distance_penalty": self.config.tier_distance_penalty, + "chosen_model": chosen_model, + "candidates": [ + { + "model": model, + "total_samples": adaptive._cells[(request_type, model)].total_samples, + } + for model in cold_start_candidates + ], + } + return chosen_model + if self.config.adaptive_eligible == "classified_tier": + candidates = list(classified_candidates) + if not candidates: + return self.get_model_for_tier(classified_tier) + else: + candidates = list(adaptive.config.available_models) + + all_costs = [adaptive.model_to_cost.get(m, 0.0) for m in candidates] + quality_weight = self.config.adaptive_weights.quality + cost_weight = self.config.adaptive_weights.cost + penalty_weight = self.config.tier_distance_penalty + + best_model: str | None = None + best_score = float("-inf") + candidate_scores: list[dict[str, Any]] = [] + for model in candidates: + cell = adaptive._cells[(request_type, model)] + quality_sample = thompson_sample(cell) + cost_score = normalized_cost(adaptive.model_to_cost.get(model, 0.0), all_costs) + if self.config.adaptive_eligible == "classified_tier": + distance = 0 + else: + model_tiers = self._model_tiers.get(model, (classified_tier,)) + distance = min( + abs(TIER_SEVERITY_ORDER.index(model_tier) - classified_idx) for model_tier in model_tiers + ) + score = quality_weight * quality_sample + cost_weight * cost_score - penalty_weight * distance + candidate_scores.append( + { + "model": model, + "quality_sample": quality_sample, + "cost_score": cost_score, + "tier_distance": distance, + "score": score, + } + ) + if score > best_score: + best_score = score + best_model = model + if best_model is None: + return self.get_model_for_tier(classified_tier) + if request_kwargs is not None: + metadata = request_kwargs.setdefault("metadata", {}) + if isinstance(metadata, dict): + metadata["adaptive_router_decision"] = { + "phase": "adaptive", + "classified_tier": classified_tier.value, + "request_type": request_type.value, + "eligible_mode": self.config.adaptive_eligible, + "quality_weight": quality_weight, + "cost_weight": cost_weight, + "tier_distance_penalty": penalty_weight, + "chosen_model": best_model, + "candidates": candidate_scores, + } + return best_model + + def _lexical_tier_override(self, user_message: str) -> ComplexityTier | None: """When keyword_tier_rules match literally, the most-severe matched tier wins. Escalating to the highest tier (rather than the first rule in the list) keeps @@ -476,7 +652,7 @@ class ComplexityRouter(CustomLogger): return None return max(matched_tiers, key=TIER_SEVERITY_ORDER.index) - def _get_or_create_semantic_routelayer(self) -> "SemanticRouter": + def _get_or_create_semantic_routelayer(self) -> SemanticRouter: """Build (once) a SemanticRouter with one route per tier, utterances = that tier's keywords.""" if self._semantic_routelayer is not None: return self._semantic_routelayer @@ -515,7 +691,7 @@ class ComplexityRouter(CustomLogger): self._semantic_routelayer = routelayer return routelayer - async def _ensure_semantic_routelayer(self) -> "SemanticRouter": + async def _ensure_semantic_routelayer(self) -> SemanticRouter: """Return the cached route layer, building it once under a lock if needed. The build embeds the static route utterances via the encoder's synchronous path, @@ -531,7 +707,7 @@ class ComplexityRouter(CustomLogger): routelayer = await asyncio.to_thread(self._get_or_create_semantic_routelayer) return routelayer - async def _semantic_tier_override(self, user_message: str, request_kwargs: Dict) -> Optional[ComplexityTier]: + async def _semantic_tier_override(self, user_message: str, request_kwargs: dict) -> ComplexityTier | None: """Match the prompt against keyword_tier_rules by embedding similarity. Embeds the query ourselves (instead of letting SemanticRouter.acall embed it @@ -571,7 +747,7 @@ class ComplexityRouter(CustomLogger): except ValueError: return None - async def _resolve_keyword_tier_override(self, user_message: str, request_kwargs: Dict) -> Optional[ComplexityTier]: + async def _resolve_keyword_tier_override(self, user_message: str, request_kwargs: dict) -> ComplexityTier | None: """Resolve a keyword_tier_rule override, semantically or lexically per config. Returns None (no override -> fall through to the scorer) not only when no rule @@ -592,9 +768,9 @@ class ComplexityRouter(CustomLogger): def _resolve_messages( self, - messages: Optional[List[Dict[str, Any]]], - request_kwargs: Dict, - ) -> Optional[List[Dict[str, Any]]]: + messages: list[dict[str, Any]] | None, + request_kwargs: dict, + ) -> list[dict[str, Any]] | None: """ Resolve messages from the request, converting from other formats if needed. @@ -609,11 +785,11 @@ class ComplexityRouter(CustomLogger): @staticmethod def _extract_user_message_and_system_prompt( - messages: List[Dict[str, Any]], - ) -> Tuple[Optional[str], Optional[str]]: + messages: list[dict[str, Any]], + ) -> tuple[str | None, str | None]: """Extract the last user message text and last system prompt from messages.""" - user_message: Optional[str] = None - system_prompt: Optional[str] = None + user_message: str | None = None + system_prompt: str | None = None for msg in reversed(messages): role = msg.get("role", "") @@ -636,11 +812,11 @@ class ComplexityRouter(CustomLogger): async def async_pre_routing_hook( self, model: str, - request_kwargs: Dict, - messages: Optional[List[Dict[str, Any]]] = None, - input: Optional[Union[str, List]] = None, - specific_deployment: Optional[bool] = False, - ) -> Optional["PreRoutingHookResponse"]: + request_kwargs: dict, + messages: list[dict[str, Any]] | None = None, + input: Union[str, list] | None = None, + specific_deployment: bool | None = False, + ) -> Optional[PreRoutingHookResponse]: """ Pre-routing hook called before the routing decision. @@ -692,12 +868,25 @@ class ComplexityRouter(CustomLogger): ) tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs) - routed_model = self.get_model_for_tier(tier) - - verbose_router_logger.info( - f"ComplexityRouter: routing decision cause=complexity_scorer, tier={tier.value}, " - f"score={score:.3f}, signals={signals}, routed_model={routed_model}" - ) + if self.config.adaptive: + routed_model = self._soft_floor_pick(tier, user_message, request_kwargs) + adaptive = self._ensure_adaptive_router() + if adaptive is not None: + kwargs_metadata = request_kwargs.setdefault("metadata", {}) + if isinstance(kwargs_metadata, dict): + chosen_key = getattr(self, "_adaptive_chosen_model_key", "adaptive_router_chosen_model") + kwargs_metadata[chosen_key] = routed_model + verbose_router_logger.info( + f"ComplexityRouter[adaptive]: routing decision cause=complexity_scorer, " + f"tier={tier.value}, score={score:.3f}, " + f"signals={signals}, routed_model={routed_model}" + ) + else: + routed_model = self.get_model_for_tier(tier) + verbose_router_logger.info( + f"ComplexityRouter: routing decision cause=complexity_scorer, tier={tier.value}, " + f"score={score:.3f}, signals={signals}, routed_model={routed_model}" + ) return PreRoutingHookResponse( model=routed_model, diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 8c8e5acb51f..df699d1a059 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -6,10 +6,12 @@ All values are configurable via proxy config.yaml. """ from enum import Enum -from typing import Dict, List, Literal, Optional +from typing import Literal from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from litellm.types.router import AdaptiveRouterWeights + class ComplexityTier(str, Enum): """Complexity tiers for routing decisions.""" @@ -27,11 +29,13 @@ TIER_SEVERITY_ORDER: tuple[ComplexityTier, ...] = ( ComplexityTier.REASONING, ) +DEFAULT_TIER_DISTANCE_PENALTY: float = 0.5 + class KeywordTierRule(BaseModel): """A deterministic override: if any keyword matches, route to this tier.""" - keywords: List[str] = Field( + keywords: list[str] = Field( min_length=1, description="Keywords/phrases that trigger this rule (lexical or semantic match)", ) @@ -56,7 +60,7 @@ class KeywordTierRule(BaseModel): # Note: Keywords should be full words/phrases to avoid substring false positives. # The matching logic uses word boundary detection for single-word keywords. -DEFAULT_CODE_KEYWORDS: List[str] = [ +DEFAULT_CODE_KEYWORDS: list[str] = [ "function", "class", "def", @@ -104,7 +108,7 @@ DEFAULT_CODE_KEYWORDS: List[str] = [ "pull request", ] -DEFAULT_REASONING_KEYWORDS: List[str] = [ +DEFAULT_REASONING_KEYWORDS: list[str] = [ "step by step", "think through", "let's think", @@ -126,7 +130,7 @@ DEFAULT_REASONING_KEYWORDS: List[str] = [ "conclude", ] -DEFAULT_TECHNICAL_KEYWORDS: List[str] = [ +DEFAULT_TECHNICAL_KEYWORDS: list[str] = [ "architecture", "distributed", "scalable", @@ -158,7 +162,7 @@ DEFAULT_TECHNICAL_KEYWORDS: List[str] = [ # Note: "async", "kubernetes", "docker" are in DEFAULT_CODE_KEYWORDS ] -DEFAULT_SIMPLE_KEYWORDS: List[str] = [ +DEFAULT_SIMPLE_KEYWORDS: list[str] = [ "what is", "what's", "define", @@ -191,7 +195,7 @@ DEFAULT_SIMPLE_KEYWORDS: List[str] = [ # ─── Default Dimension Weights ─── -DEFAULT_DIMENSION_WEIGHTS: Dict[str, float] = { +DEFAULT_DIMENSION_WEIGHTS: dict[str, float] = { "tokenCount": 0.10, # Reduced - length is less important than content "codePresence": 0.30, # High - code requests need capable models "reasoningMarkers": 0.25, # High - explicit reasoning requests @@ -204,7 +208,7 @@ DEFAULT_DIMENSION_WEIGHTS: Dict[str, float] = { # ─── Default Tier Boundaries ─── -DEFAULT_TIER_BOUNDARIES: Dict[str, float] = { +DEFAULT_TIER_BOUNDARIES: dict[str, float] = { "simple_medium": 0.15, # Lower threshold to catch more MEDIUM cases "medium_complex": 0.35, # Lower threshold to catch technical COMPLEX cases "complex_reasoning": 0.60, # Reasoning tier reserved for explicit reasoning markers @@ -213,7 +217,7 @@ DEFAULT_TIER_BOUNDARIES: Dict[str, float] = { # ─── Default Token Thresholds ─── -DEFAULT_TOKEN_THRESHOLDS: Dict[str, int] = { +DEFAULT_TOKEN_THRESHOLDS: dict[str, int] = { "simple": 15, # Only very short prompts (<15 tokens) are penalized "complex": 400, # Long prompts (>400 tokens) get complexity boost } @@ -221,7 +225,7 @@ DEFAULT_TOKEN_THRESHOLDS: Dict[str, int] = { # ─── Default Tier to Model Mapping ─── -DEFAULT_TIER_MODELS: Dict[str, str] = { +DEFAULT_TIER_MODELS: dict[str, str] = { "SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o", "COMPLEX": "claude-sonnet-4-20250514", @@ -244,46 +248,47 @@ class ClassifierLLMConfig(BaseModel): class ComplexityRouterConfig(BaseModel): """Configuration for the ComplexityRouter.""" - # string = pin; list = random pick from the tier pool + # string = pin; list = random pick when adaptive=False, soft-floor home pool when adaptive=True tiers: dict[str, str | list[str]] = Field( default_factory=lambda: DEFAULT_TIER_MODELS.copy(), description=( - "Mapping of complexity tiers to a model or model pool. A list is randomly picked from for that tier" + "Mapping of complexity tiers to a model or model pool. " + "A list is randomly picked from when adaptive=False, and used as a soft-floor home pool when adaptive=True" ), ) # Tier boundaries (normalized scores) - tier_boundaries: Dict[str, float] = Field( + tier_boundaries: dict[str, float] = Field( default_factory=lambda: DEFAULT_TIER_BOUNDARIES.copy(), description="Score boundaries between tiers", ) # Token count thresholds - token_thresholds: Dict[str, int] = Field( + token_thresholds: dict[str, int] = Field( default_factory=lambda: DEFAULT_TOKEN_THRESHOLDS.copy(), description="Token count thresholds for simple/complex classification", ) # Dimension weights - dimension_weights: Dict[str, float] = Field( + dimension_weights: dict[str, float] = Field( default_factory=lambda: DEFAULT_DIMENSION_WEIGHTS.copy(), description="Weights for each scoring dimension", ) # Keyword lists (overridable) - code_keywords: Optional[List[str]] = Field( + code_keywords: list[str] | None = Field( default=None, description="Keywords indicating code-related content", ) - reasoning_keywords: Optional[List[str]] = Field( + reasoning_keywords: list[str] | None = Field( default=None, description="Keywords indicating reasoning-required content", ) - technical_keywords: Optional[List[str]] = Field( + technical_keywords: list[str] | None = Field( default=None, description="Keywords indicating technical content", ) - custom_technical_keywords: Optional[list[str]] = Field( + custom_technical_keywords: list[str] | None = Field( default=None, description=( "Domain-specific technical keywords appended to the effective base list " @@ -292,13 +297,13 @@ class ComplexityRouterConfig(BaseModel): "the base list and within this list." ), ) - simple_keywords: Optional[List[str]] = Field( + simple_keywords: list[str] | None = Field( default=None, description="Keywords indicating simple/basic queries", ) # Default model if scoring fails - default_model: Optional[str] = Field( + default_model: str | None = Field( default=None, description="Default model to use if tier cannot be determined", ) @@ -308,13 +313,34 @@ class ComplexityRouterConfig(BaseModel): default="heuristic", description="Classification strategy: local regex/keyword scoring, or an LLM call", ) - classifier_llm_config: Optional[ClassifierLLMConfig] = Field( + classifier_llm_config: ClassifierLLMConfig | None = Field( default=None, description="Configuration for the LLM classifier; required when classifier_type is 'llm'", ) + adaptive: bool = Field( + default=False, + description="Enable adaptive bandit selection with soft complexity floors", + ) + adaptive_weights: AdaptiveRouterWeights = Field( + default_factory=lambda: AdaptiveRouterWeights(quality=0.3, cost=0.7), + description="Quality vs cost weights for adaptive selection (used when adaptive=True)", + ) + tier_distance_penalty: float = Field( + default=DEFAULT_TIER_DISTANCE_PENALTY, + ge=0.0, + description="Score penalty per tier-step away from the classified tier when adaptive=True", + ) + adaptive_eligible: Literal["all", "classified_tier"] = Field( + default="all", + description=( + "When adaptive=True: 'all' scores every pool model with a tier-distance penalty (soft floors); " + "'classified_tier' Thompson-samples only inside the classified tier's pool" + ), + ) + # Deterministic keyword -> tier overrides, evaluated before weighted scoring - keyword_tier_rules: Optional[List[KeywordTierRule]] = Field( + keyword_tier_rules: list[KeywordTierRule] | None = Field( default=None, description="Rules that force a specific tier when their keywords match the prompt", ) @@ -324,7 +350,7 @@ class ComplexityRouterConfig(BaseModel): default=False, description="Match keyword_tier_rules by embedding similarity instead of literal text", ) - embedding_model: Optional[str] = Field( + embedding_model: str | None = Field( default=None, description="Embedding model (LiteLLM model name) used when semantic_keyword_matching is enabled", ) @@ -358,6 +384,19 @@ class ComplexityRouterConfig(BaseModel): raise ValueError("classifier_llm_config is required when classifier_type is 'llm'") return self + @model_validator(mode="after") + def _validate_adaptive_pools(self) -> "ComplexityRouterConfig": + if not self.adaptive: + return self + normalized = {tier: (models if isinstance(models, list) else [models]) for tier, models in self.tiers.items()} + if not any(normalized.values()): + raise ValueError("adaptive=True requires at least one non-empty tier pool") + empty = [tier for tier, models in normalized.items() if not models] + if empty: + raise ValueError(f"adaptive=True tier pools must be non-empty; empty tiers: {empty}") + self.tiers = normalized + return self + @model_validator(mode="after") def _validate_semantic_matching(self) -> "ComplexityRouterConfig": if not self.semantic_keyword_matching: diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 7750ac6628a..dcde6fd1641 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -306,7 +306,7 @@ "limit": 9 }, "TID251": { - "limit": 2710 + "limit": 2701 }, "TRY002": { "limit": 548 @@ -324,7 +324,7 @@ "limit": 883 }, "UP006": { - "limit": 12869 + "limit": 12792 }, "UP007": { "limit": 2570 @@ -354,7 +354,7 @@ "limit": 4 }, "UP035": { - "limit": 2295 + "limit": 2284 }, "UP036": { "limit": 4 @@ -363,6 +363,6 @@ "limit": 105 }, "UP045": { - "limit": 18517 + "limit": 18462 } } diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py index 93c4db90dad..cbf5635a5ae 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py @@ -7,9 +7,6 @@ from litellm.router_strategy.adaptive_router import adaptive_router as ar_module import pytest from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter -from litellm.router_strategy.adaptive_router.config import ( - OWNER_CACHE_TTL_SECONDS, -) from litellm.router_strategy.adaptive_router.signals import Turn from litellm.types.router import ( AdaptiveRouterConfig, @@ -22,9 +19,7 @@ def _make_router() -> AdaptiveRouter: cfg = AdaptiveRouterConfig(available_models=["fast", "smart"]) prefs = { "fast": AdaptiveRouterPreferences(quality_tier=1, strengths=[]), - "smart": AdaptiveRouterPreferences( - quality_tier=3, strengths=[RequestType.CODE_GENERATION] - ), + "smart": AdaptiveRouterPreferences(quality_tier=3, strengths=[RequestType.CODE_GENERATION]), } costs = {"fast": 0.0001, "smart": 0.001} return AdaptiveRouter( @@ -58,85 +53,6 @@ async def test_pick_model_min_quality_tier_filter_raises_when_no_eligible(): await r.pick_model(RequestType.GENERAL, min_quality_tier=4) -@pytest.mark.asyncio -async def test_pick_model_is_stateless_no_owner_cache_writes(): - """pick_model must not touch the owner cache — that's gated post-call.""" - r = _make_router() - for _ in range(5): - await r.pick_model(RequestType.GENERAL) - assert r._owner_cache == {} - - -# ---- claim_or_check_owner ----------------------------------------------- - - -def test_claim_or_check_owner_first_call_claims_and_returns_true(monkeypatch): - r = _make_router() - monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0) - - assert r.claim_or_check_owner("sess-A", "fast") is True - assert r._owner_cache["sess-A"] == ("fast", 1_000.0 + OWNER_CACHE_TTL_SECONDS) - assert r._skipped_updates_total == 0 - - -def test_claim_or_check_owner_same_model_returns_true_without_extending_ttl( - monkeypatch, -): - r = _make_router() - monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0) - r.claim_or_check_owner("sess-A", "fast") - original_expiry = r._owner_cache["sess-A"][1] - - monkeypatch.setattr(ar_module.time, "time", lambda: 1_500.0) - assert r.claim_or_check_owner("sess-A", "fast") is True - # No extension on hit — owner cache snapshots the first claim. - assert r._owner_cache["sess-A"][1] == original_expiry - - -def test_claim_or_check_owner_mismatch_skips_and_increments_counter(monkeypatch): - r = _make_router() - monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0) - r.claim_or_check_owner("sess-A", "fast") - - assert r.claim_or_check_owner("sess-A", "smart") is False - assert r._skipped_updates_total == 1 - # Owner unchanged. - assert r._owner_cache["sess-A"][0] == "fast" - - -def test_claim_or_check_owner_expired_owner_reclaims_for_new_model(monkeypatch): - r = _make_router() - monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0) - r.claim_or_check_owner("sess-A", "fast") - - monkeypatch.setattr( - ar_module.time, "time", lambda: 1_000.0 + OWNER_CACHE_TTL_SECONDS + 1 - ) - assert r.claim_or_check_owner("sess-A", "smart") is True - assert r._owner_cache["sess-A"][0] == "smart" - # Reclaim isn't a skip. - assert r._skipped_updates_total == 0 - - -def test_owner_cache_evicts_expired_entries_when_threshold_crossed(monkeypatch): - """Past _OWNER_CACHE_SWEEP_THRESHOLD live entries, new claims sweep stale.""" - r = _make_router() - monkeypatch.setattr(ar_module, "_OWNER_CACHE_SWEEP_THRESHOLD", 5) - monkeypatch.setattr(ar_module.time, "time", lambda: 1_000.0) - for i in range(5): - r.claim_or_check_owner(f"old-{i}", "fast") - assert len(r._owner_cache) == 5 - - # Jump past TTL so all "old-*" entries are now expired. - monkeypatch.setattr( - ar_module.time, "time", lambda: 1_000.0 + OWNER_CACHE_TTL_SECONDS + 1 - ) - r.claim_or_check_owner("new-1", "fast") - # Sweep ran -> only the new entry remains. - assert "new-1" in r._owner_cache - assert all(k.startswith("new-") for k in r._owner_cache) - - # ---- record_turn -------------------------------------------------------- @@ -185,9 +101,7 @@ async def test_record_turn_satisfaction_increments_alpha(): # Prime with 2 prior turns to clear the MIN_TURNS_FOR_CLEAN_CREDIT gate. # Use distinct content to avoid incidentally firing stagnation/misalignment. priming_turns = [ - Turn( - user_content="alpha bravo charlie", assistant_content="delta echo foxtrot" - ), + Turn(user_content="alpha bravo charlie", assistant_content="delta echo foxtrot"), Turn( user_content="golf hotel india juliet", assistant_content="kilo lima mike november", @@ -232,6 +146,128 @@ async def test_record_turn_failure_increments_beta(): assert cell_after.alpha == pytest.approx(cell_before.alpha) +@pytest.mark.asyncio +async def test_record_turn_detects_exhaustion_in_tool_results(): + r = _make_router() + + delta = await r.record_turn( + session_id="exhausted", + model_name="smart", + request_type=RequestType.GENERAL, + turn=Turn(tool_results=[{"content": "rate limit exceeded"}]), + ) + + assert delta.exhaustion == 1 + assert r._session_states[("exhausted", "smart")].exhaustion_count == 1 + + +@pytest.mark.asyncio +async def test_record_turn_attributes_user_feedback_to_previous_response_model(): + r = _make_router() + fast_before = r._cells[(RequestType.CODE_GENERATION, "fast")] + smart_before = r._cells[(RequestType.GENERAL, "smart")] + + await r.record_turn( + session_id="feedback-switch", + model_name="fast", + request_type=RequestType.CODE_GENERATION, + turn=Turn( + user_content="fix this python retry bug", + assistant_content="clear the cache on every retry", + ), + ) + await r.record_turn( + session_id="feedback-switch", + model_name="smart", + request_type=RequestType.GENERAL, + turn=Turn( + user_content="the python fix is still broken", + assistant_content="keep successful cache entries", + ), + ) + + fast_after = r._cells[(RequestType.CODE_GENERATION, "fast")] + smart_after = r._cells[(RequestType.GENERAL, "smart")] + assert fast_after.beta == pytest.approx(fast_before.beta + 1.0) + assert smart_after.beta == pytest.approx(smart_before.beta) + snapshot = await r.get_state_snapshot() + assert snapshot["feedback_attributed_total"] == 1 + assert snapshot["cross_model_feedback_total"] == 1 + assert snapshot["feedback_without_context_total"] == 0 + + +@pytest.mark.asyncio +async def test_record_turn_attributes_satisfaction_to_previous_response_model(): + r = _make_router() + await r.record_turn( + session_id="satisfaction-switch", + model_name="smart", + request_type=RequestType.CODE_GENERATION, + turn=Turn( + user_content="write a python retry helper", + assistant_content="first draft", + ), + ) + await r.record_turn( + session_id="satisfaction-switch", + model_name="fast", + request_type=RequestType.CODE_GENERATION, + turn=Turn( + user_content="add exponential backoff to the python helper", + assistant_content="updated draft", + ), + ) + fast_before = r._cells[(RequestType.CODE_GENERATION, "fast")] + smart_before = r._cells[(RequestType.GENERAL, "smart")] + + await r.record_turn( + session_id="satisfaction-switch", + model_name="smart", + request_type=RequestType.GENERAL, + turn=Turn( + user_content="thanks, that worked", + assistant_content="glad to help", + ), + ) + + fast_after = r._cells[(RequestType.CODE_GENERATION, "fast")] + smart_after = r._cells[(RequestType.GENERAL, "smart")] + assert fast_after.alpha == pytest.approx(fast_before.alpha + 1.0) + assert smart_after.alpha == pytest.approx(smart_before.alpha) + + +@pytest.mark.asyncio +async def test_record_turn_bounds_feedback_contexts_and_evicts_least_recent_session(): + r = _make_router() + context_limit = ar_module._FEEDBACK_CONTEXT_MAX_ENTRIES + + for index in range(context_limit): + await r.record_turn( + session_id=f"session-{index}", + model_name="fast", + request_type=RequestType.GENERAL, + turn=Turn(user_content="question", assistant_content="answer"), + ) + + await r.record_turn( + session_id="session-0", + model_name="fast", + request_type=RequestType.GENERAL, + turn=Turn(user_content="follow up", assistant_content="updated answer"), + ) + await r.record_turn( + session_id="overflow", + model_name="fast", + request_type=RequestType.GENERAL, + turn=Turn(user_content="question", assistant_content="answer"), + ) + + assert len(r._feedback_contexts) == context_limit + assert "session-0" in r._feedback_contexts + assert "session-1" not in r._feedback_contexts + assert "overflow" in r._feedback_contexts + + @pytest.mark.asyncio async def test_load_state_from_db_overrides_cold_start(): r = _make_router() @@ -270,9 +306,7 @@ async def test_load_state_from_db_handles_unknown_request_type(): good_row.beta = 3.0 prisma = MagicMock() - prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock( - return_value=[bad_row, good_row] - ) + prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[bad_row, good_row]) await r.load_state_from_db(prisma) # Unknown skipped; good applied. diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py index 9786832b4ae..3071f916ef1 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py @@ -110,23 +110,6 @@ async def test_pick_record_flush_full_cycle(): assert session_call.kwargs["data"]["create"]["model_name"] == chosen -@pytest.mark.asyncio -async def test_owner_cache_pins_attribution_to_first_picked_model(): - """First call claims ownership; matching model returns True, mismatch False.""" - router = _make_router() - chosen = await router.pick_model(RequestType.GENERAL) - assert router.claim_or_check_owner("sess-own", chosen) is True - - # Same model on later turns keeps attributing. - for _ in range(5): - assert router.claim_or_check_owner("sess-own", chosen) is True - - # A different model on a later turn is rejected. - other = "gpt-4o" if chosen == "gpt-4o-mini" else "gpt-4o-mini" - assert router.claim_or_check_owner("sess-own", other) is False - assert router._skipped_updates_total == 1 - - @pytest.mark.asyncio async def test_pick_model_returns_valid_models_without_error(): router = _make_router() diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py b/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py index a2b85f2ce53..ad61f43c5a0 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py @@ -16,10 +16,9 @@ from litellm.router_strategy.adaptive_router.hooks import ( from litellm.router_strategy.adaptive_router.signals import Turn -def _make_hook(claim: bool = True) -> AdaptiveRouterPostCallHook: +def _make_hook() -> AdaptiveRouterPostCallHook: fake_router = MagicMock() fake_router.record_turn = AsyncMock() - fake_router.claim_or_check_owner = MagicMock(return_value=claim) return AdaptiveRouterPostCallHook(adaptive_router=fake_router) @@ -151,7 +150,24 @@ async def test_hook_skips_when_below_signal_gate(): kwargs = _kwargs(messages=short) await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0) hook.adaptive_router.record_turn.assert_not_awaited() - hook.adaptive_router.claim_or_check_owner.assert_not_called() + + +@pytest.mark.asyncio +async def test_hook_tracks_short_conversation_with_explicit_session_id(): + hook = _make_hook() + kwargs = _kwargs( + messages=[{"role": "user", "content": "hi"}], + extra_litellm_params={"litellm_session_id": "explicit-short"}, + ) + await hook.async_log_success_event( + kwargs, + _resp_with_content("hello"), + 0.0, + 1.0, + ) + assert hook.adaptive_router.record_turn.await_args.kwargs["session_id"] == ( + "explicit-short" + ) @pytest.mark.asyncio @@ -168,22 +184,19 @@ async def test_hook_skips_when_chosen_model_missing_from_metadata(): kwargs = _kwargs(chosen=None) await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0) hook.adaptive_router.record_turn.assert_not_awaited() - hook.adaptive_router.claim_or_check_owner.assert_not_called() @pytest.mark.asyncio -async def test_hook_skips_when_owner_cache_mismatch(): - """A different model owns this conversation -> no attribution.""" - hook = _make_hook(claim=False) +async def test_hook_records_when_model_changes(): + hook = _make_hook() kwargs = _kwargs(chosen="fast") await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0) - hook.adaptive_router.claim_or_check_owner.assert_called_once() - hook.adaptive_router.record_turn.assert_not_awaited() + hook.adaptive_router.record_turn.assert_awaited_once() @pytest.mark.asyncio -async def test_hook_records_turn_when_owner_claims(): - hook = _make_hook(claim=True) +async def test_hook_records_turn(): + hook = _make_hook() kwargs = _kwargs(chosen="smart", messages=_long_messages("ask")) await hook.async_log_success_event( kwargs, _resp_with_content("answer here"), 0.0, 1.0 @@ -205,8 +218,6 @@ async def test_hook_uses_explicit_session_id_when_provided(): extra_litellm_params={"litellm_session_id": "explicit-sess"}, ) await hook.async_log_success_event(kwargs, _resp_with_content("ok"), 0.0, 1.0) - args, _ = hook.adaptive_router.claim_or_check_owner.call_args - assert args[0] == "explicit-sess" assert hook.adaptive_router.record_turn.await_args.kwargs["session_id"] == ( "explicit-sess" ) diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py b/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py index 753a449791b..d6d89c8e811 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py @@ -1,7 +1,6 @@ """Tests for the GET /adaptive_router/state introspection endpoint and the underlying `AdaptiveRouter.get_state_snapshot()` helper.""" -import time from unittest.mock import MagicMock import pytest @@ -9,7 +8,7 @@ from fastapi import HTTPException from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.router_strategy.adaptive_router.adaptive_router import AdaptiveRouter -from litellm.router_strategy.adaptive_router.bandit import BanditCell, apply_delta +from litellm.router_strategy.adaptive_router.bandit import apply_delta from litellm.types.router import ( AdaptiveRouterConfig, AdaptiveRouterPreferences, @@ -47,8 +46,6 @@ async def test_get_state_snapshot_returns_cell_per_request_type_per_model(): assert snap["available_models"] == ["fast", "smart"] assert snap["weights"] == {"quality": 0.7, "cost": 0.3} assert snap["model_costs"] == {"fast": 0.0001, "smart": 0.001} - assert snap["owner_cache_live"] == 0 - assert snap["skipped_updates_total"] == 0 assert set(snap["queue"].keys()) == { "state_pending", "session_pending", @@ -95,26 +92,6 @@ async def test_get_state_snapshot_quality_mean_matches_alpha_over_total(): assert cell["quality_mean"] == pytest.approx(expected_mean) -@pytest.mark.asyncio -async def test_get_state_snapshot_counts_only_live_owner_cache_entries(): - r = _make_router() - now = time.time() - r._owner_cache["live-1"] = ("fast", now + 3600) - r._owner_cache["live-2"] = ("smart", now + 3600) - r._owner_cache["expired-1"] = ("fast", now - 1) - - snap = await r.get_state_snapshot() - assert snap["owner_cache_live"] == 2 - - -@pytest.mark.asyncio -async def test_get_state_snapshot_exposes_skipped_updates_total(): - r = _make_router() - r._skipped_updates_total = 7 - snap = await r.get_state_snapshot() - assert snap["skipped_updates_total"] == 7 - - # ---- endpoint -------------------------------------------------------- diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index e1133620a57..da02b774e41 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -948,6 +948,60 @@ class TestRouterComplexityDeploymentMethods: router.init_complexity_router_deployment(deployment) assert "auto_router/complexity_router/test-router" in router.complexity_routers + def test_hybrid_initialization_waits_for_later_pool_deployments(self): + router = Router( + model_list=[ + { + "model_name": "hybrid", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_default_model": "cheap", + "complexity_router_config": { + "adaptive": True, + "tiers": { + "SIMPLE": ["cheap"], + "MEDIUM": ["cheap", "premium"], + }, + }, + }, + }, + { + "model_name": "cheap", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "input_cost_per_token": 0.00000015, + }, + "model_info": { + "adaptive_router_preferences": { + "quality_tier": 1, + "strengths": [], + } + }, + }, + { + "model_name": "premium", + "litellm_params": { + "model": "openai/gpt-4o", + "input_cost_per_token": 0.000005, + }, + "model_info": { + "adaptive_router_preferences": { + "quality_tier": 3, + "strengths": [], + } + }, + }, + ] + ) + + adaptive = router.adaptive_routers["hybrid"] + assert adaptive.model_to_cost == { + "cheap": pytest.approx(0.00000015), + "premium": pytest.approx(0.000005), + } + assert adaptive.model_to_prefs["cheap"].quality_tier == 1 + assert adaptive.model_to_prefs["premium"].quality_tier == 3 + class TestAsyncPreRoutingHookMultiFormat: """Test async_pre_routing_hook with multiple input formats.""" @@ -1356,6 +1410,240 @@ class TestLLMClassifier: assert call_kwargs["metadata"] == request_metadata +class TestAdaptiveSoftFloors: + def test_adaptive_defaults_use_cost_weighted_cold_policy(self): + config = ComplexityRouterConfig( + adaptive=True, + tiers={"SIMPLE": ["cheap"]}, + ) + assert config.adaptive_weights.quality == pytest.approx(0.3) + assert config.adaptive_weights.cost == pytest.approx(0.7) + assert config.tier_distance_penalty == pytest.approx(0.5) + + @pytest.fixture + def adaptive_router_instance(self): + router = MagicMock() + router.model_list = [ + { + "model_name": "cheap", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "input_cost_per_token": 0.00000015, + }, + "model_info": { + "adaptive_router_preferences": {"quality_tier": 1, "strengths": []} + }, + }, + { + "model_name": "premium", + "litellm_params": { + "model": "openai/gpt-4o", + "input_cost_per_token": 0.000005, + }, + "model_info": { + "adaptive_router_preferences": {"quality_tier": 3, "strengths": []} + }, + }, + ] + router.model_name_to_deployment_indices = {"cheap": [0], "premium": [1]} + return router + + @pytest.fixture + def hybrid_config(self) -> Dict: + return { + "adaptive": True, + "adaptive_weights": {"quality": 0.7, "cost": 0.3}, + "tier_distance_penalty": 0.15, + "tiers": { + "SIMPLE": ["cheap"], + "MEDIUM": ["cheap"], + "COMPLEX": ["premium"], + "REASONING": ["premium"], + }, + "default_model": "cheap", + } + + def test_adaptive_config_requires_non_empty_pools(self): + with pytest.raises(ValidationError): + ComplexityRouterConfig(adaptive=True, tiers={"SIMPLE": []}) + + def test_cold_start_randomly_samples_unobserved_classified_tier_models( + self, adaptive_router_instance + ): + cr = ComplexityRouter( + model_name="hybrid", + litellm_router_instance=adaptive_router_instance, + complexity_router_config={ + "adaptive": True, + "tiers": { + "SIMPLE": ["cheap", "premium"], + "MEDIUM": ["premium"], + }, + }, + ) + request_kwargs: Dict = {"metadata": {}} + + with patch( + "litellm.router_strategy.complexity_router.complexity_router.random.choice", + return_value="premium", + ) as choice: + picked = cr._soft_floor_pick(ComplexityTier.SIMPLE, "hi", request_kwargs) + + assert picked == "premium" + choice.assert_called_once_with(("cheap", "premium")) + decision = request_kwargs["metadata"]["adaptive_router_decision"] + assert decision["phase"] == "cold_start" + assert {candidate["model"] for candidate in decision["candidates"]} == { + "cheap", + "premium", + } + + def test_get_model_for_tier_list_without_adaptive_random_choice( + self, mock_router_instance + ): + router = ComplexityRouter( + model_name="test", + litellm_router_instance=mock_router_instance, + complexity_router_config={ + "adaptive": False, + "tiers": {"SIMPLE": ["cheap", "premium"], "MEDIUM": "mid"}, + "default_model": "mid", + }, + ) + pool = ["cheap", "premium"] + with patch( + "litellm.router_strategy.complexity_router.complexity_router.random.choice", + return_value="premium", + ) as choice: + assert router.get_model_for_tier(ComplexityTier.SIMPLE) == "premium" + choice.assert_called_once_with(pool) + assert router.get_model_for_tier(ComplexityTier.MEDIUM) == "mid" + + def test_soft_floor_prefers_home_tier_when_posteriors_equal( + self, adaptive_router_instance, hybrid_config + ): + from litellm.router_strategy.adaptive_router.bandit import BanditCell + from litellm.types.router import RequestType + + cr = ComplexityRouter( + model_name="hybrid", + litellm_router_instance=adaptive_router_instance, + complexity_router_config=hybrid_config, + ) + adaptive = cr._ensure_adaptive_router() + assert adaptive is not None + for model in ("cheap", "premium"): + adaptive._cells[(RequestType.GENERAL, model)] = BanditCell( + alpha=5.0, beta=5.0 + ) + + # Equal quality samples; home-tier penalty should favor cheap for SIMPLE. + with patch( + "litellm.router_strategy.adaptive_router.bandit.thompson_sample", + return_value=0.5, + ): + picked = cr._soft_floor_pick(ComplexityTier.SIMPLE, "hi") + assert picked == "cheap" + + def test_soft_floor_allows_cross_tier_when_posterior_dominates( + self, adaptive_router_instance, hybrid_config + ): + from litellm.router_strategy.adaptive_router.bandit import BanditCell + from litellm.types.router import RequestType + + cr = ComplexityRouter( + model_name="hybrid", + litellm_router_instance=adaptive_router_instance, + complexity_router_config=hybrid_config, + ) + adaptive = cr._ensure_adaptive_router() + assert adaptive is not None + adaptive._cells[(RequestType.GENERAL, "cheap")] = BanditCell( + alpha=1.0, beta=20.0 + ) + adaptive._cells[(RequestType.GENERAL, "premium")] = BanditCell( + alpha=20.0, beta=1.0 + ) + + with patch( + "litellm.router_strategy.adaptive_router.bandit.thompson_sample", + side_effect=lambda cell, rng=None: cell.alpha / (cell.alpha + cell.beta), + ): + picked = cr._soft_floor_pick(ComplexityTier.SIMPLE, "hi") + assert picked == "premium" + + def test_reused_model_has_zero_distance_in_each_configured_tier( + self, adaptive_router_instance + ): + from litellm.router_strategy.adaptive_router.bandit import BanditCell + from litellm.types.router import RequestType + + cr = ComplexityRouter( + model_name="hybrid", + litellm_router_instance=adaptive_router_instance, + complexity_router_config={ + "adaptive": True, + "tiers": { + "SIMPLE": ["cheap"], + "MEDIUM": ["cheap", "premium"], + "COMPLEX": ["premium"], + }, + }, + ) + adaptive = cr._ensure_adaptive_router() + assert adaptive is not None + for model in ("cheap", "premium"): + adaptive._cells[(RequestType.GENERAL, model)] = BanditCell( + alpha=6.0, beta=5.0 + ) + request_kwargs: Dict = {"metadata": {}} + + with patch( + "litellm.router_strategy.adaptive_router.bandit.thompson_sample", + return_value=0.5, + ): + cr._soft_floor_pick(ComplexityTier.MEDIUM, "hi", request_kwargs) + + candidates = request_kwargs["metadata"]["adaptive_router_decision"][ + "candidates" + ] + assert { + candidate["model"]: candidate["tier_distance"] for candidate in candidates + } == { + "cheap": 0, + "premium": 0, + } + + @pytest.mark.asyncio + async def test_pre_routing_hook_adaptive_stashes_chosen_model( + self, adaptive_router_instance, hybrid_config + ): + cr = ComplexityRouter( + model_name="hybrid", + litellm_router_instance=adaptive_router_instance, + complexity_router_config=hybrid_config, + ) + request_kwargs: Dict = {"metadata": {}} + result = await cr.async_pre_routing_hook( + model="hybrid", + request_kwargs=request_kwargs, + messages=[{"role": "user", "content": "hi"}], + ) + assert result is not None + assert result.model in {"cheap", "premium"} + assert ( + request_kwargs["metadata"].get("adaptive_router_chosen_model") + == result.model + ) + decision = request_kwargs["metadata"]["adaptive_router_decision"] + assert decision["phase"] == "cold_start" + assert decision["classified_tier"] == "SIMPLE" + assert decision["request_type"] == "general" + assert decision["eligible_mode"] == "classified_tier" + assert decision["chosen_model"] == result.model + assert {candidate["model"] for candidate in decision["candidates"]} == {"cheap"} + + class TestLexicalKeywordTierRules: """Test deterministic (literal) keyword_tier_rules overrides.""" diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 50168056eaa..c31ee41a6ec 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -111,7 +111,7 @@ const ComplexityRouterConfig: React.FC = ({ classifier_type: classifierType, classifier_llm_config: classifierType === "llm" - ? value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS } + ? (value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }) : undefined, }); }; diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index 3eddca8c35b..a4a8ee6b074 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -43,8 +43,7 @@ export const getSemanticConfigError = ({ embeddingModel, keywordTierRules, }: Pick): - | string - | null => { + string | null => { if (!semanticMatchingEnabled) return null; if (!embeddingModel) return "Select an embedding model to use semantic keyword matching"; if (keywordTierRules.length === 0) return "Add at least one keyword tier rule to use semantic keyword matching"; From 3a2d14e1a6f19d7b0cc30169f8b5260aa642b73c Mon Sep 17 00:00:00 2001 From: Thibault Serot Date: Mon, 13 Jul 2026 11:33:02 +1000 Subject: [PATCH 058/123] fix(responses): continue MCP gateway tool turns from the final response and surface failures When a /responses request uses a hosted MCP tool (server_url: litellm_proxy/
)} +
+ Model Aliases + {(() => { + const aliasEntries = Object.entries(info.litellm_model_table?.model_aliases ?? {}); + if (aliasEntries.length === 0) { + return
No model aliases configured
; + } + return ( +
+ {aliasEntries.map(([alias, target]) => ( +
+ {alias} + {" -> "} + {target} +
+ ))} +
+ ); + })()} +
Rate Limits
TPM: {info.tpm_limit || "Unlimited"}
From 20e646c49a6c3ef5ce5c6957b807fa52d0fa7fe3 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Mon, 13 Jul 2026 10:20:07 -0700 Subject: [PATCH 063/123] fix(ci): bump pillow to 12.3.0 to resolve osv-scan CVEs (#33093) --- pyproject.toml | 2 +- uv.lock | 117 ++++++++++++++++++++----------------------------- 2 files changed, 49 insertions(+), 70 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index c2a0e85d13a..2c796d14c16 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -205,7 +205,7 @@ ci = [ # protobuf, Pillow is a compiled C extension). "tenacity==8.5.0", "google-generativeai==0.8.6", - "Pillow==12.2.0", + "Pillow==12.3.0", # Azure batch E2E tests still import psycopg2 directly. "psycopg2-binary==2.9.11", "pytest-codspeed==4.3.0", diff --git a/uv.lock b/uv.lock index 00db09ef4ec..b120547c536 100644 --- a/uv.lock +++ b/uv.lock @@ -9,7 +9,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-07-08T23:20:11.959202Z" +exclude-newer = "2026-07-10T16:47:58.286372Z" exclude-newer-span = "P3D" [manifest] @@ -3588,7 +3588,7 @@ ci = [ { name = "logfire", specifier = "==4.6.0" }, { name = "lunary", marker = "python_full_version == '3.10.*'", specifier = "==1.4.36" }, { name = "lunary", marker = "python_full_version >= '3.11'", specifier = "==1.4.37" }, - { name = "pillow", specifier = "==12.2.0" }, + { name = "pillow", specifier = "==12.3.0" }, { name = "psycopg2-binary", specifier = "==2.9.11" }, { name = "pyarrow", specifier = "==23.0.1" }, { name = "pygithub", specifier = "==2.8.1" }, @@ -5279,75 +5279,54 @@ wheels = [ [[package]] name = "pillow" -version = "12.2.0" +version = "12.3.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/8c/21/c2bcdd5906101a30244eaffc1b6e6ce71a31bd0742a01eb89e660ebfac2d/pillow-12.2.0.tar.gz", hash = "sha256:a830b1a40919539d07806aa58e1b114df53ddd43213d9c8b75847eee6c0182b5", size = 46987819, upload-time = "2026-04-01T14:46:17.687Z" } +sdist = { url = "https://files.pythonhosted.org/packages/1c/3d/bb7fca845737cf9d7dbde16ed1843984665ff2e0a518f5db43e77ec540b9/pillow-12.3.0.tar.gz", hash = "sha256:3b8182a766685eaa002637e28b4ec8d6b18819a0c71f579bf0dbaa5830297cce", size = 47025035, upload-time = "2026-07-01T11:56:38.965Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/3a/aa/d0b28e1c811cd4d5f5c2bfe2e022292bd255ae5744a3b9ac7d6c8f72dd75/pillow-12.2.0-cp310-cp310-macosx_10_10_x86_64.whl", hash = "sha256:a4e8f36e677d3336f35089648c8955c51c6d386a13cf6ee9c189c5f5bd713a9f", size = 5354355, upload-time = "2026-04-01T14:42:15.402Z" }, - { url = "https://files.pythonhosted.org/packages/27/8e/1d5b39b8ae2bd7650d0c7b6abb9602d16043ead9ebbfef4bc4047454da2a/pillow-12.2.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:2e589959f10d9824d39b350472b92f0ce3b443c0a3442ebf41c40cb8361c5b97", size = 4695871, upload-time = "2026-04-01T14:42:18.234Z" }, - { url = "https://files.pythonhosted.org/packages/f0/c5/dcb7a6ca6b7d3be41a76958e90018d56c8462166b3ef223150360850c8da/pillow-12.2.0-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:a52edc8bfff4429aaabdf4d9ee0daadbbf8562364f940937b941f87a4290f5ff", size = 6269734, upload-time = "2026-04-01T14:42:20.608Z" }, - { url = "https://files.pythonhosted.org/packages/ea/f1/aa1bb13b2f4eba914e9637893c73f2af8e48d7d4023b9d3750d4c5eb2d0c/pillow-12.2.0-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:975385f4776fafde056abb318f612ef6285b10a1f12b8570f3647ad0d74b48ec", size = 8076080, upload-time = "2026-04-01T14:42:23.095Z" }, - { url = "https://files.pythonhosted.org/packages/a1/2a/8c79d6a53169937784604a8ae8d77e45888c41537f7f6f65ed1f407fe66d/pillow-12.2.0-cp310-cp310-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:bd9c0c7a0c681a347b3194c500cb1e6ca9cab053ea4d82a5cf45b6b754560136", size = 6382236, upload-time = "2026-04-01T14:42:25.82Z" }, - { url = "https://files.pythonhosted.org/packages/b5/42/bbcb6051030e1e421d103ce7a8ecadf837aa2f39b8f82ef1a8d37c3d4ebc/pillow-12.2.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:88d387ff40b3ff7c274947ed3125dedf5262ec6919d83946753b5f3d7c67ea4c", size = 7070220, upload-time = "2026-04-01T14:42:28.68Z" }, - { url = "https://files.pythonhosted.org/packages/3f/e1/c2a7d6dd8cfa6b231227da096fd2d58754bab3603b9d73bf609d3c18b64f/pillow-12.2.0-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:51c4167c34b0d8ba05b547a3bb23578d0ba17b80a5593f93bd8ecb123dd336a3", size = 6493124, upload-time = "2026-04-01T14:42:31.579Z" }, - { url = "https://files.pythonhosted.org/packages/5f/41/7c8617da5d32e1d2f026e509484fdb6f3ad7efaef1749a0c1928adbb099e/pillow-12.2.0-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:34c0d99ecccea270c04882cb3b86e7b57296079c9a4aff88cb3b33563d95afaa", size = 7194324, upload-time = "2026-04-01T14:42:34.615Z" }, - { url = "https://files.pythonhosted.org/packages/2d/de/a777627e19fd6d62f84070ee1521adde5eeda4855b5cf60fe0b149118bca/pillow-12.2.0-cp310-cp310-win32.whl", hash = "sha256:b85f66ae9eb53e860a873b858b789217ba505e5e405a24b85c0464822fe88032", size = 6376363, upload-time = "2026-04-01T14:42:37.19Z" }, - { url = "https://files.pythonhosted.org/packages/e7/34/fc4cb5204896465842767b96d250c08410f01f2f28afc43b257de842eed5/pillow-12.2.0-cp310-cp310-win_amd64.whl", hash = "sha256:673aa32138f3e7531ccdbca7b3901dba9b70940a19ccecc6a37c77d5fdeb05b5", size = 7083523, upload-time = "2026-04-01T14:42:39.62Z" }, - { url = "https://files.pythonhosted.org/packages/2d/a0/32852d36bc7709f14dc3f64f929a275e958ad8c19a6deba9610d458e28b3/pillow-12.2.0-cp310-cp310-win_arm64.whl", hash = "sha256:3e080565d8d7c671db5802eedfb438e5565ffa40115216eabb8cd52d0ecce024", size = 2463318, upload-time = "2026-04-01T14:42:42.063Z" }, - { url = "https://files.pythonhosted.org/packages/68/e1/748f5663efe6edcfc4e74b2b93edfb9b8b99b67f21a854c3ae416500a2d9/pillow-12.2.0-cp311-cp311-macosx_10_10_x86_64.whl", hash = "sha256:8be29e59487a79f173507c30ddf57e733a357f67881430449bb32614075a40ab", size = 5354347, upload-time = "2026-04-01T14:42:44.255Z" }, - { url = "https://files.pythonhosted.org/packages/47/a1/d5ff69e747374c33a3b53b9f98cca7889fce1fd03d79cdc4e1bccc6c5a87/pillow-12.2.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:71cde9a1e1551df7d34a25462fc60325e8a11a82cc2e2f54578e5e9a1e153d65", size = 4695873, upload-time = "2026-04-01T14:42:46.452Z" }, - { url = "https://files.pythonhosted.org/packages/df/21/e3fbdf54408a973c7f7f89a23b2cb97a7ef30c61ab4142af31eee6aebc88/pillow-12.2.0-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:f490f9368b6fc026f021db16d7ec2fbf7d89e2edb42e8ec09d2c60505f5729c7", size = 6280168, upload-time = "2026-04-01T14:42:49.228Z" }, - { url = "https://files.pythonhosted.org/packages/d3/f1/00b7278c7dd52b17ad4329153748f87b6756ec195ff786c2bdf12518337d/pillow-12.2.0-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:8bd7903a5f2a4545f6fd5935c90058b89d30045568985a71c79f5fd6edf9b91e", size = 8088188, upload-time = "2026-04-01T14:42:51.735Z" }, - { url = "https://files.pythonhosted.org/packages/ad/cf/220a5994ef1b10e70e85748b75649d77d506499352be135a4989c957b701/pillow-12.2.0-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:3997232e10d2920a68d25191392e3a4487d8183039e1c74c2297f00ed1c50705", size = 6394401, upload-time = "2026-04-01T14:42:54.343Z" }, - { url = "https://files.pythonhosted.org/packages/e9/bd/e51a61b1054f09437acfbc2ff9106c30d1eb76bc1453d428399946781253/pillow-12.2.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:e74473c875d78b8e9d5da2a70f7099549f9eb37ded4e2f6a463e60125bccd176", size = 7079655, upload-time = "2026-04-01T14:42:56.954Z" }, - { url = "https://files.pythonhosted.org/packages/6b/3d/45132c57d5fb4b5744567c3817026480ac7fc3ce5d4c47902bc0e7f6f853/pillow-12.2.0-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:56a3f9c60a13133a98ecff6197af34d7824de9b7b38c3654861a725c970c197b", size = 6503105, upload-time = "2026-04-01T14:42:59.847Z" }, - { url = "https://files.pythonhosted.org/packages/7d/2e/9df2fc1e82097b1df3dce58dc43286aa01068e918c07574711fcc53e6fb4/pillow-12.2.0-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:90e6f81de50ad6b534cab6e5aef77ff6e37722b2f5d908686f4a5c9eba17a909", size = 7203402, upload-time = "2026-04-01T14:43:02.664Z" }, - { url = "https://files.pythonhosted.org/packages/bd/2e/2941e42858ebb67e50ae741473de81c2984e6eff7b397017623c676e2e8d/pillow-12.2.0-cp311-cp311-win32.whl", hash = "sha256:8c984051042858021a54926eb597d6ee3012393ce9c181814115df4c60b9a808", size = 6378149, upload-time = "2026-04-01T14:43:05.274Z" }, - { url = "https://files.pythonhosted.org/packages/69/42/836b6f3cd7f3e5fa10a1f1a5420447c17966044c8fbf589cc0452d5502db/pillow-12.2.0-cp311-cp311-win_amd64.whl", hash = "sha256:6e6b2a0c538fc200b38ff9eb6628228b77908c319a005815f2dde585a0664b60", size = 7082626, upload-time = "2026-04-01T14:43:08.557Z" }, - { url = "https://files.pythonhosted.org/packages/c2/88/549194b5d6f1f494b485e493edc6693c0a16f4ada488e5bd974ed1f42fad/pillow-12.2.0-cp311-cp311-win_arm64.whl", hash = "sha256:9a8a34cc89c67a65ea7437ce257cea81a9dad65b29805f3ecee8c8fe8ff25ffe", size = 2463531, upload-time = "2026-04-01T14:43:10.743Z" }, - { url = "https://files.pythonhosted.org/packages/58/be/7482c8a5ebebbc6470b3eb791812fff7d5e0216c2be3827b30b8bb6603ed/pillow-12.2.0-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:2d192a155bbcec180f8564f693e6fd9bccff5a7af9b32e2e4bf8c9c69dbad6b5", size = 5308279, upload-time = "2026-04-01T14:43:13.246Z" }, - { url = "https://files.pythonhosted.org/packages/d8/95/0a351b9289c2b5cbde0bacd4a83ebc44023e835490a727b2a3bd60ddc0f4/pillow-12.2.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:f3f40b3c5a968281fd507d519e444c35f0ff171237f4fdde090dd60699458421", size = 4695490, upload-time = "2026-04-01T14:43:15.584Z" }, - { url = "https://files.pythonhosted.org/packages/de/af/4e8e6869cbed569d43c416fad3dc4ecb944cb5d9492defaed89ddd6fe871/pillow-12.2.0-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:03e7e372d5240cc23e9f07deca4d775c0817bffc641b01e9c3af208dbd300987", size = 6284462, upload-time = "2026-04-01T14:43:18.268Z" }, - { url = "https://files.pythonhosted.org/packages/e9/9e/c05e19657fd57841e476be1ab46c4d501bffbadbafdc31a6d665f8b737b6/pillow-12.2.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:b86024e52a1b269467a802258c25521e6d742349d760728092e1bc2d135b4d76", size = 8094744, upload-time = "2026-04-01T14:43:20.716Z" }, - { url = "https://files.pythonhosted.org/packages/2b/54/1789c455ed10176066b6e7e6da1b01e50e36f94ba584dc68d9eebfe9156d/pillow-12.2.0-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:7371b48c4fa448d20d2714c9a1f775a81155050d383333e0a6c15b1123dda005", size = 6398371, upload-time = "2026-04-01T14:43:23.443Z" }, - { url = "https://files.pythonhosted.org/packages/43/e3/fdc657359e919462369869f1c9f0e973f353f9a9ee295a39b1fea8ee1a77/pillow-12.2.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:62f5409336adb0663b7caa0da5c7d9e7bdbaae9ce761d34669420c2a801b2780", size = 7087215, upload-time = "2026-04-01T14:43:26.758Z" }, - { url = "https://files.pythonhosted.org/packages/8b/f8/2f6825e441d5b1959d2ca5adec984210f1ec086435b0ed5f52c19b3b8a6e/pillow-12.2.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:01afa7cf67f74f09523699b4e88c73fb55c13346d212a59a2db1f86b0a63e8c5", size = 6509783, upload-time = "2026-04-01T14:43:29.56Z" }, - { url = "https://files.pythonhosted.org/packages/67/f9/029a27095ad20f854f9dba026b3ea6428548316e057e6fc3545409e86651/pillow-12.2.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:fc3d34d4a8fbec3e88a79b92e5465e0f9b842b628675850d860b8bd300b159f5", size = 7212112, upload-time = "2026-04-01T14:43:32.091Z" }, - { url = "https://files.pythonhosted.org/packages/be/42/025cfe05d1be22dbfdb4f264fe9de1ccda83f66e4fc3aac94748e784af04/pillow-12.2.0-cp312-cp312-win32.whl", hash = "sha256:58f62cc0f00fd29e64b29f4fd923ffdb3859c9f9e6105bfc37ba1d08994e8940", size = 6378489, upload-time = "2026-04-01T14:43:34.601Z" }, - { url = "https://files.pythonhosted.org/packages/5d/7b/25a221d2c761c6a8ae21bfa3874988ff2583e19cf8a27bf2fee358df7942/pillow-12.2.0-cp312-cp312-win_amd64.whl", hash = "sha256:7f84204dee22a783350679a0333981df803dac21a0190d706a50475e361c93f5", size = 7084129, upload-time = "2026-04-01T14:43:37.213Z" }, - { url = "https://files.pythonhosted.org/packages/10/e1/542a474affab20fd4a0f1836cb234e8493519da6b76899e30bcc5d990b8b/pillow-12.2.0-cp312-cp312-win_arm64.whl", hash = "sha256:af73337013e0b3b46f175e79492d96845b16126ddf79c438d7ea7ff27783a414", size = 2463612, upload-time = "2026-04-01T14:43:39.421Z" }, - { url = "https://files.pythonhosted.org/packages/4a/01/53d10cf0dbad820a8db274d259a37ba50b88b24768ddccec07355382d5ad/pillow-12.2.0-cp313-cp313-ios_13_0_arm64_iphoneos.whl", hash = "sha256:8297651f5b5679c19968abefd6bb84d95fe30ef712eb1b2d9b2d31ca61267f4c", size = 4100837, upload-time = "2026-04-01T14:43:41.506Z" }, - { url = "https://files.pythonhosted.org/packages/0f/98/f3a6657ecb698c937f6c76ee564882945f29b79bad496abcba0e84659ec5/pillow-12.2.0-cp313-cp313-ios_13_0_arm64_iphonesimulator.whl", hash = "sha256:50d8520da2a6ce0af445fa6d648c4273c3eeefbc32d7ce049f22e8b5c3daecc2", size = 4176528, upload-time = "2026-04-01T14:43:43.773Z" }, - { url = "https://files.pythonhosted.org/packages/69/bc/8986948f05e3ea490b8442ea1c1d4d990b24a7e43d8a51b2c7d8b1dced36/pillow-12.2.0-cp313-cp313-ios_13_0_x86_64_iphonesimulator.whl", hash = "sha256:766cef22385fa1091258ad7e6216792b156dc16d8d3fa607e7545b2b72061f1c", size = 3640401, upload-time = "2026-04-01T14:43:45.87Z" }, - { url = "https://files.pythonhosted.org/packages/34/46/6c717baadcd62bc8ed51d238d521ab651eaa74838291bda1f86fe1f864c9/pillow-12.2.0-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:5d2fd0fa6b5d9d1de415060363433f28da8b1526c1c129020435e186794b3795", size = 5308094, upload-time = "2026-04-01T14:43:48.438Z" }, - { url = "https://files.pythonhosted.org/packages/71/43/905a14a8b17fdb1ccb58d282454490662d2cb89a6bfec26af6d3520da5ec/pillow-12.2.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:56b25336f502b6ed02e889f4ece894a72612fe885889a6e8c4c80239ff6e5f5f", size = 4695402, upload-time = "2026-04-01T14:43:51.292Z" }, - { url = "https://files.pythonhosted.org/packages/73/dd/42107efcb777b16fa0393317eac58f5b5cf30e8392e266e76e51cff28c3d/pillow-12.2.0-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:f1c943e96e85df3d3478f7b691f229887e143f81fedab9b20205349ab04d73ed", size = 6280005, upload-time = "2026-04-01T14:43:54.242Z" }, - { url = "https://files.pythonhosted.org/packages/a8/68/b93e09e5e8549019e61acf49f65b1a8530765a7f812c77a7461bca7e4494/pillow-12.2.0-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:03f6fab9219220f041c74aeaa2939ff0062bd5c364ba9ce037197f4c6d498cd9", size = 8090669, upload-time = "2026-04-01T14:43:57.335Z" }, - { url = "https://files.pythonhosted.org/packages/4b/6e/3ccb54ce8ec4ddd1accd2d89004308b7b0b21c4ac3d20fa70af4760a4330/pillow-12.2.0-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5cdfebd752ec52bf5bb4e35d9c64b40826bc5b40a13df7c3cda20a2c03a0f5ed", size = 6395194, upload-time = "2026-04-01T14:43:59.864Z" }, - { url = "https://files.pythonhosted.org/packages/67/ee/21d4e8536afd1a328f01b359b4d3997b291ffd35a237c877b331c1c3b71c/pillow-12.2.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:eedf4b74eda2b5a4b2b2fb4c006d6295df3bf29e459e198c90ea48e130dc75c3", size = 7082423, upload-time = "2026-04-01T14:44:02.74Z" }, - { url = "https://files.pythonhosted.org/packages/78/5f/e9f86ab0146464e8c133fe85df987ed9e77e08b29d8d35f9f9f4d6f917ba/pillow-12.2.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:00a2865911330191c0b818c59103b58a5e697cae67042366970a6b6f1b20b7f9", size = 6505667, upload-time = "2026-04-01T14:44:05.381Z" }, - { url = "https://files.pythonhosted.org/packages/ed/1e/409007f56a2fdce61584fd3acbc2bbc259857d555196cedcadc68c015c82/pillow-12.2.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:1e1757442ed87f4912397c6d35a0db6a7b52592156014706f17658ff58bbf795", size = 7208580, upload-time = "2026-04-01T14:44:08.39Z" }, - { url = "https://files.pythonhosted.org/packages/23/c4/7349421080b12fb35414607b8871e9534546c128a11965fd4a7002ccfbee/pillow-12.2.0-cp313-cp313-win32.whl", hash = "sha256:144748b3af2d1b358d41286056d0003f47cb339b8c43a9ea42f5fea4d8c66b6e", size = 6375896, upload-time = "2026-04-01T14:44:11.197Z" }, - { url = "https://files.pythonhosted.org/packages/3f/82/8a3739a5e470b3c6cbb1d21d315800d8e16bff503d1f16b03a4ec3212786/pillow-12.2.0-cp313-cp313-win_amd64.whl", hash = "sha256:390ede346628ccc626e5730107cde16c42d3836b89662a115a921f28440e6a3b", size = 7081266, upload-time = "2026-04-01T14:44:13.947Z" }, - { url = "https://files.pythonhosted.org/packages/c3/25/f968f618a062574294592f668218f8af564830ccebdd1fa6200f598e65c5/pillow-12.2.0-cp313-cp313-win_arm64.whl", hash = "sha256:8023abc91fba39036dbce14a7d6535632f99c0b857807cbbbf21ecc9f4717f06", size = 2463508, upload-time = "2026-04-01T14:44:16.312Z" }, - { url = "https://files.pythonhosted.org/packages/4d/a4/b342930964e3cb4dce5038ae34b0eab4653334995336cd486c5a8c25a00c/pillow-12.2.0-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:042db20a421b9bafecc4b84a8b6e444686bd9d836c7fd24542db3e7df7baad9b", size = 5309927, upload-time = "2026-04-01T14:44:18.89Z" }, - { url = "https://files.pythonhosted.org/packages/9f/de/23198e0a65a9cf06123f5435a5d95cea62a635697f8f03d134d3f3a96151/pillow-12.2.0-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:dd025009355c926a84a612fecf58bb315a3f6814b17ead51a8e48d3823d9087f", size = 4698624, upload-time = "2026-04-01T14:44:21.115Z" }, - { url = "https://files.pythonhosted.org/packages/01/a6/1265e977f17d93ea37aa28aa81bad4fa597933879fac2520d24e021c8da3/pillow-12.2.0-cp313-cp313t-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:88ddbc66737e277852913bd1e07c150cc7bb124539f94c4e2df5344494e0a612", size = 6321252, upload-time = "2026-04-01T14:44:23.663Z" }, - { url = "https://files.pythonhosted.org/packages/3c/83/5982eb4a285967baa70340320be9f88e57665a387e3a53a7f0db8231a0cd/pillow-12.2.0-cp313-cp313t-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:d362d1878f00c142b7e1a16e6e5e780f02be8195123f164edf7eddd911eefe7c", size = 8126550, upload-time = "2026-04-01T14:44:26.772Z" }, - { url = "https://files.pythonhosted.org/packages/4e/48/6ffc514adce69f6050d0753b1a18fd920fce8cac87620d5a31231b04bfc5/pillow-12.2.0-cp313-cp313t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:2c727a6d53cb0018aadd8018c2b938376af27914a68a492f59dfcaca650d5eea", size = 6433114, upload-time = "2026-04-01T14:44:29.615Z" }, - { url = "https://files.pythonhosted.org/packages/36/a3/f9a77144231fb8d40ee27107b4463e205fa4677e2ca2548e14da5cf18dce/pillow-12.2.0-cp313-cp313t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:efd8c21c98c5cc60653bcb311bef2ce0401642b7ce9d09e03a7da87c878289d4", size = 7115667, upload-time = "2026-04-01T14:44:32.773Z" }, - { url = "https://files.pythonhosted.org/packages/c1/fc/ac4ee3041e7d5a565e1c4fd72a113f03b6394cc72ab7089d27608f8aaccb/pillow-12.2.0-cp313-cp313t-musllinux_1_2_aarch64.whl", hash = "sha256:9f08483a632889536b8139663db60f6724bfcb443c96f1b18855860d7d5c0fd4", size = 6538966, upload-time = "2026-04-01T14:44:35.252Z" }, - { url = "https://files.pythonhosted.org/packages/c0/a8/27fb307055087f3668f6d0a8ccb636e7431d56ed0750e07a60547b1e083e/pillow-12.2.0-cp313-cp313t-musllinux_1_2_x86_64.whl", hash = "sha256:dac8d77255a37e81a2efcbd1fc05f1c15ee82200e6c240d7e127e25e365c39ea", size = 7238241, upload-time = "2026-04-01T14:44:37.875Z" }, - { url = "https://files.pythonhosted.org/packages/ad/4b/926ab182c07fccae9fcb120043464e1ff1564775ec8864f21a0ebce6ac25/pillow-12.2.0-cp313-cp313t-win32.whl", hash = "sha256:ee3120ae9dff32f121610bb08e4313be87e03efeadfc6c0d18f89127e24d0c24", size = 6379592, upload-time = "2026-04-01T14:44:40.336Z" }, - { url = "https://files.pythonhosted.org/packages/c2/c4/f9e476451a098181b30050cc4c9a3556b64c02cf6497ea421ac047e89e4b/pillow-12.2.0-cp313-cp313t-win_amd64.whl", hash = "sha256:325ca0528c6788d2a6c3d40e3568639398137346c3d6e66bb61db96b96511c98", size = 7085542, upload-time = "2026-04-01T14:44:43.251Z" }, - { url = "https://files.pythonhosted.org/packages/00/a4/285f12aeacbe2d6dc36c407dfbbe9e96d4a80b0fb710a337f6d2ad978c75/pillow-12.2.0-cp313-cp313t-win_arm64.whl", hash = "sha256:2e5a76d03a6c6dcef67edabda7a52494afa4035021a79c8558e14af25313d453", size = 2465765, upload-time = "2026-04-01T14:44:45.996Z" }, - { url = "https://files.pythonhosted.org/packages/4e/b7/2437044fb910f499610356d1352e3423753c98e34f915252aafecc64889f/pillow-12.2.0-pp311-pypy311_pp73-macosx_10_15_x86_64.whl", hash = "sha256:0538bd5e05efec03ae613fd89c4ce0368ecd2ba239cc25b9f9be7ed426b0af1f", size = 5273969, upload-time = "2026-04-01T14:45:55.538Z" }, - { url = "https://files.pythonhosted.org/packages/f6/f4/8316e31de11b780f4ac08ef3654a75555e624a98db1056ecb2122d008d5a/pillow-12.2.0-pp311-pypy311_pp73-macosx_11_0_arm64.whl", hash = "sha256:394167b21da716608eac917c60aa9b969421b5dcbbe02ae7f013e7b85811c69d", size = 4659674, upload-time = "2026-04-01T14:45:58.093Z" }, - { url = "https://files.pythonhosted.org/packages/d4/37/664fca7201f8bb2aa1d20e2c3d5564a62e6ae5111741966c8319ca802361/pillow-12.2.0-pp311-pypy311_pp73-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:5d04bfa02cc2d23b497d1e90a0f927070043f6cbf303e738300532379a4b4e0f", size = 5288479, upload-time = "2026-04-01T14:46:01.141Z" }, - { url = "https://files.pythonhosted.org/packages/49/62/5b0ed78fce87346be7a5cfcfaaad91f6a1f98c26f86bdbafa2066c647ef6/pillow-12.2.0-pp311-pypy311_pp73-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:0c838a5125cee37e68edec915651521191cef1e6aa336b855f495766e77a366e", size = 7032230, upload-time = "2026-04-01T14:46:03.874Z" }, - { url = "https://files.pythonhosted.org/packages/c3/28/ec0fc38107fc32536908034e990c47914c57cd7c5a3ece4d8d8f7ffd7e27/pillow-12.2.0-pp311-pypy311_pp73-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:4a6c9fa44005fa37a91ebfc95d081e8079757d2e904b27103f4f5fa6f0bf78c0", size = 5355404, upload-time = "2026-04-01T14:46:06.33Z" }, - { url = "https://files.pythonhosted.org/packages/5e/8b/51b0eddcfa2180d60e41f06bd6d0a62202b20b59c68f5a132e615b75aecf/pillow-12.2.0-pp311-pypy311_pp73-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:25373b66e0dd5905ed63fa3cae13c82fbddf3079f2c8bf15c6fb6a35586324c1", size = 6002215, upload-time = "2026-04-01T14:46:08.83Z" }, - { url = "https://files.pythonhosted.org/packages/bc/60/5382c03e1970de634027cee8e1b7d39776b778b81812aaf45b694dfe9e28/pillow-12.2.0-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:bfa9c230d2fe991bed5318a5f119bd6780cda2915cca595393649fc118ab895e", size = 7080946, upload-time = "2026-04-01T14:46:11.734Z" }, + { url = "https://files.pythonhosted.org/packages/25/c2/669d88644cddb1485bd9534e63e8cf476c8e51cb3c3a1297677023505c0e/pillow-12.3.0-cp310-cp310-macosx_10_10_x86_64.whl", hash = "sha256:6c0016e7b354317c4e9e525b937ac8596c38d2d232b419529b9cd7a1cd46e39a", size = 5392418, upload-time = "2026-07-01T11:53:27.808Z" }, + { url = "https://files.pythonhosted.org/packages/6b/ba/3762f376a2948e3036488d773a146e0ae6ecc2ca03ac20e2615bd0b2ba02/pillow-12.3.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:bcc33feacfaefce60c12fd500a277533bdc02b10a19f7f6d348763d8140bbba7", size = 4785287, upload-time = "2026-07-01T11:53:29.761Z" }, + { url = "https://files.pythonhosted.org/packages/07/50/b5d688cc9c52d4482f3d5bcab6ce20bc2a74a85d2343841c907444a3be2c/pillow-12.3.0-cp310-cp310-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5594fc43d548a7ed94949d139aa1341b270f1863f11cfd37f5a6c8b778a6b67f", size = 6253754, upload-time = "2026-07-01T11:53:32.298Z" }, + { url = "https://files.pythonhosted.org/packages/4e/89/36f4cd76cf4baf05c50ababb976249153f18c959171c7f6ba09a6f217260/pillow-12.3.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f0606c8bf2cdefea14a43530f7657cbbb7ecf1c4222512492ef4a4434a9501ec", size = 6925605, upload-time = "2026-07-01T11:53:34.487Z" }, + { url = "https://files.pythonhosted.org/packages/eb/c0/4de58cf6633b9e3a6061ef4be6fb91fc3c90b812ece886f531e3c523d777/pillow-12.3.0-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:85f998ea1848bc6757289e739cfbdda3a04adfd58b02fc018ce54d754a5ce468", size = 6327788, upload-time = "2026-07-01T11:53:36.433Z" }, + { url = "https://files.pythonhosted.org/packages/87/3c/14d53682a19550dbbaf3b598f807d5457646c510805a44c7d7891cd1cd1a/pillow-12.3.0-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:25b9b82bb22e6e2b3cd07b39c68b7b862001226cb3dff7130d1cb914121b39ed", size = 7036288, upload-time = "2026-07-01T11:53:38.712Z" }, + { url = "https://files.pythonhosted.org/packages/38/1d/36279e3c77efe034e4cc2b0393ee74ffdb5a62391dacbf9b916154f5f0b8/pillow-12.3.0-cp310-cp310-win32.whl", hash = "sha256:37dc8f7bbb66efe481bb60defacef820c950c24713fb44962ed6aa2a50966de1", size = 6472396, upload-time = "2026-07-01T11:53:40.781Z" }, + { url = "https://files.pythonhosted.org/packages/48/7c/8fa0039574c476d7c6fa57dd7c32a130436877c6ec1e5ce1cc8ec44878c1/pillow-12.3.0-cp310-cp310-win_amd64.whl", hash = "sha256:300557495eb45ebb8aec96c2da9c4be642fbf7cd937278b4013ba894ea8eb0eb", size = 7226887, upload-time = "2026-07-01T11:53:42.764Z" }, + { url = "https://files.pythonhosted.org/packages/fa/17/e324be141d173c1c919428066c3259f21c1b8982e564e01a4a81e96dbdcf/pillow-12.3.0-cp310-cp310-win_arm64.whl", hash = "sha256:514435a37670e3e5e08f3945b68718b6ed329bb84367777e16f9f4dfe1e61a0f", size = 2568039, upload-time = "2026-07-01T11:53:45.372Z" }, + { url = "https://files.pythonhosted.org/packages/fb/c8/0a78b0e02d7ac54bc03e5321c9220da52f0c2ea83b21f7c40e7f3169c502/pillow-12.3.0-cp311-cp311-macosx_10_10_x86_64.whl", hash = "sha256:00808c5e14ef63ac5161091d242999076604ff74b883423a11e5d7bbb38bf756", size = 5392415, upload-time = "2026-07-01T11:53:47.162Z" }, + { url = "https://files.pythonhosted.org/packages/b2/5b/a02d30018abd97ced9f5a6c63d28597694a00d066516b9c1c6de45859fc9/pillow-12.3.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:37d6d0a00072fd2948eb22bce7e1475f34569d90c87c59f7a2ec59541b77f7a6", size = 4785266, upload-time = "2026-07-01T11:53:49.079Z" }, + { url = "https://files.pythonhosted.org/packages/c8/98/766667a4be768150a202836acd9fad19c06824ca86c4286d3cf6b274964e/pillow-12.3.0-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:bcb46e2f9feff8d06323983bd83ed00c201fdcab3d74973e7072a889b3979fcd", size = 6263814, upload-time = "2026-07-01T11:53:51.32Z" }, + { url = "https://files.pythonhosted.org/packages/3b/2d/ede717bc1144f63886c21fd349bb95860b0d1a21149ff16f2bb362b612b6/pillow-12.3.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:23d27a3e0307ec2244cc51e7287b919aa68d097504ebe19df4e76a98a3eea5bd", size = 6934408, upload-time = "2026-07-01T11:53:53.487Z" }, + { url = "https://files.pythonhosted.org/packages/a3/48/9c58b685e69d49c31af6c8eb9012055fab7e665785165c84796e2c73ce72/pillow-12.3.0-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:4f883547d4b7f0495ebe7056b0cc2aea76094e7a4abc8e933540f3271df27d9c", size = 6337160, upload-time = "2026-07-01T11:53:55.457Z" }, + { url = "https://files.pythonhosted.org/packages/ff/fa/dc2a5c0ba6df93f67c31d34b808b7ce440b40cdbf96f0b81cde1d1e6fa93/pillow-12.3.0-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:236ff70b9312fb68943c703aa842ca6a758abfa45ac187a5e7c1452e96ef72b5", size = 7045172, upload-time = "2026-07-01T11:53:57.736Z" }, + { url = "https://files.pythonhosted.org/packages/86/a5/444817a4d4c4c2417df00513086ca196f388d8f9ef40c2e4ccd1ad1af54b/pillow-12.3.0-cp311-cp311-win32.whl", hash = "sha256:10e41f0fbf1eec8cfd234b8fe17a4caac7c9d0db4c204d3c173a8f9f6ef3232b", size = 6472232, upload-time = "2026-07-01T11:53:59.767Z" }, + { url = "https://files.pythonhosted.org/packages/63/c6/4bad1b18d132a50b27e1365e1ab163616f7a5bb56d330f66f9d1d9d4f9d4/pillow-12.3.0-cp311-cp311-win_amd64.whl", hash = "sha256:8e95e1385e4998ae9694eeaa4730ba5457ff61185b3a55e2e7bea0880aef452a", size = 7233653, upload-time = "2026-07-01T11:54:02.066Z" }, + { url = "https://files.pythonhosted.org/packages/fd/16/00f91ab7760dc842f5aad55217e80fc4a7067a0604535249bc8a2d6d9870/pillow-12.3.0-cp311-cp311-win_arm64.whl", hash = "sha256:ebaea975e03d3141d9d3a507df75c9b3ec90fa9d2ffd07567b3a978d9d790b26", size = 2568195, upload-time = "2026-07-01T11:54:04.622Z" }, + { url = "https://files.pythonhosted.org/packages/37/bf/fb3ebff8ddcb76aac5a01389251bbbb9519922a9b520d8247c1ca864a25d/pillow-12.3.0-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:ba09209fbe443b4acccebe845d8a138b89a8f4fbaeedd44953490b5315d5e965", size = 5345969, upload-time = "2026-07-01T11:54:06.397Z" }, + { url = "https://files.pythonhosted.org/packages/d8/66/9a386a92561f402389a4fc70c18838bf6d35eb5eb5c6850b4b2dc64f5048/pillow-12.3.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:ffd0c5368496f41b0944be820fcb7a838aa6e623d250b01acf2643939c3f99d7", size = 4780323, upload-time = "2026-07-01T11:54:09.351Z" }, + { url = "https://files.pythonhosted.org/packages/25/27/ac8f99618ffd3dde21db0f4d4b1d2ab00c0880595bfd17df103f7f39fd0c/pillow-12.3.0-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:d9c7f76c0673154f044e9d78c8655fb4213f6ca31a836df48b40fe5d187717b9", size = 6266838, upload-time = "2026-07-01T11:54:11.71Z" }, + { url = "https://files.pythonhosted.org/packages/84/21/a35af28dcc61f37ed850a2d64c65c701321dfbf25085e469d5559360cbbf/pillow-12.3.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:78cb2c6865a35ab8ff8b75fd122f6033b92a62c82801110e48ddd6c936a45d91", size = 6940830, upload-time = "2026-07-01T11:54:13.732Z" }, + { url = "https://files.pythonhosted.org/packages/eb/51/8b08617af3ad95e33ce6d7dd2c99ed6c8298f7fb131636303956be022e25/pillow-12.3.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:e491916b378fba47242221bb9ead245211b70d504f495d105d17b14a24b4907c", size = 6344383, upload-time = "2026-07-01T11:54:15.756Z" }, + { url = "https://files.pythonhosted.org/packages/1d/72/cf78ac9780bb93c28328f408973845a309d4d145041665f734572ced1b52/pillow-12.3.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:0dd2064cbc55aaec028ef5fbb60fa47bb6c3e7918e07ff17935284b227a9d2df", size = 7052934, upload-time = "2026-07-01T11:54:17.721Z" }, + { url = "https://files.pythonhosted.org/packages/20/20/25e0f4dc178a6bc0696793720055519a0de89e7661dae886992decbd2f81/pillow-12.3.0-cp312-cp312-win32.whl", hash = "sha256:dbce0b29841537a2fa4a214c2bbf14de3587c9680caa9b4e217568472490b28f", size = 6472684, upload-time = "2026-07-01T11:54:19.839Z" }, + { url = "https://files.pythonhosted.org/packages/45/89/da2f7971a317f83d807fdd4065c0af40208e59e692cc43d315a71a0e96d1/pillow-12.3.0-cp312-cp312-win_amd64.whl", hash = "sha256:a2b55dd6b2a4c4b7d87ffa56bdb33fdc5fdb9a462173861a7bc097f17d91cb09", size = 7227137, upload-time = "2026-07-01T11:54:22.025Z" }, + { url = "https://files.pythonhosted.org/packages/de/47/4845a0a6c0dbf1db8456bd9fc791f13c5ced7ced20606d08a0aacfd25b49/pillow-12.3.0-cp312-cp312-win_arm64.whl", hash = "sha256:331b624368d4f1d069149002f25f44bc61c8919ce8ddb3c45bdad8f6e2d89510", size = 2568267, upload-time = "2026-07-01T11:54:24.051Z" }, + { url = "https://files.pythonhosted.org/packages/9d/ac/31fb64e1e7efb5a4b50cd3d92049ba89ac6e4d8d3bb6a74e15048ca3353e/pillow-12.3.0-cp313-cp313-ios_13_0_arm64_iphoneos.whl", hash = "sha256:21900ce7ba264168cd50defae43cd75d25c833ad4ad6e73ffc5596d12e25ac89", size = 4161684, upload-time = "2026-07-01T11:54:25.934Z" }, + { url = "https://files.pythonhosted.org/packages/87/b4/9805e23d2b4d77842b468513841fda254ee42f0289d25088340e4ff46e2d/pillow-12.3.0-cp313-cp313-ios_13_0_arm64_iphonesimulator.whl", hash = "sha256:4e8c2a84d977f50b9daed6eeaf3baef67d00d5d74d932288f02cb94518ee3ace", size = 4255487, upload-time = "2026-07-01T11:54:27.935Z" }, + { url = "https://files.pythonhosted.org/packages/df/39/ecf519435a200c693fe053a6ee4d835b41cf963a4dfc2551c4e637cb2a71/pillow-12.3.0-cp313-cp313-ios_13_0_x86_64_iphonesimulator.whl", hash = "sha256:ae26d61dfa7a47befdc7572b521024e8745f3d809bd95ca9505a7bba9ef849ec", size = 3696433, upload-time = "2026-07-01T11:54:29.813Z" }, + { url = "https://files.pythonhosted.org/packages/42/92/2fc3ffad878ae8dd5469ec1bc8eb83b71f48e13efdf68f02709003982a32/pillow-12.3.0-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:7a743ff716f746fc19a9557f60dab1600d4613255f8a7aeb3cdde4db7eb15a66", size = 5345889, upload-time = "2026-07-01T11:54:31.97Z" }, + { url = "https://files.pythonhosted.org/packages/10/76/8803c13605b763d33d156c4678fc77f8443389c0c51c8aef707bb02015f4/pillow-12.3.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:d69141514cc30b774ceea5e3ed3a6635c8d8a96edf664689b890f4089111fb35", size = 4780109, upload-time = "2026-07-01T11:54:34.026Z" }, + { url = "https://files.pythonhosted.org/packages/1f/01/e18aff37cb0b4aac47ac90f016d347a49aca667ef97f190b06ac2aabc928/pillow-12.3.0-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:f7401aebd7f581d7f83a439d87d474999317ee099218e5ad25d125290990ba65", size = 6263736, upload-time = "2026-07-01T11:54:36.131Z" }, + { url = "https://files.pythonhosted.org/packages/f7/62/de5bdd77d935331f4f802edc11e4d82950f642caad6cb2f949837b8560e2/pillow-12.3.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0847a763afefb695bc912d7c131e7e0632d4edc1d8698f58ddabec8e46b8b6d3", size = 6937129, upload-time = "2026-07-01T11:54:38.216Z" }, + { url = "https://files.pythonhosted.org/packages/70/4d/105627a13300c5e0df1d174230b32fd1273062c96f7745fd552b945d1e1d/pillow-12.3.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:571b9fcb07b97ef3a492028fb3d2dc0993ca23a06138b0315286566d29ef718a", size = 6339562, upload-time = "2026-07-01T11:54:40.354Z" }, + { url = "https://files.pythonhosted.org/packages/6b/1d/f13de01a553988ab895ba1c722e06cf3144d4f57656fd5b81b6d881f1179/pillow-12.3.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:756c768d0c9c2955feb7a56c37ea24aea2e369f8d36a88da270b6a9f19e62b5e", size = 7049439, upload-time = "2026-07-01T11:54:42.489Z" }, + { url = "https://files.pythonhosted.org/packages/c9/f9/066794cca041b969964f779ee5fa66a9498bbf34248ac39c5d7954e4198f/pillow-12.3.0-cp313-cp313-win32.whl", hash = "sha256:a876864214e136f0eb367788dbd7df045f4806801518e2cfe9e13229cfe06d8f", size = 6473287, upload-time = "2026-07-01T11:54:44.9Z" }, + { url = "https://files.pythonhosted.org/packages/a6/9b/7a58e61d62be561da3a356fe2384d4059a6345fc130e23ef1c36a5b81d24/pillow-12.3.0-cp313-cp313-win_amd64.whl", hash = "sha256:1cca606cd25738df4ed873d5ad46bbdb3d83b5cbca291f6b4ff13a4df6b0bbe8", size = 7239691, upload-time = "2026-07-01T11:54:47.141Z" }, + { url = "https://files.pythonhosted.org/packages/aa/b0/c4ed4f0ef8f8fa5ee8351537db6650bb8189f7e118842978dd6589065692/pillow-12.3.0-cp313-cp313-win_arm64.whl", hash = "sha256:b629de27fda84b42cde7edef0d85f13b958b47f6e9bbcbba9b673c562a89bd8b", size = 2568185, upload-time = "2026-07-01T11:54:49.137Z" }, + { url = "https://files.pythonhosted.org/packages/75/18/2e8b40223153ccbc60df07f9e8928dc0c76202aa4e55ae9f53962b6510d6/pillow-12.3.0-pp311-pypy311_pp73-macosx_10_15_x86_64.whl", hash = "sha256:b3c777e849237620b022f7f297dd67705f9f5cf1685f09f02e46f93e92725468", size = 5302510, upload-time = "2026-07-01T11:56:25.736Z" }, + { url = "https://files.pythonhosted.org/packages/46/3e/51fabf59d5ab801ceab709453d3ab6b180083496579549de4c45ced6528a/pillow-12.3.0-pp311-pypy311_pp73-macosx_11_0_arm64.whl", hash = "sha256:b343699e8308bdc51978310e1c959c584e7869cc8c40780058c87da7781a1e94", size = 4736058, upload-time = "2026-07-01T11:56:28.041Z" }, + { url = "https://files.pythonhosted.org/packages/bf/20/22fe9384b7949e25fb1293bcfc84fb82590ff4ea6b37c95b24d26d793d86/pillow-12.3.0-pp311-pypy311_pp73-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:fbd139c8447d25dd750ab79ee274cc5e1fe80fc56340ab10b18a195e1b6eca3e", size = 5237776, upload-time = "2026-07-01T11:56:30.263Z" }, + { url = "https://files.pythonhosted.org/packages/08/14/f6ba68107680ffa74b39985f3f30884e41318fbc4250caa423c79b4788bb/pillow-12.3.0-pp311-pypy311_pp73-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:e7e480451b9fa137494bccd3a7d69adbe8ac65a87d97be61e11f1b1050a5bac3", size = 5860358, upload-time = "2026-07-01T11:56:32.68Z" }, + { url = "https://files.pythonhosted.org/packages/36/54/0169bc772ec491108b62f644f8ecf1fe5d8ae5ebafde2ee2142210166903/pillow-12.3.0-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:04f01d28a6aaff387bf842a13be313df23ba0597a44f1a976c9feb3c6ff4711a", size = 7231786, upload-time = "2026-07-01T11:56:35.046Z" }, ] [[package]] From 8936d07be887f61ae10ba854ca8777b6d2e88d17 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 13 Jul 2026 13:39:38 -0400 Subject: [PATCH 064/123] fix(proxy): track unauthenticated pass-through requests in spend logs (#32410) Pass-through endpoints configured with auth=false reach the cost-tracking callback with no key/user/team/end-user, so _should_track_cost_callback returned False and the spend-log write was skipped, leaving the request out of request/usage logs. Track pass-through call types even when unauthenticated so the SpendLog row is still written. Co-authored-by: Mubashir Osmani Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/hooks/proxy_track_cost_callback.py | 19 +++- .../hooks/test_proxy_track_cost_callback.py | 86 +++++++++++++++++++ 2 files changed, 104 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index b6342f4fa1a..b839426fcda 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -28,11 +28,20 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( ) from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( + CallTypes, StandardLoggingPayload, StandardLoggingPayloadErrorInformation, ) from litellm.utils import get_end_user_id_for_cost_tracking +_PASS_THROUGH_CALL_TYPES: frozenset[str] = frozenset( + { + CallTypes.pass_through.value, + CallTypes.llm_passthrough_route.value, + CallTypes.allm_passthrough_route.value, + } +) + class _ProxyDBLogger(CustomLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -219,11 +228,13 @@ class _ProxyDBLogger(CustomLogger): verbose_proxy_logger.debug( f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" ) + call_type: Optional[str] = kwargs.get("call_type") if _should_track_cost_callback( user_api_key=user_api_key, user_id=user_id, team_id=team_id, end_user_id=end_user_id, + call_type=call_type, ): ## UPDATE DATABASE await _update_database_and_spend_counters( @@ -412,9 +423,15 @@ def _should_track_cost_callback( user_id: Optional[str], team_id: Optional[str], end_user_id: Optional[str], + call_type: Optional[str] = None, ) -> bool: """ Determine if the cost callback should be tracked based on the kwargs + + Pass-through endpoints can be configured with ``auth=false``, which leaves + the request with no key/user/team/end-user to attribute spend to. Those + requests still forward real provider traffic that operators expect to see + in request/usage logs, so they are tracked even when unauthenticated. """ # don't run track cost callback if user opted into disabling spend @@ -423,7 +440,7 @@ def _should_track_cost_callback( if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None: return True - return False + return call_type in _PASS_THROUGH_CALL_TYPES def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 813a0c5e38f..f289148101a 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -14,6 +14,7 @@ from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.proxy_track_cost_callback import ( _ProxyDBLogger, _get_budget_reservation_from_metadata, + _should_track_cost_callback, _update_database_and_spend_counters, ) @@ -1177,3 +1178,88 @@ async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata(): kwargs["litellm_params"]["metadata"]["user_api_key_user_id"] == "mcp-user@example.com" ) + + +@pytest.mark.parametrize( + "call_type, expected", + [ + ("pass_through_endpoint", True), + ("llm_passthrough_route", True), + ("allm_passthrough_route", True), + ("acompletion", False), + ("call_mcp_tool", False), + (None, False), + ], +) +def test_should_track_cost_callback_pass_through_without_owner(call_type, expected): + """Regression for LIT-3782: unauthenticated pass-through requests (auth=false) + carry no key/user/team/end-user, yet must still be tracked so they land in + LiteLLM_SpendLogs. Other call types with no owner stay untracked.""" + assert ( + _should_track_cost_callback( + user_api_key=None, + user_id=None, + team_id=None, + end_user_id=None, + call_type=call_type, + ) + is expected + ) + + +@pytest.mark.parametrize( + "call_type, expect_spend_log", + [ + ("pass_through_endpoint", True), + ("acompletion", False), + (None, False), + ], +) +@pytest.mark.asyncio +async def test_track_cost_callback_logs_unauthenticated_pass_through_request( + call_type, expect_spend_log +): + """Regression for LIT-3782: a pass-through request with auth=false reaches the + cost callback with no key/user/team/end-user. Before the fix the spend-log + write was skipped and the request never appeared in request/usage logs. It + must now be written for pass-through call types while other unauthenticated + calls remain skipped.""" + logger = _ProxyDBLogger() + + kwargs = { + "call_type": call_type, + "model": "unknown", + "litellm_params": {"metadata": {}}, + "standard_logging_object": { + "response_cost": 0.0, + "request_tags": None, + }, + "stream": False, + } + + with ( + patch( + "litellm.proxy.proxy_server.increment_spend_counters", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.proxy_server.update_cache", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + ) as mock_proxy_logging, + ): + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() + mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == ( + 1 if expect_spend_log else 0 + ) From 78e5c4330124622bf8293d3a35498c8c37ff2b24 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Mon, 13 Jul 2026 10:39:41 -0700 Subject: [PATCH 065/123] feat(lasso): send source.type=litellm for Used By attribution (#33090) Co-authored-by: Or Gershoni --- litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py | 8 +++++++- .../proxy/guardrails/guardrail_hooks/test_lasso.py | 3 +++ 2 files changed, 10 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index 9c4cef2f06b..31ce0cc2214 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -772,7 +772,13 @@ class LassoGuardrail(CustomGuardrail): data: Request data (used for conversation_id generation and tools extraction) cache: Cache instance for storing conversation_id (optional for post-call) """ - payload: Dict[str, Any] = {"messages": messages, "messageType": message_type} + payload: Dict[str, Any] = { + "messages": messages, + "messageType": message_type, + # Drives the "Used By" badge on Lasso Application API Keys: every call from this + # integration is attributed as "litellm" on the keys list. + "source": {"type": "litellm"}, + } # Add optional parameters if available if self.user_id: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lasso.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lasso.py index 5a84b6ebecd..16185cadbdf 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lasso.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_lasso.py @@ -693,6 +693,8 @@ class TestLassoGuardrail: assert prompt_payload["messages"] == messages assert prompt_payload["userId"] == "test-user" assert prompt_payload["sessionId"] == "test-conversation" + # Every call is attributed to the "litellm" integration for the "Used By" badge. + assert prompt_payload["source"] == {"type": "litellm"} # Test COMPLETION payload completion_messages = [{"role": "assistant", "content": "Test response"}] @@ -703,6 +705,7 @@ class TestLassoGuardrail: assert completion_payload["messages"] == completion_messages assert completion_payload["userId"] == "test-user" assert completion_payload["sessionId"] == "test-conversation" + assert completion_payload["source"] == {"type": "litellm"} def test_header_preparation(self): """Test header preparation.""" From 45fed6a50a231822aeb616c99d4b6f17ffd48da0 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 14:44:10 -0700 Subject: [PATCH 066/123] feat(mcp): generalize the bridge envelope identity to a key_hash or user_id subject The scripted two-header client mints under a virtual key it presents at the token endpoint (key_hash), but the interactive DCR client authenticates via SSO at the bridged authorize, which yields a user, not a key. Make EnvelopeIdentity a discriminated subject (subject_type key_hash | user_id) with key_hash_identity / user_identity constructors, and dispatch admission on it: a key_hash reloads the key, a user_id reloads the user and admits them as themselves (user-level budget and SCIM enforced via the same centralized gate; no team bound, since a user belongs to many teams or none). The interactive producer that mints a user_id envelope lands in the follow-up commit. --- .../mcp_server/auth/user_api_key_auth_mcp.py | 57 +++++++++- .../mcp_server/discoverable_endpoints.py | 4 +- .../outbound_credentials/envelope.py | 48 ++++++-- .../auth/test_user_api_key_auth_mcp.py | 105 +++++++++++++++++- .../test_bridge_credentials.py | 7 +- .../outbound_credentials/test_envelope.py | 35 ++++-- .../mcp_server/test_discoverable_endpoints.py | 3 +- 7 files changed, 230 insertions(+), 29 deletions(-) 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 e300a22e5db..faec35db41a 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 @@ -18,6 +18,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credenti is_bridge_envelope_shaped, resolve_bridge_envelope, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( + EnvelopeIdentity, +) from litellm.proxy._types import ( UI_TEAM_ID, LiteLLM_TeamTable, @@ -543,7 +546,7 @@ class MCPRequestHandler: header_key = server.alias or server.server_name if header_key is None: raise HTTPException(status_code=500, detail="Server misconfigured: MCP server has no routable name") - admitted = await MCPRequestHandler._reload_admitted_key(result.identity.key_hash) + admitted = await MCPRequestHandler._reload_admitted_principal(result.identity) await MCPRequestHandler._enforce_admitted_live_policy(admitted=admitted, request=request, route=route) injected = {header_key: {"Authorization": result.upstream_authorization.get_secret_value()}} new_headers = {**(mcp_server_auth_headers or {}), **injected} @@ -572,6 +575,58 @@ class MCPRequestHandler: route=route, ) + @staticmethod + async def _reload_admitted_principal(identity: EnvelopeIdentity) -> UserAPIKeyAuth: + """Reload the live litellm record the envelope's subject references. + + Dispatches on the sealed subject type: a ``key_hash`` reloads the virtual key that + minted the envelope (the scripted two-header client that presents a litellm key at the + token endpoint), a ``user_id`` reloads the user that authenticated interactively (the + DCR client, whose SSO login at the bridged authorize yields a user, not a key). Both + return a ``UserAPIKeyAuth`` the caller runs through the centralized policy gate, so + team/project/org/budget/SCIM enforcement is identical to the principal presenting + itself directly.""" + match identity.subject_type: + case "key_hash": + return await MCPRequestHandler._reload_admitted_key(identity.subject) + case "user_id": + return await MCPRequestHandler._reload_admitted_user(identity.subject) + case _: + assert_never(identity.subject_type) + + @staticmethod + async def _reload_admitted_user(user_id: str) -> UserAPIKeyAuth: + """Reload the live user an interactively-minted envelope references and admit them as + themselves. + + The DCR client authenticates via SSO at the bridged authorize, which yields a user + subject rather than a virtual key, so the envelope admits under the user's own + identity: the reloaded ``user_id`` rides on the returned ``UserAPIKeyAuth`` and the + caller's centralized policy gate then enforces the user's live budget and org state, + and a SCIM-deactivated owner fails closed here exactly as the key path enforces it. No + team is bound; a user may belong to many teams or none, so the envelope grants the + user's own access rather than silently selecting one team's scope. A missing user + fails closed with a 401 rather than admitting an unresolved identity.""" + from litellm.proxy.auth.auth_checks import get_user_object + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Server misconfigured: no database connection") + try: + user_object = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + ) + except (ProxyException, HTTPException): + raise HTTPException(status_code=401, detail="Invalid or expired credential") from None + if user_object is None: + raise HTTPException(status_code=401, detail="Invalid or expired credential") + if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False: + raise HTTPException(status_code=401, detail="Invalid or expired credential") + return UserAPIKeyAuth(user_id=user_object.user_id) + @staticmethod async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth: """Reload the live key record an admitted envelope references and re-check live policy. diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 8d1713a5911..3e727ce95bb 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -961,15 +961,15 @@ def _finish_bridge_mint( build_bridge_token_response, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import - EnvelopeIdentity, SealedEnvelope, UpstreamTokenGrant, + key_hash_identity, ) grant = _bridge_grant_from_token_response(token_response) if not isinstance(grant, UpstreamTokenGrant): return _upstream_rejection_to_mint_error(grant) - identity = EnvelopeIdentity(server_id=mcp_server.server_id, key_hash=ready.key_hash) + identity = key_hash_identity(server_id=mcp_server.server_id, key_hash=ready.key_hash) sealed = build_bridge_token_response(identity, grant, ready.keys, now) if not isinstance(sealed, SealedEnvelope): return "too_large" diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py index 517c2ef5c8f..783e64d13e2 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/envelope.py @@ -67,20 +67,42 @@ typed error, never truncated.""" _ENVELOPE_JWT_ALGORITHM = "HS256" -class EnvelopeIdentity(BaseModel): - """The litellm identity the envelope binds the inner grant to. +EnvelopeSubjectType: TypeAlias = Literal["key_hash", "user_id"] +"""Discriminator for what litellm principal the envelope binds the grant to. - ``key_hash`` is the hashed litellm key that authorized the mint, never a raw - credential (and the edge rejects a bare hash presented as a bearer). Admission - reloads the live key record by it, so the key's current team/org/object-permission - restrictions and its revocation state are enforced at use time rather than frozen at - mint time. ``server_id`` binds the envelope to one MCP server so it cannot be replayed - across a server boundary. +``key_hash`` is a hashed virtual key (the scripted two-header client mints under the key it +presents at the token endpoint); ``user_id`` is a litellm user subject (the interactive DCR +client mints under the SSO-authenticated user, which is the only identity that browser login +yields). Admission reloads a key record for the first and a user record for the second, then +runs both through the same live-policy gate, so team/org/budget/revocation enforcement is +identical either way.""" + + +class EnvelopeIdentity(BaseModel): + """The litellm principal the envelope binds the inner grant to. + + ``subject`` is the principal identifier and ``subject_type`` says how to resolve it: a + hashed litellm key (``key_hash``) or a litellm user id (``user_id``), never a raw + credential (and the edge rejects a bare hash or id presented as a bearer). Admission + reloads the live record by it, so the principal's current team/org restrictions and its + revocation state are enforced at use time rather than frozen at mint time. ``server_id`` + binds the envelope to one MCP server so it cannot be replayed across a server boundary. """ model_config = ConfigDict(frozen=True) server_id: str = Field(min_length=1) - key_hash: str = Field(min_length=1) + subject_type: EnvelopeSubjectType + subject: str = Field(min_length=1) + + +def key_hash_identity(server_id: str, key_hash: str) -> EnvelopeIdentity: + """The identity for the scripted client that mints under a presented virtual key.""" + return EnvelopeIdentity(server_id=server_id, subject_type="key_hash", subject=key_hash) + + +def user_identity(server_id: str, user_id: str) -> EnvelopeIdentity: + """The identity for the interactive DCR client that mints under its SSO user subject.""" + return EnvelopeIdentity(server_id=server_id, subject_type="user_id", subject=user_id) class UpstreamTokenGrant(BaseModel): @@ -200,7 +222,8 @@ class _EnvelopeClaims(BaseModel): iat: int exp: int server_id: str = Field(min_length=1) - key_hash: str = Field(min_length=1) + subject_type: EnvelopeSubjectType + subject: str = Field(min_length=1) grant: str = Field(min_length=1) @@ -236,7 +259,8 @@ def mint_envelope( iat=int(now.timestamp()), exp=int(expires_at.timestamp()), server_id=identity.server_id, - key_hash=identity.key_hash, + subject_type=identity.subject_type, + subject=identity.subject, grant=_encrypt_grant_blob(_grant_plaintext(grant), keys.encryption_key), ) token = ENVELOPE_PREFIX + jwt.encode( @@ -281,7 +305,7 @@ def open_envelope( if not isinstance(grant, UpstreamTokenGrant): return grant return OpenedEnvelope( - identity=EnvelopeIdentity(server_id=claims.server_id, key_hash=claims.key_hash), + identity=EnvelopeIdentity(server_id=claims.server_id, subject_type=claims.subject_type, subject=claims.subject), grant=grant, ) 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 c785ac577f7..a6affe5496c 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 @@ -4910,6 +4910,7 @@ class TestMCPDcrBridgeDelegateAdmission: cls, *, key_hash=None, + user_id=None, server_id="bridge-server-id", access_token="inner-upstream-access-token", token_type="Bearer", @@ -4921,17 +4922,23 @@ class TestMCPDcrBridgeDelegateAdmission: envelope_keys_from_master_key, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( - EnvelopeIdentity, SealedEnvelope, UpstreamTokenGrant, + key_hash_identity, mint_envelope, + user_identity, ) from pydantic import SecretStr + identity = ( + user_identity(server_id=server_id, user_id=user_id) + if user_id is not None + else key_hash_identity(server_id=server_id, key_hash=key_hash or cls._KEY_HASH) + ) keys = envelope_keys_from_master_key(master_key or cls._MASTER_KEY) now = minted_at or datetime.now(timezone.utc) sealed = mint_envelope( - identity=EnvelopeIdentity(server_id=server_id, key_hash=key_hash or cls._KEY_HASH), + identity=identity, grant=UpstreamTokenGrant( access_token=SecretStr(access_token), token_type=token_type, @@ -4999,6 +5006,22 @@ class TestMCPDcrBridgeDelegateAdmission: stack.enter_context(patcher) yield get_key_object + @staticmethod + @contextlib.contextmanager + def _patch_user_reload(*, return_value=None, side_effect=None): + """Patch the user-subject reload path an interactively-minted envelope takes: the + ``get_user_object`` lookup ``_reload_admitted_user`` runs (which also drives the SCIM gate), + plus the ``prisma_client`` / ``user_api_key_cache`` globals. The centralized gate's own + fetches fail-safe to None under the MagicMock prisma, so an unblocked user admits. Yields the + ``get_user_object`` mock so a caller can assert the sealed user_id was the reload key.""" + get_user_object = AsyncMock(return_value=return_value, side_effect=side_effect) + with ( + patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + ): + yield get_user_object + @staticmethod def _mcp_request(path="/mcp/bridge_delegate_server"): """A minimal ``Request`` for direct ``_admit_dcr_bridge_delegate`` calls, mirroring how @@ -5060,6 +5083,84 @@ class TestMCPDcrBridgeDelegateAdmission: "bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"} } + async def test_user_subject_envelope_admits_under_the_reloaded_user(self): + """An interactively-minted (user_id) envelope admits under the reloaded USER, not a key: the + reload is keyed by the sealed user_id, the admitted auth carries that user_id, the raw-key + pipeline is never invoked, and the inner upstream token is injected for egress. This is the + interactive-DCR admission the whole flow exists for.""" + envelope = self._mint_bridge_envelope(user_id="sso-user-7") + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/bridge_delegate_server", + "headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))], + } + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ) as mock_auth, + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + self._patch_user_reload( + return_value=MagicMock(user_id="sso-user-7", metadata={"scim_active": True}) + ) as get_user_object, + ): + mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() + (auth_result, _h, _s, mcp_server_auth_headers, _o, _r) = await MCPRequestHandler.process_mcp_request(scope) + + assert get_user_object.await_args.kwargs["user_id"] == "sso-user-7" + assert auth_result.user_id == "sso-user-7" + mock_auth.assert_not_called() + assert mcp_server_auth_headers == { + "bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"} + } + + async def test_user_subject_envelope_missing_user_fails_closed_401(self): + """A user_id envelope whose user has since been deleted must fail closed: get_user_object + resolves None, so admission 401s instead of admitting an unresolved identity.""" + envelope = self._mint_bridge_envelope(user_id="ghost-user") + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/bridge_delegate_server", + "headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))], + } + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + self._patch_user_reload(return_value=None), + ): + mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + + async def test_user_subject_envelope_scim_deactivated_user_fails_closed_401(self): + """SCIM-deactivating the envelope's user revokes it immediately: the reloaded user carries + scim_active False, so admission 401s rather than letting an offboarded user keep tool access + until the envelope expires.""" + envelope = self._mint_bridge_envelope(user_id="offboarded-user") + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/bridge_delegate_server", + "headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))], + } + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + self._patch_user_reload( + return_value=MagicMock(user_id="offboarded-user", metadata={"scim_active": False}) + ), + ): + mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + async def test_revoked_key_envelope_fails_closed_401(self): """An envelope whose key has since been deleted must fail closed: ``get_key_object`` raises for the missing row, so admission 401s instead of admitting the caller as an unrestricted diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_bridge_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_bridge_credentials.py index 82e8e2aae89..ecea86bbed4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_bridge_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_bridge_credentials.py @@ -28,13 +28,14 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import EnvelopeTooLarge, SealedEnvelope, UpstreamTokenGrant, + key_hash_identity, mint_envelope, ) _NOW = datetime(2026, 7, 9, 12, 0, 0, tzinfo=timezone.utc) _MASTER_KEY = "sk-master-key-for-derivation-tests-0123456789" _ACCESS_TOKEN = "upstream-access-token-do-not-leak-8f14e45fceea" -_IDENTITY = EnvelopeIdentity(server_id="srv-456", key_hash="hashed-key-123") +_IDENTITY = key_hash_identity(server_id="srv-456", key_hash="hashed-key-123") _SERVER_ID = _IDENTITY.server_id @@ -138,7 +139,7 @@ def test_resolve_envelope_minted_for_another_server_is_invalid(): captured or misrouted envelope cannot forward one server's upstream credential to another. The valid access token stays sealed; the mismatch alone fails the resolve.""" keys = envelope_keys_from_master_key(_MASTER_KEY) - other_server_identity = EnvelopeIdentity(server_id="srv-OTHER", key_hash=_IDENTITY.key_hash) + other_server_identity = key_hash_identity(server_id="srv-OTHER", key_hash=_IDENTITY.subject) token = _sealed_token(keys, identity=other_server_identity) result = resolve_bridge_envelope(token, keys, _NOW, _SERVER_ID) assert isinstance(result, BridgeEnvelopeInvalid) @@ -155,7 +156,7 @@ def test_resolve_non_ascii_server_id_stays_total_and_does_not_raise(): unicode server_id); it stays total and returns a typed result. A matching non-ASCII id admits, a mismatching one is BridgeEnvelopeInvalid, and neither raises.""" keys = envelope_keys_from_master_key(_MASTER_KEY) - unicode_identity = EnvelopeIdentity(server_id="srv-café", key_hash=_IDENTITY.key_hash) + unicode_identity = key_hash_identity(server_id="srv-café", key_hash=_IDENTITY.subject) token = _sealed_token(keys, identity=unicode_identity) assert isinstance(resolve_bridge_envelope(token, keys, _NOW, "srv-café"), BridgeEnvelopeAdmitted) assert isinstance(resolve_bridge_envelope(token, keys, _NOW, "srv-cafe"), BridgeEnvelopeInvalid) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_envelope.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_envelope.py index b44f3f84cc9..7a2b51c2a95 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_envelope.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_envelope.py @@ -36,8 +36,10 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import SealedEnvelope, UpstreamTokenGrant, is_envelope, + key_hash_identity, mint_envelope, open_envelope, + user_identity, ) from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value, encrypt_value @@ -51,7 +53,7 @@ _WRONG_SIGNING = EnvelopeKeys(signing_key=SecretStr(_OTHER_SIGNING_KEY), encrypt _WRONG_ENCRYPTION = EnvelopeKeys(signing_key=SecretStr(_SIGNING_KEY), encryption_key=SecretStr(_OTHER_ENCRYPTION_KEY)) _ACCESS_TOKEN = "upstream-access-token-do-not-leak-8f14e45fceea" _REFRESH_TOKEN = "upstream-refresh-token-do-not-leak-1d0aa4b7" -_IDENTITY = EnvelopeIdentity(server_id="srv-456", key_hash="hashed-key-123") +_IDENTITY = key_hash_identity(server_id="srv-456", key_hash="hashed-key-123") def _full_grant() -> UpstreamTokenGrant: @@ -137,12 +139,13 @@ def test_minimal_grant_round_trips_without_none_leakage_into_claims(): def test_claim_layout_and_no_plaintext_token_in_envelope(): token = _sealed_token(_full_grant()) claims = _unverified_claims(token) - assert set(claims) == {"iss", "iat", "exp", "server_id", "key_hash", "grant"} + assert set(claims) == {"iss", "iat", "exp", "server_id", "subject_type", "subject", "grant"} assert claims["iss"] == ENVELOPE_ISSUER assert claims["iat"] == int(_NOW.timestamp()) assert claims["exp"] == int(_NOW.timestamp()) + 600 assert claims["server_id"] == "srv-456" - assert claims["key_hash"] == "hashed-key-123" + assert claims["subject_type"] == "key_hash" + assert claims["subject"] == "hashed-key-123" assert _ACCESS_TOKEN not in token assert _ACCESS_TOKEN not in json.dumps(claims) assert _REFRESH_TOKEN not in json.dumps(claims) @@ -226,11 +229,11 @@ def test_wrong_issuer_is_malformed_payload(): def test_missing_identity_claim_is_malformed_payload(): claims = _unverified_claims(_sealed_token(_full_grant())) - forged = _forge({key: value for key, value in claims.items() if key != "key_hash"}) + forged = _forge({key: value for key, value in claims.items() if key != "subject"}) assert isinstance(open_envelope(forged, _KEYS, _NOW), MalformedPayload) -@pytest.mark.parametrize("identity_claim", ["server_id", "key_hash"]) +@pytest.mark.parametrize("identity_claim", ["server_id", "subject"]) def test_signed_empty_identity_claim_is_malformed_payload_not_a_raise(identity_claim): claims = _unverified_claims(_sealed_token(_full_grant())) forged = _forge({**claims, identity_claim: ""}) @@ -463,9 +466,11 @@ def test_non_positive_expires_in_is_rejected_at_construction_without_leaking(): def test_empty_identity_and_key_fields_are_rejected_at_construction(): with pytest.raises(ValidationError): - EnvelopeIdentity(server_id="", key_hash="hashed-key-123") + EnvelopeIdentity(server_id="", subject_type="key_hash", subject="hashed-key-123") with pytest.raises(ValidationError): - EnvelopeIdentity(server_id="srv-456", key_hash="") + EnvelopeIdentity(server_id="srv-456", subject_type="key_hash", subject="") + with pytest.raises(ValidationError): + EnvelopeIdentity(server_id="srv-456", subject_type="not-a-subject-type", subject="x") with pytest.raises(ValidationError): EnvelopeKeys(signing_key=SecretStr(""), encryption_key=SecretStr(_ENCRYPTION_KEY)) with pytest.raises(ValidationError): @@ -474,6 +479,20 @@ def test_empty_identity_and_key_fields_are_rejected_at_construction(): UpstreamTokenGrant(access_token=SecretStr(""), token_type="Bearer") +def test_user_subject_identity_round_trips(): + """The user_id subject variant seals and opens with its discriminator intact, so the edge can + tell an interactively-minted (user) envelope from a scripted (key_hash) one and reload the right + kind of record.""" + identity = user_identity(server_id="srv-456", user_id="user-42") + sealed = mint_envelope(identity, _full_grant(), _KEYS, _NOW) + assert isinstance(sealed, SealedEnvelope) + opened = open_envelope(sealed.token.get_secret_value(), _KEYS, _NOW) + assert isinstance(opened, OpenedEnvelope) + assert opened.identity.server_id == "srv-456" + assert opened.identity.subject_type == "user_id" + assert opened.identity.subject == "user-42" + + def test_public_models_are_frozen(): sealed = mint_envelope(_IDENTITY, _full_grant(), _KEYS, _NOW) assert isinstance(sealed, SealedEnvelope) @@ -484,4 +503,4 @@ def test_public_models_are_frozen(): with pytest.raises(ValidationError): opened.grant = _minimal_grant() with pytest.raises(ValidationError): - _IDENTITY.key_hash = "someone-elses-hash" + _IDENTITY.subject = "someone-elses-hash" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 68466e624ec..e43312e400e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4441,7 +4441,8 @@ async def test_oauth_delegate_bridge_token_exchange_mints_envelope_not_raw_token keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY) opened = resolve_bridge_envelope(token, keys, datetime.now(timezone.utc), server.server_id) assert isinstance(opened, BridgeEnvelopeAdmitted) - assert opened.identity.key_hash == "hashed-litellm-key-77" + assert opened.identity.subject_type == "key_hash" + assert opened.identity.subject == "hashed-litellm-key-77" assert opened.upstream_authorization.get_secret_value() == "Bearer UPSTREAM-SECRET-TOKEN" From 02e9c5631a88d0bdb52d1d4ccb1de21e9c29bede Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Sat, 11 Jul 2026 14:58:52 -0700 Subject: [PATCH 067/123] feat(mcp): interactive SSO sign-in for dcr_bridge oauth_delegate DCR clients Completes the oauth_delegate bridge for real DCR clients (Claude Code, Claude Desktop), which send no litellm key and cannot use the scripted two-header path. On the short-circuit bridge arm the gateway now captures the SSO-authenticated litellm user from the browser session at /authorize and seals it into the OAuth state; at /callback it seals that user plus the upstream code into a gateway authorization code the client echoes back; at /token it recovers the user, exchanges the real upstream code, and mints a user-subject envelope. The user identity captured in the browser thus rides to the back-channel token call with nothing stored server-side, and admission opens the envelope under that user. The scripted key_hash path is unchanged (raw upstream code, key from the request); without a session the browser is sent through login first. --- .../mcp_server/discoverable_endpoints.py | 176 +++++++++++++-- .../mcp_server/test_discoverable_endpoints.py | 200 +++++++++++++++++- 2 files changed, 352 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 3e727ce95bb..bd26abe0b26 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -12,7 +12,7 @@ from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse import httpx from fastapi import APIRouter, Form, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response -from pydantic import BaseModel, SecretStr, ValidationError +from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError from typing_extensions import assert_never from litellm._logging import verbose_logger @@ -41,6 +41,7 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( + EnvelopeIdentity, EnvelopeKeys, UpstreamTokenGrant, ) @@ -98,6 +99,8 @@ def encode_state_with_base_url( code_challenge: Optional[str] = None, code_challenge_method: Optional[str] = None, client_redirect_uri: Optional[str] = None, + litellm_user_id: str | None = None, + mcp_server_id: str | None = None, ) -> str: """ Encode the base_url, original state, and PKCE parameters using encryption. @@ -108,6 +111,11 @@ def encode_state_with_base_url( code_challenge: PKCE code challenge from client code_challenge_method: PKCE code challenge method from client client_redirect_uri: Original redirect_uri from client + litellm_user_id: The SSO-authenticated litellm user captured at the bridge authorize + (interactive dcr_bridge oauth_delegate only); the callback seals it into the gateway + authorization code so the token mint can bind the envelope to this user + mcp_server_id: The bridge server the interactive flow targets, sealed alongside + litellm_user_id so the gateway code cannot be replayed against another server Returns: An encrypted string that encodes all values @@ -118,6 +126,8 @@ def encode_state_with_base_url( "code_challenge": code_challenge, "code_challenge_method": code_challenge_method, "client_redirect_uri": client_redirect_uri, + "litellm_user_id": litellm_user_id, + "mcp_server_id": mcp_server_id, } state_json = json.dumps(state_data, sort_keys=True) encrypted_state = encrypt_value_helper(state_json) @@ -145,6 +155,68 @@ def decode_state_hash(encrypted_state: str) -> dict: return state_data +_BRIDGE_AUTH_CODE_PREFIX = "llm_bcode_" + + +class _BridgeAuthorizationCode(BaseModel): + """The identity and upstream code the gateway seals into the authorization code it hands a DCR + client for an interactive dcr_bridge oauth_delegate sign-in, recovered at the token endpoint.""" + + model_config = ConfigDict(frozen=True) + upstream_code: str = Field(min_length=1) + litellm_user_id: str = Field(min_length=1) + mcp_server_id: str = Field(min_length=1) + + +def is_bridge_authorization_code(code: str) -> bool: + """Cheap prefix check that ``code`` is a gateway-sealed bridge authorization code rather than a + raw upstream code, so the token endpoint can route without decrypting.""" + return code.startswith(_BRIDGE_AUTH_CODE_PREFIX) + + +def seal_bridge_authorization_code(upstream_code: str, litellm_user_id: str, mcp_server_id: str) -> str: + """Seal the upstream authorization code and the SSO-captured litellm user into a gateway + authorization code. The DCR client only echoes this opaque value back at the token endpoint; the + gateway decrypts it there to recover the user (to bind the envelope) and the upstream code (to + exchange with the upstream), so a litellm identity captured in the browser at authorize survives + to the back-channel token call with nothing stored server-side. Encrypted with the repo's + authenticated symmetric helper (the same family the OAuth state uses), so the client can neither + read nor forge it.""" + payload = json.dumps( + {"upstream_code": upstream_code, "litellm_user_id": litellm_user_id, "mcp_server_id": mcp_server_id}, + sort_keys=True, + ) + return _BRIDGE_AUTH_CODE_PREFIX + encrypt_value_helper(payload) + + +def open_bridge_authorization_code(code: str) -> _BridgeAuthorizationCode | None: + """Recover the sealed identity and upstream code, or ``None`` when ``code`` is not a gateway + bridge code or does not decrypt / validate. Total over hostile input: a raw upstream code (the + scripted two-header path) returns ``None`` and the caller falls through to the existing + behavior.""" + if not is_bridge_authorization_code(code): + return None + decrypted = decrypt_value_helper( + code[len(_BRIDGE_AUTH_CODE_PREFIX) :], "bridge_authorization_code", return_original_value=False + ) + if not isinstance(decrypted, str): + return None + try: + return _BridgeAuthorizationCode.model_validate_json(decrypted) + except ValidationError: + return None + + +def _redirect_to_litellm_login(request: Request) -> RedirectResponse: + """Send an unauthenticated browser through litellm login before the interactive bridge authorize + can capture its identity. The bridge oauth_delegate flow seals the SSO user into the gateway code, + so a session is required; without one there is nothing to bind. After login the user re-initiates + the connection, which then finds the session cookie (the seamless return-to round-trip, which is + origin-validated against the control-plane URL, is a follow-up).""" + base_url = get_request_base_url(request) + return RedirectResponse(f"{base_url}/sso/key/generate") + + # LIT-4197: some upstream authorization servers reject an over-long ``state`` # (the encrypted OAuth session blob routinely exceeds their limit). The upstream # only needs an opaque value it echoes back on ``/callback``, so we forward a @@ -697,12 +769,31 @@ async def authorize_with_server( parsed = urlparse(redirect_uri) base_url = urlunparse(parsed._replace(query="")) request_base_url = get_request_base_url(request) + + # Interactive dcr_bridge oauth_delegate sign-in: this arm runs the gateway /callback and /token in + # the loop, so the gateway can capture the litellm user here (from the browser's UI session) and + # carry it to the back-channel token mint. Seal the SSO user and the target server into the state; + # the callback reads them back to mint the gateway authorization code. A DCR client cannot present a + # litellm key, so the browser session is the only identity source; without one there is nothing to + # bind, so send the user through login first. Every other oauth2 server keeps the identity-less state. + litellm_user_id: str | None = None + if mcp_server.is_dcr_bridge and mcp_server.is_oauth_delegate: + from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # inline import avoids a module-load circular import + _user_id_from_session_cookie, + ) + + litellm_user_id = _user_id_from_session_cookie(request) + if litellm_user_id is None: + return _redirect_to_litellm_login(request) + encoded_state = encode_state_with_base_url( base_url=base_url, original_state=state, code_challenge=code_challenge, code_challenge_method=code_challenge_method, client_redirect_uri=redirect_uri, + litellm_user_id=litellm_user_id, + mcp_server_id=mcp_server.server_id if litellm_user_id else None, ) relay_state = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES) @@ -824,11 +915,14 @@ _BridgeMintError = Literal[ @dataclass(frozen=True, slots=True) class _BridgeMintReady: - """Everything the seal needs, resolved once before the exchange: the authorizing key hash and the - master-key-derived envelope keys. Passing this forward means identity resolution and key derivation - happen exactly once, and ``_finish_bridge_mint`` has no preconditions left that could fail.""" + """Everything the seal needs, resolved once before the exchange: the identity to bind the envelope + to and the master-key-derived envelope keys. The identity is a key_hash subject for the scripted + two-header client (resolved from the litellm key it presents) or a user_id subject for the + interactive SSO client (the user recovered from the gateway authorization code), so one phase-3 seal + serves both. Resolving identity here means ``_finish_bridge_mint`` has no preconditions left to + fail.""" - key_hash: str + identity: "EnvelopeIdentity" keys: "EnvelopeKeys" @@ -844,8 +938,8 @@ def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse: status, code, desc = ( 400, "invalid_request", - "this server issues a gateway-bound credential; send a litellm credential " - "(x-litellm-api-key or Authorization) on the token request", + "this server issues a gateway-bound credential; complete the interactive sign-in, or " + "send a litellm credential (x-litellm-api-key or Authorization) on the token request", ) case "unsupported_grant": status, code, desc = ( @@ -923,18 +1017,30 @@ def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _Br assert_never(rejection) -async def _prepare_bridge_mint(request: Request, grant_type: str) -> "_BridgeMintReady | _BridgeMintError": +async def _prepare_bridge_mint( + request: Request, + grant_type: str, + mcp_server: MCPServer, + bridge_identity: _BridgeAuthorizationCode | None = None, +) -> "_BridgeMintReady | _BridgeMintError": """Phase 1, BEFORE the upstream exchange: reject a grant this mint does not support, confirm the gateway can mint (master_key set), resolve the litellm identity, and derive the envelope keys. Returns a ready context or a precise failure value. Running before the exchange is what makes every - failure here fail closed without consuming the single-use code or rotating a refresh token. A bridge - server issues only envelopes and seals no upstream refresh_token, so the client holds none to - present: the refresh_token grant is rejected up front rather than exchanged (which could rotate the - upstream credential) and its result then discarded. Identity-resolution failures keep their origin - so the mapper statuses each truthfully.""" + failure here fail closed without consuming the single-use code. + + Two identity sources, one envelope. The interactive DCR client authenticates via SSO at the bridged + authorize, so its identity arrives as ``bridge_identity`` (the user recovered from the gateway + authorization code) and mints a user subject. The scripted two-header client presents a litellm key + on the token request instead, so its identity is the active key's hash and mints a key_hash subject. + A missing or invalid presented key keeps its resolution origin so the mapper statuses it truthfully; + neither source present is ``no_identity``.""" from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import envelope_keys_from_master_key, ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import + key_hash_identity, + user_identity, + ) from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import master_key, ) @@ -943,16 +1049,21 @@ async def _prepare_bridge_mint(request: Request, grant_type: str) -> "_BridgeMin return "unsupported_grant" if not master_key: return "not_configured" + keys = envelope_keys_from_master_key(master_key) + if bridge_identity is not None: + identity = user_identity(server_id=mcp_server.server_id, user_id=bridge_identity.litellm_user_id) + return _BridgeMintReady(identity=identity, keys=keys) resolved = await _resolve_active_litellm_key(request) if not isinstance(resolved, _ResolvedKey): return _key_resolution_failure_to_mint_error(resolved) - return _BridgeMintReady(key_hash=resolved.key_hash, keys=envelope_keys_from_master_key(master_key)) + identity = key_hash_identity(server_id=mcp_server.server_id, key_hash=resolved.key_hash) + return _BridgeMintReady(identity=identity, keys=keys) def _finish_bridge_mint( ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime ) -> "JSONResponse | _BridgeMintError": - """Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held envelope using + """Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held envelope under the pre-resolved identity and keys, so the client holds one bearer that admits it and forwards the upstream token with nothing stored server-side. The only failures here are properties of the upstream response (no usable token, an already-expired lifetime, or a token too large to seal), @@ -963,14 +1074,12 @@ def _finish_bridge_mint( from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import SealedEnvelope, UpstreamTokenGrant, - key_hash_identity, ) grant = _bridge_grant_from_token_response(token_response) if not isinstance(grant, UpstreamTokenGrant): return _upstream_rejection_to_mint_error(grant) - identity = key_hash_identity(server_id=mcp_server.server_id, key_hash=ready.key_hash) - sealed = build_bridge_token_response(identity, grant, ready.keys, now) + sealed = build_bridge_token_response(ready.identity, grant, ready.keys, now) if not isinstance(sealed, SealedEnvelope): return "too_large" # Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the @@ -1014,6 +1123,7 @@ async def exchange_token_with_server( except TokenEndpointAuthConfigError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc + bridge_identity: _BridgeAuthorizationCode | None = None if grant_type == "refresh_token": if not refresh_token: raise HTTPException( @@ -1033,6 +1143,19 @@ async def exchange_token_with_server( status_code=400, detail="code is required for authorization_code grant", ) + # Interactive dcr_bridge oauth_delegate: the client presents the gateway authorization code the + # callback sealed. Recover the SSO user and the real upstream code from it; the upstream exchange + # below uses the upstream code, and the mint binds the envelope to the recovered user. Bind the + # sealed server to this request so a code minted for one bridge server cannot be spent at another. + # A raw upstream code (scripted path) opens to None and the code is used as-is. + bridge_identity = open_bridge_authorization_code(code) + if bridge_identity is not None: + if bridge_identity.mcp_server_id != mcp_server.server_id: + raise HTTPException( + status_code=400, + detail="Authorization code was issued for a different MCP server", + ) + code = bridge_identity.upstream_code bridge_token_relay = _dcr_bridge_relays_client_registration(mcp_server) if bridge_token_relay and not redirect_uri: raise HTTPException( @@ -1058,7 +1181,7 @@ async def exchange_token_with_server( # phase 3. A failure here returns without ever touching the upstream credential. bridge_mint_ready: _BridgeMintReady | None = None if mcp_server.is_oauth_delegate and mcp_server.is_dcr_bridge: - prepared = await _prepare_bridge_mint(request, grant_type) + prepared = await _prepare_bridge_mint(request, grant_type, mcp_server, bridge_identity) if not isinstance(prepared, _BridgeMintReady): return _bridge_mint_error_response(prepared) bridge_mint_ready = prepared @@ -1706,7 +1829,20 @@ async def callback( # states while permitting same-origin / allowlisted clients. redirect_uri = _get_validated_client_redirect_uri(request, state_data) - params = {"code": code, "state": original_state} + # Interactive dcr_bridge oauth_delegate: the state carries the litellm user the authorize step + # captured. Instead of forwarding the raw upstream code (which the client would present at the + # token endpoint with no way to prove who signed in), seal the user and the upstream code into a + # gateway authorization code and forward THAT. The token endpoint decrypts it to bind the + # envelope to this user. Every other flow forwards the raw code unchanged. + litellm_user_id = state_data.get("litellm_user_id") + mcp_server_id = state_data.get("mcp_server_id") + forwarded_code = code + if isinstance(litellm_user_id, str) and litellm_user_id and isinstance(mcp_server_id, str) and mcp_server_id: + forwarded_code = seal_bridge_authorization_code( + upstream_code=code, litellm_user_id=litellm_user_id, mcp_server_id=mcp_server_id + ) + + params = {"code": forwarded_code, "state": original_state} complete_returned_url = _append_query_params(redirect_uri, params) response = RedirectResponse(url=complete_returned_url, status_code=302) _clear_oauth_state_cookie(response, request, state) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index e43312e400e..966619ee6b8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4365,7 +4365,7 @@ async def test_register_bridge_relay_never_persists(): _BRIDGE_MASTER_KEY = "sk-bridge-producer-master-key-0123456789abcdef" -async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_client_out=None): +async def _exchange_for_bridge_server(server, upstream_body, key_hash, code="auth-code", fake_client_out=None): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( _ResolvedKey, exchange_token_with_server, @@ -4398,13 +4398,17 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash, fake_clie request=_bridge_mock_request(), mcp_server=server, grant_type="authorization_code", - code="auth-code", + code=code, redirect_uri="https://claude.ai/api/mcp/auth_callback", client_id="dcr-client-123", client_secret=None, code_verifier="verifier", ) - if server.is_oauth_delegate and server.is_dcr_bridge: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import is_bridge_authorization_code + + # The key_hash path resolves the presented litellm key; the interactive SSO path recovers identity + # from the gateway authorization code instead, so it never awaits the resolver. + if server.is_oauth_delegate and server.is_dcr_bridge and not is_bridge_authorization_code(code): key_resolver.assert_awaited_once() else: key_resolver.assert_not_awaited() @@ -4446,6 +4450,193 @@ async def test_oauth_delegate_bridge_token_exchange_mints_envelope_not_raw_token assert opened.upstream_authorization.get_secret_value() == "Bearer UPSTREAM-SECRET-TOKEN" +def test_bridge_authorization_code_round_trips_and_rejects_hostile_input(): + """The gateway authorization code seals and recovers the upstream code and the SSO user, and is + total over hostile input: a raw upstream code (scripted path) opens to None, and a tampered or + non-gateway value opens to None rather than raising.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + is_bridge_authorization_code, + open_bridge_authorization_code, + seal_bridge_authorization_code, + ) + + with patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY): + sealed = seal_bridge_authorization_code( + upstream_code="up-code", litellm_user_id="sso-user-9", mcp_server_id="srv-1" + ) + assert is_bridge_authorization_code(sealed) + opened = open_bridge_authorization_code(sealed) + assert opened is not None + assert opened.upstream_code == "up-code" + assert opened.litellm_user_id == "sso-user-9" + assert opened.mcp_server_id == "srv-1" + assert open_bridge_authorization_code("raw-upstream-code") is None + assert open_bridge_authorization_code(sealed[:-4] + "aaaa") is None + + +@pytest.mark.asyncio +async def test_interactive_bridge_token_exchange_mints_user_subject_envelope(): + """An interactive dcr_bridge oauth_delegate exchange (the client presents the gateway code the + callback sealed, and NO litellm key) mints an envelope bound to the SSO-captured user: it opens + to a user_id subject, and the upstream exchange used the real upstream code recovered from the + gateway code, not the sealed wrapper.""" + from datetime import datetime, timezone + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + seal_bridge_authorization_code, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( + BridgeEnvelopeAdmitted, + envelope_keys_from_master_key, + resolve_bridge_envelope, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + with patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY): + gateway_code = seal_bridge_authorization_code( + upstream_code="REAL-UPSTREAM-CODE", litellm_user_id="sso-user-42", mcp_server_id=server.server_id + ) + upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} + captured: dict = {} + response = await _exchange_for_bridge_server( + server, upstream, key_hash=None, code=gateway_code, fake_client_out=captured + ) + + token = json.loads(response.body)["access_token"] + keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY) + opened = resolve_bridge_envelope(token, keys, datetime.now(timezone.utc), server.server_id) + assert isinstance(opened, BridgeEnvelopeAdmitted) + assert opened.identity.subject_type == "user_id" + assert opened.identity.subject == "sso-user-42" + assert opened.upstream_authorization.get_secret_value() == "Bearer UPSTREAM-SECRET-TOKEN" + assert captured["client"].post.call_args.kwargs["data"]["code"] == "REAL-UPSTREAM-CODE" + + +@pytest.mark.asyncio +async def test_interactive_bridge_gateway_code_for_another_server_is_rejected_400(): + """A gateway authorization code is bound to the server it was minted for: presenting it at another + server's token endpoint is a 400, so a code cannot be replayed across a server boundary.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + seal_bridge_authorization_code, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate) + with patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY): + gateway_code = seal_bridge_authorization_code( + upstream_code="up-code", litellm_user_id="sso-user-42", mcp_server_id="a-different-server-id" + ) + upstream = {"access_token": "UPSTREAM-SECRET-TOKEN", "token_type": "Bearer", "expires_in": 3600} + with pytest.raises(HTTPException) as exc: + await _exchange_for_bridge_server(server, upstream, key_hash=None, code=gateway_code) + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_interactive_bridge_authorize_seals_sso_user_into_state(): + """On the short-circuit bridge oauth_delegate arm, authorize captures the SSO user from the UI + session cookie and seals it (and the target server) into the encrypted OAuth state, so the + callback can later mint a user-bound gateway code; it still proceeds to the upstream redirect.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import authorize_with_server + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate, client_id="admin-client", registration_url=None) + captured: dict = {} + + def _capture(**kwargs): + captured.update(kwargs) + return "mocked_encrypted_state" + + with ( + patch( + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + return_value="sso-user-42", + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encode_state_with_base_url", + side_effect=_capture, + ), + ): + response = await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=server, + client_id="ignored", + redirect_uri="http://127.0.0.1:60108/callback", + state="s", + code_challenge="chal", + code_challenge_method="S256", + ) + + assert captured["litellm_user_id"] == "sso-user-42" + assert captured["mcp_server_id"] == server.server_id + assert "/sso/key/generate" not in response.headers["location"] + + +@pytest.mark.asyncio +async def test_interactive_bridge_authorize_without_session_redirects_to_login(): + """Without a UI session there is no identity to bind, so the short-circuit bridge oauth_delegate + authorize sends the browser through litellm login instead of proceeding to the upstream.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import authorize_with_server + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.oauth_delegate, client_id="admin-client", registration_url=None) + with patch( + "litellm.proxy._experimental.mcp_server.byok_oauth_endpoints._user_id_from_session_cookie", + return_value=None, + ): + response = await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=server, + client_id="ignored", + redirect_uri="http://127.0.0.1:60108/callback", + state="s", + code_challenge="chal", + code_challenge_method="S256", + ) + assert "/sso/key/generate" in response.headers["location"] + + +@pytest.mark.asyncio +async def test_interactive_bridge_callback_seals_user_into_gateway_code(): + """When the OAuth state carries the captured SSO user, the callback forwards a gateway + authorization code (sealing the user and upstream code) to the client instead of the raw upstream + code, so the client's later token call can prove who signed in.""" + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + callback, + is_bridge_authorization_code, + ) + + state_data = { + "original_state": "client-state", + "client_redirect_uri": "http://127.0.0.1:60108/cb", + "base_url": "http://127.0.0.1:60108/cb", + "litellm_user_id": "sso-user-42", + "mcp_server_id": "bridge_srv", + } + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_encoded_oauth_state", + return_value="enc", + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash", + return_value=state_data, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._get_validated_client_redirect_uri", + return_value="http://127.0.0.1:60108/cb", + ), + patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), + ): + response = await callback(request=_bridge_mock_request(), code="REAL-UPSTREAM-CODE", state="relay") + + forwarded_code = parse_qs(urlparse(response.headers["location"]).query)["code"][0] + assert is_bridge_authorization_code(forwarded_code) + + @pytest.mark.asyncio async def test_oauth_delegate_bridge_token_exchange_fails_closed_without_litellm_identity(): """Without a resolvable litellm identity on the token request, the exchange must not mint an @@ -4727,10 +4918,11 @@ def test_bridge_reported_expires_in_can_be_zero_at_jwt_exp_boundary(): from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( envelope_keys_from_master_key, ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import key_hash_identity from litellm.types.mcp import MCPAuth ready = _BridgeMintReady( - key_hash="hashed-litellm-key-77", + identity=key_hash_identity(server_id="bridge_srv", key_hash="hashed-litellm-key-77"), keys=envelope_keys_from_master_key(_BRIDGE_MASTER_KEY), ) response = _finish_bridge_mint( From f96899ae2b793f049a01a91db648980cd85021d2 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Mon, 13 Jul 2026 11:08:08 -0700 Subject: [PATCH 068/123] fix(mcp): classify the user-subject reload's errors like the key path (503 outage, 401 missing) _reload_admitted_user mirrored only part of _reload_admitted_key's error contract: it caught ProxyException and HTTPException but had no arm for anything else, so a transient DB outage surfaced as an opaque 500 instead of the retryable 503 the key path guarantees, and a missing user surfaced as a 500 too. The missing-user case is the subtle one: get_user_object raises a bare Exception for a deleted user (not a ProxyException like get_key_object does for a missing key), so the ProxyException/HTTPException clause never caught it and the user_object-is-None branch it was supposed to hit is unreachable on the production path. Add the same except-Exception arm the key path uses, with the one deliberate difference the differing get_user_object contract requires: a database-service-unavailable error still raises the retryable 503, while a missing user or any other non-outage resolution failure fails closed as a 401 rather than propagating as a 500. The regression tests now drive the real behavior (get_user_object raising) rather than a None return that never happens in production, and cover both the 503 outage and the 401 missing-user paths. --- .../mcp_server/auth/user_api_key_auth_mcp.py | 13 +++++-- .../auth/test_user_api_key_auth_mcp.py | 34 +++++++++++++++---- 2 files changed, 39 insertions(+), 8 deletions(-) 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 faec35db41a..01861abbbc0 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 @@ -605,8 +605,14 @@ class MCPRequestHandler: caller's centralized policy gate then enforces the user's live budget and org state, and a SCIM-deactivated owner fails closed here exactly as the key path enforces it. No team is bound; a user may belong to many teams or none, so the envelope grants the - user's own access rather than silently selecting one team's scope. A missing user - fails closed with a 401 rather than admitting an unresolved identity.""" + user's own access rather than silently selecting one team's scope. + + Error handling mirrors the key path's retryable-503 contract, with one deliberate + difference: ``get_key_object`` raises a ``ProxyException`` for a missing key, but + ``get_user_object`` raises a bare ``Exception`` for a missing user (it does not surface as a + ``ProxyException``/``HTTPException``). So a transient DB outage still surfaces as a retryable + 503 via ``_raise_503_if_db_unavailable``, while a missing user, or any other non-outage + resolution failure, fails closed as a 401 rather than propagating as an opaque 500.""" from litellm.proxy.auth.auth_checks import get_user_object from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -621,6 +627,9 @@ class MCPRequestHandler: ) except (ProxyException, HTTPException): raise HTTPException(status_code=401, detail="Invalid or expired credential") from None + except Exception as e: # noqa: BLE001 # DB outage -> retryable 503; a missing user (bare Exception) or any other resolution failure -> fail closed 401, never an opaque 500 + MCPRequestHandler._raise_503_if_db_unavailable(e) + raise HTTPException(status_code=401, detail="Invalid or expired credential") from None if user_object is None: raise HTTPException(status_code=401, detail="Invalid or expired credential") if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False: 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 a6affe5496c..a441e6a154b 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 @@ -5117,8 +5117,10 @@ class TestMCPDcrBridgeDelegateAdmission: } async def test_user_subject_envelope_missing_user_fails_closed_401(self): - """A user_id envelope whose user has since been deleted must fail closed: get_user_object - resolves None, so admission 401s instead of admitting an unresolved identity.""" + """A user_id envelope whose user has since been deleted must fail closed with a 401, not a 500. + get_user_object raises a bare Exception for a missing user (it does not return None on the + production path), so the reload must catch it and fail closed rather than let it propagate as an + opaque 500. Regression for the missing-user path surfacing as a 500.""" envelope = self._mint_bridge_envelope(user_id="ghost-user") scope = { "type": "http", @@ -5129,7 +5131,7 @@ class TestMCPDcrBridgeDelegateAdmission: with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - self._patch_user_reload(return_value=None), + self._patch_user_reload(side_effect=Exception("user not found")), ): mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() with pytest.raises(HTTPException) as exc_info: @@ -5137,6 +5139,28 @@ class TestMCPDcrBridgeDelegateAdmission: assert exc_info.value.status_code == 401 + async def test_user_subject_envelope_db_outage_is_retryable_503(self): + """A transient database outage while reloading the envelope's user is a retryable 503, not an + opaque 500, matching the key path's contract so an interactive DCR client retries instead of + treating a live identity as invalid. Regression for the user reload dropping the 503 arm.""" + envelope = self._mint_bridge_envelope(user_id="sso-user-7") + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/bridge_delegate_server", + "headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))], + } + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + self._patch_user_reload(side_effect=ConnectionError("auth database unreachable")), + ): + mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 503 + async def test_user_subject_envelope_scim_deactivated_user_fails_closed_401(self): """SCIM-deactivating the envelope's user revokes it immediately: the reloaded user carries scim_active False, so admission 401s rather than letting an offboarded user keep tool access @@ -5151,9 +5175,7 @@ class TestMCPDcrBridgeDelegateAdmission: with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), - self._patch_user_reload( - return_value=MagicMock(user_id="offboarded-user", metadata={"scim_active": False}) - ), + self._patch_user_reload(return_value=MagicMock(user_id="offboarded-user", metadata={"scim_active": False})), ): mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() with pytest.raises(HTTPException) as exc_info: From c46863b0e64a4962b84ddf41dc1a9faf7faac3dd Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Mon, 13 Jul 2026 11:18:53 -0700 Subject: [PATCH 069/123] fix(mcp): admit a user-subject envelope with the user's own MCP object permission _reload_admitted_user returned a bare UserAPIKeyAuth(user_id=...), so the shared get_allowed_mcp_servers found no key/team/object-permission grants and an interactive SSO client could admit successfully yet see zero tools on a normal (allow_all_keys=False) server. The key path returns the full key record whose object permission drives that computation; the user path dropped it. Resolve the user's own MCP object permission and put it on the returned auth, so the same get_allowed_mcp_servers the key path uses grants the user their litellm-granted servers and access groups. This reuses get_object_permission (the id-to-grants resolver keys and teams already use) and does not duplicate any permission logic; get_user_object does not load object_permission, so it is resolved from the user's object_permission_id the same way the key and team paths do. Only the user's own object permission is bound. A UserAPIKeyAuth carries a single team_id while a user may belong to many teams, so team-inherited MCP grants for a user are a follow-up: they need a many-teams union get_allowed_mcp_servers does not do off one auth object, and faking one here would be the kind of half-measure that spawns more bugs. Tests cover the user's object permission riding onto the admitted auth, and the existing admit/SCIM/missing-user/503 cases still hold. --- .../mcp_server/auth/user_api_key_auth_mcp.py | 32 ++++++++++--- .../auth/test_user_api_key_auth_mcp.py | 46 ++++++++++++++++++- 2 files changed, 70 insertions(+), 8 deletions(-) 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 01861abbbc0..ac3e4439fa1 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 @@ -601,11 +601,14 @@ class MCPRequestHandler: The DCR client authenticates via SSO at the bridged authorize, which yields a user subject rather than a virtual key, so the envelope admits under the user's own - identity: the reloaded ``user_id`` rides on the returned ``UserAPIKeyAuth`` and the - caller's centralized policy gate then enforces the user's live budget and org state, - and a SCIM-deactivated owner fails closed here exactly as the key path enforces it. No - team is bound; a user may belong to many teams or none, so the envelope grants the - user's own access rather than silently selecting one team's scope. + identity: the reloaded ``user_id`` and the user's own MCP object permission ride on the + returned ``UserAPIKeyAuth``, and the SAME ``get_allowed_mcp_servers`` the key path uses then + computes which servers the user may reach, so the user's litellm MCP grants and access groups + gate the request exactly as a key's do. Only the user's OWN object permission is bound: a + ``UserAPIKeyAuth`` carries a single ``team_id`` while a user may belong to many teams, so + team-inherited MCP grants for a user are a follow-up (they need a many-teams union + ``get_allowed_mcp_servers`` does not do off one auth object). The caller's centralized policy + gate enforces the user's live budget and org state, and a SCIM-deactivated owner fails closed. Error handling mirrors the key path's retryable-503 contract, with one deliberate difference: ``get_key_object`` raises a ``ProxyException`` for a missing key, but @@ -613,7 +616,7 @@ class MCPRequestHandler: ``ProxyException``/``HTTPException``). So a transient DB outage still surfaces as a retryable 503 via ``_raise_503_if_db_unavailable``, while a missing user, or any other non-outage resolution failure, fails closed as a 401 rather than propagating as an opaque 500.""" - from litellm.proxy.auth.auth_checks import get_user_object + from litellm.proxy.auth.auth_checks import get_object_permission, get_user_object from litellm.proxy.proxy_server import prisma_client, user_api_key_cache if prisma_client is None: @@ -634,7 +637,22 @@ class MCPRequestHandler: raise HTTPException(status_code=401, detail="Invalid or expired credential") if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False: raise HTTPException(status_code=401, detail="Invalid or expired credential") - return UserAPIKeyAuth(user_id=user_object.user_id) + # Resolve the user's own MCP object permission (get_user_object does not load it) so the shared + # get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same + # get_object_permission resolver the key and team paths use; no permission logic is duplicated. + object_permission = user_object.object_permission + if user_object.object_permission_id and object_permission is None: + object_permission = await get_object_permission( + object_permission_id=user_object.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + return UserAPIKeyAuth( + user_id=user_object.user_id, + user_role=user_object.user_role, + object_permission=object_permission, + object_permission_id=user_object.object_permission_id, + ) @staticmethod async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth: 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 a441e6a154b..0b1905689ac 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 @@ -5103,7 +5103,13 @@ class TestMCPDcrBridgeDelegateAdmission: patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), self._patch_user_reload( - return_value=MagicMock(user_id="sso-user-7", metadata={"scim_active": True}) + return_value=MagicMock( + user_id="sso-user-7", + metadata={"scim_active": True}, + user_role=None, + object_permission=None, + object_permission_id=None, + ) ) as get_user_object, ): mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() @@ -5116,6 +5122,44 @@ class TestMCPDcrBridgeDelegateAdmission: "bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"} } + async def test_user_subject_envelope_carries_the_users_mcp_object_permission(self): + """The admitted user's own MCP object permission rides on the returned auth so the shared + get_allowed_mcp_servers grants the user their litellm-granted servers, rather than admitting a + bare user with no MCP access. Regression for the signed-in SSO client getting zero tools because + the reload dropped the user's object permission.""" + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="op-user-7", mcp_servers=["bridge_delegate_server"] + ) + envelope = self._mint_bridge_envelope(user_id="sso-user-7") + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/bridge_delegate_server", + "headers": [(b"authorization", f"Bearer {envelope}".encode("latin-1"))], + } + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + new_callable=AsyncMock, + ), + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr, + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + self._patch_user_reload( + return_value=MagicMock( + user_id="sso-user-7", + metadata={"scim_active": True}, + user_role=None, + object_permission=object_permission, + object_permission_id="op-user-7", + ) + ), + ): + mock_mgr.get_mcp_server_by_name.return_value = self._bridge_delegate_server() + (auth_result, _h, _s, _headers, _o, _r) = await MCPRequestHandler.process_mcp_request(scope) + + assert auth_result.object_permission is not None + assert auth_result.object_permission.mcp_servers == ["bridge_delegate_server"] + async def test_user_subject_envelope_missing_user_fails_closed_401(self): """A user_id envelope whose user has since been deleted must fail closed with a 401, not a 500. get_user_object raises a bare Exception for a missing user (it does not return None on the From 7fce761cdeca0ec8431e1b8feb1b8192df92521a Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 13 Jul 2026 12:19:38 -0700 Subject: [PATCH 070/123] fix(ui): respect litellm_key_header_name in BYOK credential save and workflow runs fetches (#33103) --- ui/litellm-dashboard/eslint-suppressions.json | 299 +++++++++--------- .../mcp-servers/_components/mcp_servers.tsx | 1 - .../playground/components/chat_ui/ChatUI.tsx | 1 - .../workflows/WorkflowRuns.test.tsx | 19 +- .../(dashboard)/workflows/WorkflowRuns.tsx | 8 +- .../mcp_tools/ByokCredentialModal.test.tsx | 72 +++++ .../mcp_tools/ByokCredentialModal.tsx | 37 +-- .../src/components/networking.test.ts | 64 ++++ .../src/components/networking.tsx | 3 + 9 files changed, 324 insertions(+), 180 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.test.tsx diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index ba2e2a375ca..953de4fe480 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -515,6 +515,152 @@ "count": 2 } }, + "src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx": { + "unused-imports/no-unused-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx": { + "react-hooks/immutability": { + "count": 2 + } + }, + "src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx": { + "react-hooks/set-state-in-effect": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/MCPToolsetsTab.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 1 + }, + "unused-imports/no-unused-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/OpenAPIQuickPicker.tsx": { + "react-hooks/set-state-in-effect": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx": { + "no-nested-ternary": { + "count": 3 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx": { + "no-nested-ternary": { + "count": 2 + } + }, + "src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 4 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx": { + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/static-components": { + "count": 4 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.tsx": { + "no-nested-ternary": { + "count": 3 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_discovery.tsx": { + "react-hooks/set-state-in-effect": { + "count": 2 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_config.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_display.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/immutability": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 5 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 2 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx": { + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 2 + } + }, "src/app/(dashboard)/memory/_components/MemoryView.tsx": { "react-hooks/set-state-in-effect": { "count": 1 @@ -1860,45 +2006,11 @@ "count": 1 } }, - "src/components/mcp_tools/ByokCredentialModal.tsx": { - "no-restricted-syntax": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx": { - "unused-imports/no-unused-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx": { - "react-hooks/immutability": { - "count": 2 - } - }, - "src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx": { - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, "src/components/mcp_tools/MCPToolArgumentsForm.tsx": { "no-nested-ternary": { "count": 5 } }, - "src/app/(dashboard)/mcp-servers/_components/MCPToolsetsTab.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 - }, - "unused-imports/no-unused-imports": { - "count": 2 - } - }, "src/components/mcp_tools/McpCrudPermissionPanel.tsx": { "no-nested-ternary": { "count": 3 @@ -1907,123 +2019,6 @@ "count": 1 } }, - "src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/OpenAPIQuickPicker.tsx": { - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx": { - "no-nested-ternary": { - "count": 3 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx": { - "no-nested-ternary": { - "count": 2 - } - }, - "src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 4 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx": { - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/static-components": { - "count": 4 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.tsx": { - "no-nested-ternary": { - "count": 3 - }, - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_discovery.tsx": { - "react-hooks/set-state-in-effect": { - "count": 2 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_config.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_display.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/immutability": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 5 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 2 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 2 - } - }, "src/components/model_add/AddCredentialModal.tsx": { "no-restricted-imports": { "count": 1 @@ -2545,4 +2540,4 @@ "count": 1 } } -} +} \ No newline at end of file diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index 0afb4bd9314..f186fef22da 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -684,7 +684,6 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) refetch(); setByokModalServer(null); }} - accessToken={accessToken || ""} /> )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx index 689abf66b41..5133fafb4a7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx @@ -2186,7 +2186,6 @@ const ChatUI: React.FC = ({ loadMCPServers(); setByokModalServer(null); }} - accessToken={accessToken || ""} /> )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx index e73abfe7cd6..1aee3fcc8ab 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.test.tsx @@ -4,7 +4,10 @@ import { afterEach, describe, expect, it, vi } from "vitest"; import WorkflowRuns from "./WorkflowRuns"; -vi.mock("@/components/networking", () => ({ proxyBaseUrl: "" })); +vi.mock("@/components/networking", () => ({ + proxyBaseUrl: "", + getGlobalLitellmHeaderName: () => "x-litellm-api-key", +})); interface FakeRun { run_id: string; @@ -78,4 +81,18 @@ describe("WorkflowRuns (migrated onto shared DataTable)", () => { expect(await screen.findByText("No workflow runs yet")).toBeInTheDocument(); }); + + it("sends the configured litellm key header on every fetch instead of hardcoding Authorization", async () => { + const user = userEvent.setup(); + const fetchSpy = mockFetch(RUNS); + vi.stubGlobal("fetch", fetchSpy); + render(); + + await user.click(await screen.findByText("First run")); + + await waitFor(() => expect(fetchSpy).toHaveBeenCalledTimes(3)); + for (const [url, init] of fetchSpy.mock.calls as [string, RequestInit][]) { + expect(init.headers, url).toEqual({ "x-litellm-api-key": "Bearer tok" }); + } + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx index 9afa07251c2..7354b7479f4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/workflows/WorkflowRuns.tsx @@ -2,7 +2,7 @@ import React, { useState, useEffect, useCallback, useMemo } from "react"; import { Button, Collapse, Drawer, Empty, Spin, Tooltip, Typography } from "antd"; import { ReloadOutlined } from "@ant-design/icons"; import type { ColumnDef, ColumnFiltersState } from "@tanstack/react-table"; -import { proxyBaseUrl } from "@/components/networking"; +import { getGlobalLitellmHeaderName, proxyBaseUrl } from "@/components/networking"; import { DataTable, DataTableFilterDrawer, @@ -507,7 +507,7 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { setLoadingRuns(true); try { const res = await fetch(`${proxyBaseUrl ?? ""}/v1/workflows/runs?limit=100`, { - headers: { Authorization: `Bearer ${accessToken}` }, + headers: { [getGlobalLitellmHeaderName()]: `Bearer ${accessToken}` }, }); if (!res.ok) throw new Error(`HTTP ${res.status}`); const data = await res.json(); @@ -531,10 +531,10 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { const base = proxyBaseUrl ?? ""; const [evRes, msgRes] = await Promise.all([ fetch(`${base}/v1/workflows/runs/${run.run_id}/events`, { - headers: { Authorization: `Bearer ${accessToken}` }, + headers: { [getGlobalLitellmHeaderName()]: `Bearer ${accessToken}` }, }), fetch(`${base}/v1/workflows/runs/${run.run_id}/messages`, { - headers: { Authorization: `Bearer ${accessToken}` }, + headers: { [getGlobalLitellmHeaderName()]: `Bearer ${accessToken}` }, }), ]); const evData = evRes.ok ? await evRes.json() : { events: [] }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.test.tsx new file mode 100644 index 00000000000..021aec5f85f --- /dev/null +++ b/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.test.tsx @@ -0,0 +1,72 @@ +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { registerAuthHeaderNameGetter, registerAuthTokenGetter, registerBaseUrlGetter } from "@/lib/http/runtime"; +import { ByokCredentialModal } from "./ByokCredentialModal"; +import type { MCPServer } from "./types"; + +const fetchSpy = vi.hoisted(() => { + const spy = vi.fn<(request: Request) => Promise>(); + vi.stubGlobal("fetch", spy); + return spy; +}); + +vi.mock("@/components/molecules/message_manager", () => ({ + default: { success: vi.fn(), error: vi.fn() }, +})); + +const SERVER = { server_id: "srv-1", alias: "Linear", server_name: "Linear" } as MCPServer; + +const jsonResponse = (body: unknown, status = 200) => + new Response(JSON.stringify(body), { status, headers: { "Content-Type": "application/json" } }); + +async function fillAndSubmit(user: ReturnType) { + await user.click(screen.getByText("Continue to Authentication")); + await user.type(screen.getByPlaceholderText("Enter your API key"), "linear-key"); + await user.click(screen.getByRole("button", { name: /Connect & Authorize/ })); +} + +beforeEach(() => { + fetchSpy.mockReset(); + registerBaseUrlGetter(() => ""); + registerAuthTokenGetter(() => "sk-session"); +}); + +describe("ByokCredentialModal", () => { + it("saves the credential with the session's configured litellm key header, not a hardcoded Authorization", async () => { + registerAuthHeaderNameGetter(() => "x-litellm-api-key"); + fetchSpy.mockResolvedValue(jsonResponse({ server_id: "srv-1", has_credential: true })); + const onSuccess = vi.fn(); + const user = userEvent.setup(); + render( {}} onSuccess={onSuccess} />); + + await fillAndSubmit(user); + + await waitFor(() => expect(onSuccess).toHaveBeenCalledWith("srv-1")); + const request = fetchSpy.mock.calls[0][0]; + expect(request.method).toBe("POST"); + expect(new URL(request.url).pathname).toBe("/v1/mcp/server/srv-1/user-credential"); + expect(request.headers.get("x-litellm-api-key")).toBe("Bearer sk-session"); + expect(request.headers.get("Authorization")).toBeNull(); + expect(await request.json()).toEqual({ credential: "linear-key", save: true }); + }); + + it("surfaces the backend's detail.error message when the save fails", async () => { + registerAuthHeaderNameGetter(() => "Authorization"); + fetchSpy.mockResolvedValue( + jsonResponse({ detail: { error: "This MCP server does not support BYOK credentials" } }, 400), + ); + const MessageManager = (await import("@/components/molecules/message_manager")).default; + const onSuccess = vi.fn(); + const user = userEvent.setup(); + render( {}} onSuccess={onSuccess} />); + + await fillAndSubmit(user); + + await waitFor(() => + expect(MessageManager.error).toHaveBeenCalledWith("This MCP server does not support BYOK credentials"), + ); + expect(onSuccess).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx b/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx index cb07db871fd..f36de019aa5 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/ByokCredentialModal.tsx @@ -3,6 +3,8 @@ import React, { useState } from "react"; import { Modal, Input, Switch } from "antd"; import MessageManager from "@/components/molecules/message_manager"; +import { fetchClient } from "@/lib/http/api"; +import { ApiError } from "@/lib/http/client"; import { KeyOutlined, LockOutlined, @@ -14,21 +16,22 @@ import { } from "@ant-design/icons"; import { MCPServer } from "./types"; +const byokSaveErrorMessage = (e: unknown): string => { + if (e instanceof ApiError) { + const detail = (e.body as { detail?: { error?: string } } | null)?.detail?.error; + if (detail) return detail; + } + return e instanceof Error && e.message ? e.message : "Failed to connect"; +}; + interface ByokCredentialModalProps { server: MCPServer; open: boolean; onClose: () => void; onSuccess: (serverId: string) => void; - accessToken: string; } -export const ByokCredentialModal: React.FC = ({ - server, - open, - onClose, - onSuccess, - accessToken, -}) => { +export const ByokCredentialModal: React.FC = ({ server, open, onClose, onSuccess }) => { const [step, setStep] = useState<1 | 2>(1); const [apiKey, setApiKey] = useState(""); const [saveKey, setSaveKey] = useState(true); @@ -52,23 +55,15 @@ export const ByokCredentialModal: React.FC = ({ } setLoading(true); try { - const response = await fetch(`/v1/mcp/server/${server.server_id}/user-credential`, { - method: "POST", - headers: { - "Content-Type": "application/json", - Authorization: `Bearer ${accessToken}`, - }, - body: JSON.stringify({ credential: apiKey.trim(), save: saveKey }), + await fetchClient.POST("/v1/mcp/server/{server_id}/user-credential", { + params: { path: { server_id: server.server_id } }, + body: { credential: apiKey.trim(), save: saveKey }, }); - if (!response.ok) { - const err = await response.json(); - throw new Error(err?.detail?.error || "Failed to save credential"); - } MessageManager.success(`Connected to ${serverDisplayName}`); onSuccess(server.server_id); handleClose(); - } catch (e: any) { - MessageManager.error(e.message || "Failed to connect"); + } catch (e) { + MessageManager.error(byokSaveErrorMessage(e)); } finally { setLoading(false); } diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts index b9d04a00f61..e6ee4de2735 100644 --- a/ui/litellm-dashboard/src/components/networking.test.ts +++ b/ui/litellm-dashboard/src/components/networking.test.ts @@ -530,3 +530,67 @@ describe("buildModelGroupTestRequest", () => { expect(body).toEqual({ model: "text-embedding-3-small", input: "test from litellm" }); }); }); + +describe("testMCPToolsListRequest auth headers", () => { + const originalFetch = global.fetch; + + const captureFetch = () => { + const mockFetch = vi.fn().mockResolvedValue({ + ok: true, + status: 200, + headers: { get: () => "application/json" }, + json: vi.fn().mockResolvedValue({ tools: [] }), + } as any); + global.fetch = mockFetch as any; + return mockFetch; + }; + + const sentHeaders = (mockFetch: ReturnType): Record => + (mockFetch.mock.calls[0][1] as RequestInit).headers as Record; + + afterEach(() => { + Networking.setGlobalLitellmHeaderName("Authorization"); + global.fetch = originalFetch; + }); + + it("sends the litellm key under a custom litellm_key_header_name even when an upstream OAuth token uses Authorization", async () => { + Networking.setGlobalLitellmHeaderName("x-litellm-key"); + const mockFetch = captureFetch(); + + await Networking.testMCPToolsListRequest("sk-key", {}, "upstream-oauth-token"); + + const headers = sentHeaders(mockFetch); + expect(headers["x-litellm-key"]).toBe("Bearer sk-key"); + expect(headers["Authorization"]).toBe("Bearer upstream-oauth-token"); + }); + + it("Bearer-prefixes x-litellm-api-key when it is the configured key header (raw values fail _get_bearer_token)", async () => { + Networking.setGlobalLitellmHeaderName("x-litellm-api-key"); + const mockFetch = captureFetch(); + + await Networking.testMCPToolsListRequest("sk-key", {}, "upstream-oauth-token"); + + const headers = sentHeaders(mockFetch); + expect(headers["x-litellm-api-key"]).toBe("Bearer sk-key"); + expect(headers["Authorization"]).toBe("Bearer upstream-oauth-token"); + }); + + it("never clobbers the upstream OAuth token on default deployments", async () => { + const mockFetch = captureFetch(); + + await Networking.testMCPToolsListRequest("sk-key", {}, "upstream-oauth-token"); + + const headers = sentHeaders(mockFetch); + expect(headers["Authorization"]).toBe("Bearer upstream-oauth-token"); + expect(headers["x-litellm-api-key"]).toBe("sk-key"); + }); + + it("sends the litellm key as the bearer on default deployments without an OAuth token", async () => { + const mockFetch = captureFetch(); + + await Networking.testMCPToolsListRequest("sk-key", {}); + + const headers = sentHeaders(mockFetch); + expect(headers["Authorization"]).toBe("Bearer sk-key"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 10d7f12604d..875a345a4c5 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -6650,6 +6650,9 @@ export const testMCPToolsListRequest = async ( }; if (accessToken) { headers["x-litellm-api-key"] = accessToken; + if (globalLitellmHeaderName.toLowerCase() !== "authorization") { + headers[globalLitellmHeaderName] = `Bearer ${accessToken}`; + } } if (oauthAccessToken) { headers["Authorization"] = `Bearer ${oauthAccessToken}`; From aa9dcb43cfaee139ad515ea2a66eab0c31f04a6c Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 13 Jul 2026 12:19:57 -0700 Subject: [PATCH 071/123] refactor(ui): standardize debounce waits behind shared DEBOUNCE_WAIT_MS constant (#33040) --- .../usage/_components/components/UsagePageView.tsx | 3 ++- .../src/app/(dashboard)/users/_components/view_users.tsx | 3 ++- .../PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx | 4 ++-- .../ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx | 4 ++-- .../src/components/VirtualKeysPage/VirtualKeysTable.tsx | 3 ++- .../src/components/common_components/team_dropdown.tsx | 4 ++-- .../src/components/common_components/team_multi_select.tsx | 4 ++-- .../src/components/team/TeamVirtualKeysTable.tsx | 3 ++- ui/litellm-dashboard/src/utils/debounceConstants.ts | 1 + 9 files changed, 17 insertions(+), 12 deletions(-) create mode 100644 ui/litellm-dashboard/src/utils/debounceConstants.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index a95d48aa75b..d2f75609d18 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -8,6 +8,7 @@ import { DownOutlined, ExportOutlined, InfoCircleOutlined, LoadingOutlined, RightOutlined } from "@ant-design/icons"; import { useDebouncedState } from "@tanstack/react-pacer/debouncer"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { Card, Col, @@ -94,7 +95,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { // Debounced search for user selector const [userSearchInput, setUserSearchInput] = useState(""); const [debouncedUserSearch, setDebouncedUserSearch] = useDebouncedState("", { - wait: 300, + wait: DEBOUNCE_WAIT_MS, }); const { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users.tsx b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users.tsx index 18709e05df8..db3b17d6af3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/view_users.tsx @@ -16,6 +16,7 @@ import { import OnboardingModal, { InvitationLink } from "@/components/onboarding_link"; import { updateExistingKeys } from "@/utils/dataUtils"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { isAdminRole, isProxyAdminRole } from "@/utils/roles"; import { useDebouncedState } from "@tanstack/react-pacer/debouncer"; import { useQuery, useQueryClient } from "@tanstack/react-query"; @@ -86,7 +87,7 @@ const ViewUserDashboard: React.FC = ({ const [userToDelete, setUserToDelete] = useState(null); const [activeTab, setActiveTab] = useState("users"); const [filters, setFilters] = useState(initialFilters); - const [debouncedFilters, setDebouncedFilters, debouncer] = useDebouncedState(filters, { wait: 300 }); + const [debouncedFilters, setDebouncedFilters, debouncer] = useDebouncedState(filters, { wait: DEBOUNCE_WAIT_MS }); const [isInvitationLinkModalVisible, setIsInvitationLinkModalVisible] = useState(false); const [invitationLinkData, setInvitationLinkData] = useState(null); const [baseUrl, setBaseUrl] = useState(null); diff --git a/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx b/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx index d42d5ab324a..1d19ba3255d 100644 --- a/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx +++ b/ui/litellm-dashboard/src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx @@ -1,4 +1,5 @@ import { useInfiniteKeyAliases } from "@/app/(dashboard)/hooks/keys/useKeyAliases"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { LoadingOutlined } from "@ant-design/icons"; import { useDebouncedState } from "@tanstack/react-pacer/debouncer"; import { Select } from "antd"; @@ -16,7 +17,6 @@ export interface PaginatedKeyAliasSelectProps { } const SCROLL_THRESHOLD = 0.8; -const DEBOUNCE_MS = 300; export const PaginatedKeyAliasSelect = ({ value, @@ -30,7 +30,7 @@ export const PaginatedKeyAliasSelect = ({ }: PaginatedKeyAliasSelectProps) => { const [searchInput, setSearchInput] = useState(""); const [debouncedSearch, setDebouncedSearch] = useDebouncedState("", { - wait: DEBOUNCE_MS, + wait: DEBOUNCE_WAIT_MS, }); const teamId = allFilters?.["Team ID"] || undefined; diff --git a/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx b/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx index da3ecf77ded..a77b2cf561e 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx @@ -1,4 +1,5 @@ import { useInfiniteModelInfo } from "@/app/(dashboard)/hooks/models/useModels"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { LoadingOutlined } from "@ant-design/icons"; import { useDebouncedState } from "@tanstack/react-pacer/debouncer"; import { Select, Space, Typography } from "antd"; @@ -17,7 +18,6 @@ export interface PaginatedModelSelectProps { } const SCROLL_THRESHOLD = 0.8; -const DEBOUNCE_MS = 300; export const PaginatedModelSelect = ({ value, @@ -30,7 +30,7 @@ export const PaginatedModelSelect = ({ }: PaginatedModelSelectProps) => { const [searchInput, setSearchInput] = useState(""); const [debouncedSearch, setDebouncedSearch] = useDebouncedState("", { - wait: DEBOUNCE_MS, + wait: DEBOUNCE_WAIT_MS, }); const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteModelInfo( diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx index cae6dc54df5..697318af62f 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx @@ -4,6 +4,7 @@ import { useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrgan import { useAllTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { formatNumberWithCommas } from "@/utils/dataUtils"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { ChevronDownIcon, ChevronRightIcon, ChevronUpIcon, SwitchVerticalIcon } from "@heroicons/react/outline"; import { ColumnDef, @@ -64,7 +65,7 @@ export function VirtualKeysTable() { pageSize: 50, }); const [filters, setFilters] = useState(DEFAULT_KEY_FILTERS); - const [debouncedFilters] = useDebouncedValue(filters, { wait: 300 }); + const [debouncedFilters] = useDebouncedValue(filters, { wait: DEBOUNCE_WAIT_MS }); const sortBy = sorting.length > 0 ? sorting[0].id : null; const sortOrder = sorting.length > 0 ? (sorting[0].desc ? "desc" : "asc") : null; diff --git a/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx b/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx index 84140135db7..7d27886c7f5 100644 --- a/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx +++ b/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx @@ -3,6 +3,7 @@ import { Select, Typography } from "antd"; import { LoadingOutlined } from "@ant-design/icons"; import { useDebouncedState } from "@tanstack/react-pacer/debouncer"; import { useInfiniteTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { Team } from "../key_team_helpers/key_list"; const { Text } = Typography; @@ -19,7 +20,6 @@ interface TeamDropdownProps { } const SCROLL_THRESHOLD = 0.8; -const DEBOUNCE_MS = 300; const TeamDropdown: React.FC = ({ value, @@ -31,7 +31,7 @@ const TeamDropdown: React.FC = ({ }) => { const [searchInput, setSearchInput] = useState(""); const [debouncedSearch, setDebouncedSearch] = useDebouncedState("", { - wait: DEBOUNCE_MS, + wait: DEBOUNCE_WAIT_MS, }); const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteTeams( diff --git a/ui/litellm-dashboard/src/components/common_components/team_multi_select.tsx b/ui/litellm-dashboard/src/components/common_components/team_multi_select.tsx index d91f83c589b..3a48b5f7b50 100644 --- a/ui/litellm-dashboard/src/components/common_components/team_multi_select.tsx +++ b/ui/litellm-dashboard/src/components/common_components/team_multi_select.tsx @@ -3,6 +3,7 @@ import { Select, Typography } from "antd"; import { LoadingOutlined } from "@ant-design/icons"; import { useDebouncedState } from "@tanstack/react-pacer/debouncer"; import { useInfiniteTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { Team } from "../key_team_helpers/key_list"; const { Text } = Typography; @@ -17,7 +18,6 @@ interface TeamMultiSelectProps { } const SCROLL_THRESHOLD = 0.8; -const DEBOUNCE_MS = 300; const TeamMultiSelect: React.FC = ({ value = [], @@ -29,7 +29,7 @@ const TeamMultiSelect: React.FC = ({ }) => { const [searchInput, setSearchInput] = useState(""); const [debouncedSearch, setDebouncedSearch] = useDebouncedState("", { - wait: DEBOUNCE_MS, + wait: DEBOUNCE_WAIT_MS, }); const { data, fetchNextPage, hasNextPage, isFetchingNextPage, isLoading } = useInfiniteTeams( diff --git a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx index 207e6f2ccfe..eeccc7482e3 100644 --- a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx @@ -9,6 +9,7 @@ import { DataTableToolbar, } from "@/components/shared/DataTable"; import { Input } from "@/components/ui/input"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { ChevronDownIcon, ChevronRightIcon } from "@heroicons/react/outline"; import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { ColumnDef, ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; @@ -43,7 +44,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi const [columnFilters, setColumnFilters] = useState([]); const [filtersOpen, setFiltersOpen] = useState(false); const [searchInput, setSearchInput] = useState(""); - const [searchQuery] = useDebouncedValue(searchInput, { wait: 300 }); + const [searchQuery] = useDebouncedValue(searchInput, { wait: DEBOUNCE_WAIT_MS }); const handleSearchChange = useCallback((value: string) => { setSearchInput(value); diff --git a/ui/litellm-dashboard/src/utils/debounceConstants.ts b/ui/litellm-dashboard/src/utils/debounceConstants.ts new file mode 100644 index 00000000000..bb8a1e4014c --- /dev/null +++ b/ui/litellm-dashboard/src/utils/debounceConstants.ts @@ -0,0 +1 @@ +export const DEBOUNCE_WAIT_MS = 300; From fa09cde3c09b68354e7f11b4654d15ff77f088cc Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 13 Jul 2026 12:49:01 -0700 Subject: [PATCH 072/123] feat(ui): rebuild the Virtual Keys table on the shared DataTable (#32991) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(ui): rebuild the Virtual Keys table on the shared DataTable Replaces the hand-rolled Tremor table and bespoke toolbar/pagination on the admin Virtual Keys page with the shared DataTable: server-side sort, paginate, and filter, a sticky scrolling body, a search plus column-visibility plus filters toolbar, a right-side filter drawer, and a rows-per-page footer. A page header with the existing key icon carries the Create New Key action. Adds reusable, shadcn-default building blocks for the tables migrating onto the DataTable next: shared IdentityCell, ModelsCell, and SpendBudgetCell in shared/table_cells, plus a shared PageHeader. The models cell reveals overflow in a hover tooltip and the spend/budget cell uses the Meter primitive. All data and domain logic is preserved, including the useKeys query, team and org alias resolution, the user popover, and the KeyInfoView detail swap. The rich async Team/Org/Alias filters move into the drawer, and the toolbar search maps to the key-alias substring search. Status now also reflects key expiry alongside blocked and SCIM-blocked. The VirtualKeysTable tests are updated to the new markup and extended with focused coverage for each new shared cell * fix(ui): address Virtual Keys redesign review feedback Fold the status badge into the clickable Key cell and drop the separate Status column so a key's alias, secret, and status read as one unit. The Key cell is now the single click target that opens the key detail; the whole-row click is removed Migrate the filter drawer off AntD to shadcn. A new Combobox composed from Popover and Input backs the Team, Organization, and Key Alias filters, keeping search and the alias infinite-scroll Show $0.00 for zero spend instead of a hyphen, and extend the shared DataTable with badge, chips, and meter skeleton shapes so the loading state matches the loaded cells (status pill, model chips, spend meter) rather than uniform bars Fix key sorting: the Key column sent its column id "key" as sort_by, which /key/list rejects with 400. It now sorts by the backend field key_alias * fix(ui): use the shadcn base combobox and refine the keys filters and skeletons Replace the hand-rolled filter combobox with the supported shadcn Base UI combobox (ui/combobox, added via the CLI and reused through a small SearchSelect wrapper). Its vended input-group and textarea deps are written for React 19 (plain functions with ref-as-prop); this app is on React 18, where those subcomponents drop the refs Base UI passes for focus and anchoring, so InputGroupInput, InputGroupButton, and ComboboxTrigger are adapted to forwardRef. Those ui/ files now diverge from the registry, and a future shadcn add would overwrite the adaptation until the app moves to React 19. Adds class-variance-authority, which input-group needs Give loading skeletons a per-column renderSkeleton escape hatch on the shared DataTable and mirror the Key cell exactly (alias line, secret, status pill), so skeleton rows match the real rows instead of being shorter and simpler Resolve the automated review: the toolbar search and the drawer Key Alias filter both mapped to the key-alias query, so the search silently overrode the drawer value while its chip stayed visible. Consolidate to a single alias search in the toolbar (placeholder now "Search by key alias…") and drop the redundant drawer field. Re-add coverage for the Created By column's alias-over-email display Refine the Team and Organization filters: they match on name and id, so the labels read "Team" and "Organization" rather than "... ID", each option shows the name with the id on a muted second line instead of "name (id)", and the active-filter chip shows the friendly name * chore(ui): drop duplicate class-variance-authority, use the repo cva package in input-group --- ui/litellm-dashboard/eslint-suppressions.json | 8 - .../VirtualKeysPage/VirtualKeysTable.test.tsx | 205 ++-- .../VirtualKeysPage/VirtualKeysTable.tsx | 970 ++++-------------- .../VirtualKeysPage/keyTableColumns.tsx | 353 +++++++ .../shared/DataTable/DataTable.test.tsx | 32 + .../components/shared/DataTable/DataTable.tsx | 26 +- .../components/shared/DataTable/columnMeta.ts | 3 + .../src/components/shared/DataTable/types.ts | 2 +- .../src/components/shared/PageHeader.test.tsx | 31 + .../src/components/shared/PageHeader.tsx | 29 + .../components/shared/SearchSelect.test.tsx | 64 ++ .../src/components/shared/SearchSelect.tsx | 76 ++ .../shared/table_cells/identity_cell.test.tsx | 38 + .../shared/table_cells/identity_cell.tsx | 48 + .../components/shared/table_cells/index.ts | 3 + .../shared/table_cells/models_cell.test.tsx | 45 + .../shared/table_cells/models_cell.tsx | 56 + .../table_cells/spend_budget_cell.test.tsx | 53 + .../shared/table_cells/spend_budget_cell.tsx | 44 + .../src/components/ui/combobox.tsx | 266 +++++ .../src/components/ui/input-group.tsx | 140 +++ .../src/components/ui/textarea.tsx | 18 + .../src/components/user_dashboard.tsx | 27 +- 23 files changed, 1647 insertions(+), 890 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/PageHeader.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/SearchSelect.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/SearchSelect.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/table_cells/models_cell.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/table_cells/models_cell.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.tsx create mode 100644 ui/litellm-dashboard/src/components/ui/combobox.tsx create mode 100644 ui/litellm-dashboard/src/components/ui/input-group.tsx create mode 100644 ui/litellm-dashboard/src/components/ui/textarea.tsx diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 953de4fe480..7c6df4e7bb9 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1690,14 +1690,6 @@ "count": 1 } }, - "src/components/VirtualKeysPage/VirtualKeysTable.tsx": { - "no-nested-ternary": { - "count": 2 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/activity_metrics.tsx": { "no-nested-ternary": { "count": 1 diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx index 02f5d588149..cf4beeed40f 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx @@ -63,7 +63,7 @@ const mockKey: KeyResponse = { key_alias: "Test Key Alias", spend: 5.5, max_budget: 100, - expires: "2024-12-31T23:59:59Z", + expires: "2999-12-31T23:59:59Z", models: ["gpt-3.5-turbo", "gpt-4"], aliases: {}, config: {}, @@ -154,6 +154,8 @@ const keysResult = (keys: KeyResponse[], data: Partial = {}, extra ...extra, }) as any; +const openFilters = () => fireEvent.click(screen.getByRole("button", { name: "Filters" })); + beforeEach(() => { vi.clearAllMocks(); @@ -170,6 +172,12 @@ it("should render VirtualKeysTable component", () => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); +it("renders the page header with the create-key action slot", () => { + renderWithProviders(Create New Key} />); + expect(screen.getByRole("heading", { name: "Virtual Keys" })).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Create New Key" })).toBeInTheDocument(); +}); + it("should display key information correctly", async () => { renderWithProviders(); @@ -177,6 +185,7 @@ it("should display key information correctly", async () => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); expect(screen.getByText("Test Team")).toBeInTheDocument(); expect(screen.getByText("$5.5000")).toBeInTheDocument(); + expect(screen.getByText("of $100")).toBeInTheDocument(); }); }); @@ -188,14 +197,49 @@ it("should display user email correctly", async () => { }); }); -it("should show loading message only on initial load (isPending)", () => { +it("shows the user alias over the email in the visible cell when both exist", async () => { + mockUseKeys.mockReturnValue( + keysResult([{ ...mockKey, user: { user_id: "user-1", user_email: "user@example.com", user_alias: "The User" } }]), + ); + + renderWithProviders(); + + const row = (await screen.findByText("Test Key Alias")).closest("tr") as HTMLElement; + expect(within(row).getByText("The User")).toBeInTheDocument(); + expect(within(row).queryByText("user@example.com")).not.toBeInTheDocument(); +}); + +it("shows created_by_user alias over email in the Created By column when it is enabled", async () => { + mockUseKeys.mockReturnValue( + keysResult([ + { + ...mockKey, + created_by: "some-uuid", + created_by_user: { user_id: "some-uuid", user_email: "creator@example.com", user_alias: "The Creator" }, + }, + ]), + ); + const user = userEvent.setup(); + renderWithProviders(); + + // Created By is hidden by default; turn it on via the Columns menu. + await user.click(screen.getByRole("button", { name: "Columns" })); + await user.click(await screen.findByText("Created By")); + await user.keyboard("{Escape}"); + + const row = (await screen.findByText("Test Key Alias")).closest("tr") as HTMLElement; + expect(within(row).getByText("The Creator")).toBeInTheDocument(); + expect(within(row).queryByText("creator@example.com")).not.toBeInTheDocument(); +}); + +it("should show a loading state on the initial load and hide the data", () => { mockUseKeys.mockReturnValue(keysResult([], {}, { data: null, isPending: true, isFetching: true })); renderWithProviders(); - expect(screen.getByText("🚅 Loading keys...")).toBeInTheDocument(); + expect(screen.getByText("Loading keys...")).toBeInTheDocument(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); expect(screen.queryByText("Test Key Alias")).not.toBeInTheDocument(); - expect(screen.queryByText("Test Team")).not.toBeInTheDocument(); }); it("should show 'No keys found' message when the key list is empty", () => { @@ -206,61 +250,52 @@ it("should show 'No keys found' message when the key list is empty", () => { expect(screen.getByText("No keys found")).toBeInTheDocument(); }); -it("should handle models with more than 3 entries to trigger expansion UI", () => { +it("collapses models beyond the visible limit into a '+N more' badge", () => { mockUseKeys.mockReturnValue( keysResult([{ ...mockKey, models: ["gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "claude-3", "claude-3-5-sonnet"] }]), ); renderWithProviders(); - expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); + expect(screen.getByText("+2 more")).toBeInTheDocument(); }); -it("should render table headers correctly", () => { +it("should render the redesigned table headers", () => { renderWithProviders(); - expect(screen.getByText("Key ID")).toBeInTheDocument(); - expect(screen.getByText("Key Alias")).toBeInTheDocument(); + expect(screen.getByText("Key")).toBeInTheDocument(); expect(screen.getByText("Team")).toBeInTheDocument(); expect(screen.getByText("Models")).toBeInTheDocument(); - expect(screen.getByText("Spend (USD)")).toBeInTheDocument(); + expect(screen.getByText("Spend / Budget")).toBeInTheDocument(); }); -it("should handle column resizing hover events", () => { +it("sorts by the backend key_alias field (not the column label) when the Key header is clicked", async () => { renderWithProviders(); - const headerCell = document.querySelector("[data-header-id]") as HTMLElement; - expect(headerCell).toBeInTheDocument(); + const keyHeader = screen.getByText("Key").closest("button") as HTMLElement; + fireEvent.click(keyHeader); - const resizer = headerCell?.querySelector(".resizer") as HTMLElement; - expect(resizer).toBeInTheDocument(); - expect(resizer.style.opacity).toBe("0"); - - fireEvent.mouseEnter(headerCell); - expect(resizer.style.opacity).toBe("0.5"); - - fireEvent.mouseLeave(headerCell); - expect(resizer.style.opacity).toBe("0"); + await waitFor(() => { + expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ sortBy: "key_alias" })); + }); }); -it("should open KeyInfoView when clicking on a key ID button", async () => { +it("should open KeyInfoView when clicking the key cell", async () => { renderWithProviders(); await waitFor(() => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); }); - expect(screen.getByText(/Showing.*results/)).toBeInTheDocument(); + expect(screen.getByTestId("pagination-range")).toBeInTheDocument(); - const keyIdButton = screen.getByText("sk-1234567890abcdef"); - fireEvent.click(keyIdButton); + fireEvent.click(screen.getByText("Test Key Alias")); await waitFor(() => { expect(screen.getByText("Back to Keys")).toBeInTheDocument(); - expect(screen.getByText("Created At")).toBeInTheDocument(); }); - expect(screen.queryByText(/Showing.*results/)).not.toBeInTheDocument(); + expect(screen.queryByTestId("pagination-range")).not.toBeInTheDocument(); }); it("should display 'Default Proxy Admin' for user_id when value is 'default_user_id'", async () => { @@ -282,44 +317,6 @@ it("should display 'Default Proxy Admin' for user_id when value is 'default_user }); }); -it("should display created_by_user email in 'Created By' column when available", async () => { - mockUseKeys.mockReturnValue( - keysResult([ - { - ...mockKey, - created_by: "some-uuid-1234", - created_by_user: { user_id: "some-uuid-1234", user_email: "creator@example.com", user_alias: null }, - }, - ]), - ); - - renderWithProviders(); - - await waitFor(() => { - expect(screen.getByText("creator@example.com")).toBeInTheDocument(); - }); -}); - -it("should display created_by_user alias over email when both are available", async () => { - mockUseKeys.mockReturnValue( - keysResult([ - { - ...mockKey, - created_by: "some-uuid-1234", - created_by_user: { user_id: "some-uuid-1234", user_email: "creator@example.com", user_alias: "The Creator" }, - }, - ]), - ); - - renderWithProviders(); - - // Scope to the key's row so we assert the visible cell value: the hover popover that - // also holds the email is portaled out of the row, not the displayed "Created By" text. - const row = (await screen.findByText("Test Key Alias")).closest("tr") as HTMLElement; - expect(within(row).getByText("The Creator")).toBeInTheDocument(); - expect(within(row).queryByText("creator@example.com")).not.toBeInTheDocument(); -}); - it("should render table without crashing when models is null", async () => { mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, models: null as unknown as string[] }])); @@ -327,6 +324,7 @@ it("should render table without crashing when models is null", async () => { await waitFor(() => { expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); + expect(screen.getByText("All Proxy Models")).toBeInTheDocument(); }); }); @@ -341,13 +339,14 @@ it("should display 'Unknown' for last_active when value is null", async () => { }); describe("server-side filtering – the LIT-4080 regression guard", () => { - it("threads an active User ID filter into the useKeys query so any refetch keeps it", async () => { + it("threads an applied User ID filter into the useKeys query so any refetch keeps it", async () => { renderWithProviders(); - fireEvent.click(screen.getByRole("button", { name: "Filters" })); + openFilters(); - const userIdInput = await screen.findByPlaceholderText("Enter User ID..."); + const userIdInput = await screen.findByPlaceholderText(/Enter User ID/); fireEvent.change(userIdInput, { target: { value: "user-42" } }); + fireEvent.click(screen.getByTestId("filter-drawer-apply")); await waitFor(() => { expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: "user-42" })); @@ -361,18 +360,19 @@ describe("server-side filtering – the LIT-4080 regression guard", () => { expect(lastCall[2] ?? {}).toMatchObject({ userID: undefined, teamID: undefined, keyHash: undefined }); }); - it("drops the filter from the useKeys query when Reset Filters is clicked", async () => { + it("drops the filter from the useKeys query when it is cleared", async () => { renderWithProviders(); - fireEvent.click(screen.getByRole("button", { name: "Filters" })); - const userIdInput = await screen.findByPlaceholderText("Enter User ID..."); + openFilters(); + const userIdInput = await screen.findByPlaceholderText(/Enter User ID/); fireEvent.change(userIdInput, { target: { value: "user-42" } }); + fireEvent.click(screen.getByTestId("filter-drawer-apply")); await waitFor(() => { expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: "user-42" })); }); - fireEvent.click(screen.getByRole("button", { name: "Reset Filters" })); + fireEvent.click(screen.getByTestId("datatable-clear-filters")); await waitFor(() => { const lastCall = mockUseKeys.mock.calls[mockUseKeys.mock.calls.length - 1]; @@ -388,8 +388,8 @@ describe("pagination display – total count comes from useKeys", () => { renderWithProviders(); await waitFor(() => { - expect(screen.getByText("Showing 1 - 50 of 509 results")).toBeInTheDocument(); - expect(screen.getByText("Page 1 of 11")).toBeInTheDocument(); + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-50 of 509"); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 1 of 11"); }); }); @@ -399,57 +399,44 @@ describe("pagination display – total count comes from useKeys", () => { renderWithProviders(); await waitFor(() => { - expect(screen.getByText("Showing 1 - 1 of 1 results")).toBeInTheDocument(); - expect(screen.getByText("Page 1 of 1")).toBeInTheDocument(); + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-1 of 1"); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 1 of 1"); }); }); }); -describe("refetch button", () => { - it("should show Fetch button in normal state", () => { +describe("refresh button", () => { + it("renders an enabled refresh control in the normal state", () => { renderWithProviders(); - const fetchButton = screen.getByTitle("Fetch data"); - expect(fetchButton).toBeInTheDocument(); - expect(fetchButton).not.toBeDisabled(); - expect(screen.getByText("Fetch")).toBeInTheDocument(); + const refresh = screen.getByTestId("datatable-refresh"); + expect(refresh).toBeInTheDocument(); + expect(refresh).not.toBeDisabled(); }); - it("should show Fetching state and keep table data visible during refetch", () => { + it("disables the refresh control while a fetch is in flight but keeps data visible", () => { mockUseKeys.mockReturnValue(keysResult([mockKey], {}, { isFetching: true })); renderWithProviders(); - expect(screen.getByText("Fetching")).toBeInTheDocument(); - expect(screen.getByTitle("Fetch data")).toBeDisabled(); + expect(screen.getByTestId("datatable-refresh")).toBeDisabled(); expect(screen.getByText("Test Key Alias")).toBeInTheDocument(); - expect(screen.queryByText("🚅 Loading keys...")).not.toBeInTheDocument(); }); - it("should call refetch when Fetch button is clicked", () => { + it("calls refetch when the refresh control is clicked", () => { const mockRefetch = vi.fn(); mockUseKeys.mockReturnValue(keysResult([mockKey], {}, { refetch: mockRefetch })); renderWithProviders(); - fireEvent.click(screen.getByTitle("Fetch data")); + fireEvent.click(screen.getByTestId("datatable-refresh")); expect(mockRefetch).toHaveBeenCalledTimes(1); }); - - it("should show Fetch button enabled on error so user can retry", () => { - mockUseKeys.mockReturnValue(keysResult([], {}, { data: null, isError: true })); - - renderWithProviders(); - - const fetchButton = screen.getByTitle("Fetch data"); - expect(fetchButton).not.toBeDisabled(); - expect(screen.getByText("Fetch")).toBeInTheDocument(); - }); }); -describe("Status column reflects key.blocked / scim_blocked metadata", () => { - it("should render Active for a non-blocked key", async () => { +describe("Status column reflects blocked / expiry / scim metadata", () => { + it("renders Active for a non-blocked, unexpired key", async () => { mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, blocked: false, metadata: {} }])); renderWithProviders(); @@ -459,7 +446,19 @@ describe("Status column reflects key.blocked / scim_blocked metadata", () => { }); }); - it("should render Blocked when key.blocked is true", async () => { + it("renders Expired when the expiry date has passed", async () => { + mockUseKeys.mockReturnValue( + keysResult([{ ...mockKey, blocked: false, metadata: {}, expires: "2020-01-01T00:00:00Z" }]), + ); + + renderWithProviders(); + + await waitFor(() => { + expect(screen.getByTestId(`key-status-${mockKey.token_id}`)).toHaveTextContent("Expired"); + }); + }); + + it("renders Blocked when key.blocked is true", async () => { mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, blocked: true, metadata: {} }])); renderWithProviders(); @@ -470,7 +469,7 @@ describe("Status column reflects key.blocked / scim_blocked metadata", () => { expect(screen.queryByText(/Blocked by SCIM/i)).not.toBeInTheDocument(); }); - it("should mark a SCIM-blocked key with the SCIM tooltip reason", async () => { + it("marks a SCIM-blocked key with the SCIM tooltip reason", async () => { mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, blocked: true, metadata: { scim_blocked: true } }])); renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx index 697318af62f..112b6cbcba4 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx @@ -1,811 +1,251 @@ "use client"; -import { useKeys, KeyListCallOptions } from "@/app/(dashboard)/hooks/keys/useKeys"; + +import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; import { useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; import { useAllTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; -import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; -import { formatNumberWithCommas } from "@/utils/dataUtils"; import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; -import { ChevronDownIcon, ChevronRightIcon, ChevronUpIcon, SwitchVerticalIcon } from "@heroicons/react/outline"; import { - ColumnDef, - flexRender, - getCoreRowModel, - PaginationState, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { Badge, Icon, Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow, Text } from "@tremor/react"; -import { InfoCircleOutlined, SyncOutlined } from "@ant-design/icons"; -import { Button as AntButton, Popover, Skeleton, Typography } from "antd"; -import { DateCell, IdCell, MoneyCell, StatusBadge } from "@/components/shared/table_cells"; -import React, { useDeferredValue, useMemo, useState } from "react"; -import { getModelDisplayName } from "../key_team_helpers/fetch_available_models_team_key"; -import { PaginatedKeyAliasSelect } from "../KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect"; + DataTable, + DataTableFilterDrawer, + DataTableFilterField, + DataTableToolbar, +} from "@/components/shared/DataTable"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { PageHeader } from "@/components/shared/PageHeader"; +import { Input } from "@/components/ui/input"; +import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; +import { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; +import { KeyRound } from "lucide-react"; +import React, { useCallback, useMemo, useState } from "react"; + import { KeyResponse, Team } from "../key_team_helpers/key_list"; -import FilterComponent, { FilterOption } from "../molecules/filter"; -import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; import KeyInfoView from "../templates/key_info_view"; +import { getKeyTableColumns, KEY_TABLE_HIDDEN_COLUMNS } from "./keyTableColumns"; -type KeyFilterState = { - "Team ID": string; - "Organization ID": string; - "Key Alias": string; - "User ID": string; - "Key Hash": string; +interface VirtualKeysTableProps { + headerActions?: React.ReactNode; +} + +const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; + +const toSortOrder = (sorting: SortingState): "asc" | "desc" | undefined => { + const active = sorting[0]; + if (!active) return undefined; + return active.desc ? "desc" : "asc"; }; -const DEFAULT_KEY_FILTERS: KeyFilterState = { - "Team ID": "", - "Organization ID": "", - "Key Alias": "", - "User ID": "", - "Key Hash": "", +const FILTER_LABELS: Record = { + team_id: "Team", + org_id: "Organization", + user_id: "User ID", + key_hash: "Key ID", }; -type KeyListFilterOptions = Pick< - KeyListCallOptions, - "teamID" | "organizationID" | "selectedKeyAlias" | "userID" | "keyHash" ->; +export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) { + const { data: fetchedOrganizations } = useOrganizations(); + const organizations = useMemo(() => fetchedOrganizations ?? [], [fetchedOrganizations]); + const { data: fetchedTeams } = useAllTeams(); + const allTeams = useMemo(() => fetchedTeams ?? [], [fetchedTeams]); -const toKeyListFilters = (filters: KeyFilterState): KeyListFilterOptions => ({ - teamID: filters["Team ID"].trim() || undefined, - organizationID: filters["Organization ID"].trim() || undefined, - selectedKeyAlias: filters["Key Alias"].trim() || undefined, - userID: filters["User ID"].trim() || undefined, - keyHash: filters["Key Hash"].trim() || undefined, -}); - -export function VirtualKeysTable() { - const { data: fetchedOrganizations, isLoading: isOrgsLoading } = useOrganizations(); - const resolvedOrganizations = useMemo(() => fetchedOrganizations ?? [], [fetchedOrganizations]); const [selectedKey, setSelectedKey] = useState(null); - const [sorting, setSorting] = React.useState([{ id: "created_at", desc: true }]); - const [tablePagination, setTablePagination] = React.useState({ - pageIndex: 0, - pageSize: 50, - }); - const [filters, setFilters] = useState(DEFAULT_KEY_FILTERS); - const [debouncedFilters] = useDebouncedValue(filters, { wait: DEBOUNCE_WAIT_MS }); + const [sorting, setSorting] = useState(DEFAULT_SORTING); + const [tablePagination, setTablePagination] = useState({ pageIndex: 0, pageSize: 50 }); + const [columnFilters, setColumnFilters] = useState([]); + const [filtersOpen, setFiltersOpen] = useState(false); + const [searchInput, setSearchInput] = useState(""); + const [searchQuery] = useDebouncedValue(searchInput, { wait: DEBOUNCE_WAIT_MS }); - const sortBy = sorting.length > 0 ? sorting[0].id : null; - const sortOrder = sorting.length > 0 ? (sorting[0].desc ? "desc" : "asc") : null; + const getFilterValue = useCallback( + (columnId: string): string | undefined => { + const entry = columnFilters.find((filter) => filter.id === columnId); + return typeof entry?.value === "string" && entry.value.trim() ? entry.value.trim() : undefined; + }, + [columnFilters], + ); + + const sortBy = sorting[0]?.id; + const sortOrder = toSortOrder(sorting); + + const keyListOptions = { + teamID: getFilterValue("team_id"), + organizationID: getFilterValue("org_id"), + selectedKeyAlias: searchQuery.trim() || undefined, + userID: getFilterValue("user_id"), + keyHash: getFilterValue("key_hash"), + sortBy, + sortOrder, + expand: "user", + }; const { data: keys, isPending: isLoading, isFetching, - isError, refetch, - } = useKeys(tablePagination.pageIndex + 1, tablePagination.pageSize, { - ...toKeyListFilters(debouncedFilters), - sortBy: sortBy || undefined, - sortOrder: sortOrder || undefined, - expand: "user", - }); - const [expandedAccordions, setExpandedAccordions] = useState>({}); + } = useKeys(tablePagination.pageIndex + 1, tablePagination.pageSize, keyListOptions); const keyList = useMemo(() => keys?.keys ?? [], [keys]); + const rowCount = keys?.total_count ?? 0; - const { data: fetchedTeams, isLoading: isTeamsLoading } = useAllTeams(); - const allTeams = useMemo(() => fetchedTeams ?? [], [fetchedTeams]); - - // Defer the transition so the button stays in loading state until the table - // has rendered with the new data (mirrors the spend-logs pattern) - const isFetchingDeferred = useDeferredValue(isFetching); - const isButtonLoading = (isFetching || isFetchingDeferred) && !isError; - - const handleRefresh = () => { - refetch(); - }; - - const handleFilterChange = (newFilters: Record) => { - setFilters({ - "Team ID": newFilters["Team ID"] || "", - "Organization ID": newFilters["Organization ID"] || "", - "Key Alias": newFilters["Key Alias"] || "", - "User ID": newFilters["User ID"] || "", - "Key Hash": newFilters["Key Hash"] || "", - }); + const handleSearchChange = useCallback((value: string) => { + setSearchInput(value); setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); - }; + }, []); - const handleFilterReset = () => { - setFilters(DEFAULT_KEY_FILTERS); + const handleSortingChange = useCallback>((updaterOrValue) => { + setSorting(updaterOrValue); setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); - }; + }, []); - const totalCount = keys?.total_count ?? 0; + const handleColumnFiltersChange = useCallback>((updaterOrValue) => { + setColumnFilters(updaterOrValue); + setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); + }, []); - const columns: ColumnDef[] = useMemo( - () => [ - { - id: "expander", - header: () => null, - size: 40, - enableSorting: false, - cell: ({ row }) => - row.getCanExpand() ? ( - - ) : null, - }, - { - id: "token", - accessorKey: "token", - header: "Key ID", - size: 100, - enableSorting: true, - cell: (info) => setSelectedKey(info.row.original)} />, - }, - { - id: "key_alias", - accessorKey: "key_alias", - header: "Key Alias", - size: 150, - enableSorting: true, - cell: (info) => { - const value = info.getValue() as string; - const width = info.cell.column.getSize(); - return ( - - {value ?? "-"} - - ); - }, - }, - { - id: "status", - header: "Status", - size: 100, - enableSorting: false, - cell: ({ row }) => { - const key = row.original; - if (key.blocked !== true) { - return ; - } - const isScimBlocked = (key.metadata as Record | null | undefined)?.scim_blocked === true; - const reason = isScimBlocked - ? "Blocked by SCIM (external identity provider deactivated or deleted the owning user)." - : "Blocked. Requests using this key will be rejected with 401."; - return ( - - ); - }, - }, - { - id: "key_name", - accessorKey: "key_name", - header: "Secret Key", - size: 120, - enableSorting: false, - cell: (info) => {info.getValue() as string}, - }, - { - id: "team_alias", - accessorKey: "team_id", - header: "Team", - size: 120, - enableSorting: false, - cell: (info) => { - const teamId = info.getValue() as string | null; - if (!teamId) return "-"; - const team = allTeams.find((t) => t.team_id === teamId); - const displayValue = team?.team_alias || teamId; - const width = info.cell.column.getSize(); - return ( - - {displayValue} - - ); - }, - }, - { - id: "organization_alias", - accessorKey: "org_id", - header: "Organization", - size: 140, - enableSorting: false, - cell: (info) => { - const orgId = info.getValue() as string | null; - if (!orgId) return "-"; - const org = resolvedOrganizations.find((o) => o.organization_id === orgId); - const displayValue = org?.organization_alias || orgId; - const width = info.cell.column.getSize(); - return ( - - {displayValue} - - ); - }, - }, - { - id: "user", - accessorKey: "user", - header: () => ( - - User - - - - - ), - size: 160, - enableSorting: false, - cell: ({ row }) => { - const key = row.original; - const userAlias = key.user?.user_alias ?? null; - const userEmail = key.user?.user_email ?? key.user_email ?? null; - const userId = key.user_id ?? null; - const isDefaultAdmin = userId === "default_user_id"; - const displayValue = userAlias || userEmail || userId; - const width = 160; - - const popoverContent = ( -
- {[ - { label: "User Alias", value: userAlias }, - { label: "User Email", value: userEmail }, - { label: "User ID", value: userId }, - ].map(({ label, value }) => ( -
- {label} - {value ? ( - - {value} - - ) : ( - - - )} -
- ))} -
- ); - - if (isDefaultAdmin && !userAlias && !userEmail) { - return ( - - - - - - ); - } - - return ( - - - {displayValue || "-"} - - - ); - }, - }, - { - id: "created_at", - accessorKey: "created_at", - header: "Created At", - size: 120, - enableSorting: true, - cell: (info) => , - }, - { - id: "created_by", - accessorKey: "created_by", - header: "Created By", - size: 160, - enableSorting: false, - cell: (info) => { - const userId = info.getValue() as string | null; - if (!userId) return "-"; - const key = info.row.original; - const createdByUser = key.created_by_user; - const userAlias = createdByUser?.user_alias ?? null; - const userEmail = createdByUser?.user_email ?? null; - const isDefaultAdmin = userId === "default_user_id"; - const displayValue = userAlias || userEmail || userId; - const width = 160; - - const popoverContent = ( -
- {[ - { label: "User Alias", value: userAlias }, - { label: "User Email", value: userEmail }, - { label: "User ID", value: userId }, - ].map(({ label, value }) => ( -
- {label} - {value ? ( - - {value} - - ) : ( - - - )} -
- ))} -
- ); - - if (isDefaultAdmin && !userAlias && !userEmail) { - return ( - - - - - - ); - } - - return ( - - - {displayValue} - - - ); - }, - }, - { - id: "updated_at", - accessorKey: "updated_at", - header: "Updated At", - size: 120, - enableSorting: true, - cell: (info) => , - }, - { - id: "last_active", - accessorKey: "last_active", - header: () => ( - - Last Active - - - - - ), - size: 130, - enableSorting: false, - cell: (info) => , - }, - { - id: "expires", - accessorKey: "expires", - header: "Expires", - size: 120, - enableSorting: false, - cell: (info) => , - }, - { - id: "spend", - accessorKey: "spend", - header: "Spend (USD)", - size: 100, - enableSorting: true, - cell: (info) => , - }, - { - id: "max_budget", - accessorKey: "max_budget", - header: "Budget (USD)", - size: 110, - enableSorting: true, - cell: (info) => { - const maxBudget = info.getValue() as number | null; - if (maxBudget !== null) { - return `$${formatNumberWithCommas(maxBudget)}`; - } - const teamId = info.row.original.team_id; - const team = allTeams.find((t) => t.team_id === teamId); - if (team?.max_budget != null) { - return `$${formatNumberWithCommas(team.max_budget)} (Team)`; - } - return "Unlimited"; - }, - }, - { - id: "budget_reset_at", - accessorKey: "budget_reset_at", - header: "Budget Reset", - size: 130, - enableSorting: false, - cell: (info) => , - }, - { - id: "models", - accessorKey: "models", - header: "Models", - size: 200, - enableSorting: false, - cell: (info) => { - const models = info.getValue() as string[]; - return ( -
- {Array.isArray(models) ? ( -
- {models.length === 0 ? ( - - All Proxy Models - - ) : ( - <> -
- {models.length > 3 && ( -
- { - setExpandedAccordions((prev) => ({ - ...prev, - [info.row.id]: !prev[info.row.id], - })); - }} - /> -
- )} -
- {models.slice(0, 3).map((model, index) => - model === "all-proxy-models" ? ( - - All Proxy Models - - ) : ( - - - {model.length > 30 - ? `${getModelDisplayName(model).slice(0, 30)}...` - : getModelDisplayName(model)} - - - ), - )} - {models.length > 3 && !expandedAccordions[info.row.id] && ( - - - +{models.length - 3} {models.length - 3 === 1 ? "more model" : "more models"} - - - )} - {expandedAccordions[info.row.id] && ( -
- {models.slice(3).map((model, index) => - model === "all-proxy-models" ? ( - - All Proxy Models - - ) : ( - - - {model.length > 30 - ? `${getModelDisplayName(model).slice(0, 30)}...` - : getModelDisplayName(model)} - - - ), - )} -
- )} -
-
- - )} -
- ) : null} -
- ); - }, - }, - { - id: "rate_limits", - header: "Rate Limits", - size: 140, - enableSorting: false, - cell: ({ row }) => { - const key = row.original; - return ( -
-
TPM: {key.tpm_limit !== null ? key.tpm_limit : "Unlimited"}
-
RPM: {key.rpm_limit !== null ? key.rpm_limit : "Unlimited"}
-
- ); - }, - }, - ], - [allTeams, resolvedOrganizations], + const columns = useMemo( + () => getKeyTableColumns({ allTeams, organizations, onSelectKey: setSelectedKey }), + [allTeams, organizations], ); - const filterOptions: FilterOption[] = [ - { - name: "Team ID", - label: "Team ID", - isSearchable: true, - loading: isTeamsLoading, - searchFn: async (searchText: string) => { - if (!allTeams || allTeams.length === 0) return []; + const teamOptions = useMemo( + () => + allTeams.map((team) => ({ + label: team.team_alias || team.team_id, + value: team.team_id, + sublabel: team.team_alias ? team.team_id : undefined, + })), + [allTeams], + ); - const filteredTeams = allTeams.filter( - (team) => - team.team_id.toLowerCase().includes(searchText.toLowerCase()) || - (team.team_alias && team.team_alias.toLowerCase().includes(searchText.toLowerCase())), - ); + const orgOptions = useMemo( + () => + organizations + .filter((org) => org.organization_id) + .map((org) => { + const id = org.organization_id as string; + return { label: org.organization_alias || id, value: id, sublabel: org.organization_alias ? id : undefined }; + }), + [organizations], + ); - return filteredTeams.map((team) => ({ - label: `${team.team_alias || team.team_id} (${team.team_id})`, - value: team.team_id, - })); - }, + const formatFilterValue = useCallback( + (columnId: string, value: unknown): string => { + const raw = String(value); + if (columnId === "team_id") { + return allTeams.find((team) => team.team_id === raw)?.team_alias || raw; + } + if (columnId === "org_id") { + return organizations.find((org) => org.organization_id === raw)?.organization_alias || raw; + } + return raw; }, - { - name: "Organization ID", - label: "Organization ID", - isSearchable: true, - loading: isOrgsLoading, - searchFn: async (searchText: string) => { - if (!resolvedOrganizations || resolvedOrganizations.length === 0) return []; + [allTeams, organizations], + ); - const filteredOrgs = resolvedOrganizations.filter( - (org) => org.organization_id?.toLowerCase().includes(searchText.toLowerCase()) ?? false, - ); - - return filteredOrgs - .filter((org) => org.organization_id !== null && org.organization_id !== undefined) - .map((org) => ({ - label: `${org.organization_id || "Unknown"} (${org.organization_id})`, - value: org.organization_id as string, - })); - }, - }, - { - name: "Key Alias", - label: "Key Alias", - customComponent: PaginatedKeyAliasSelect, - }, - { - name: "User ID", - label: "User ID", - isSearchable: false, - }, - { - name: "Key Hash", - label: "Key ID", - isSearchable: false, - }, - ]; - - const table = useReactTable({ - data: keyList, - columns: columns.filter((col) => col.id !== "expander"), - columnResizeMode: "onChange", - columnResizeDirection: "ltr", - state: { - sorting, - pagination: tablePagination, - }, - onSortingChange: (updaterOrValue) => { - const newSorting = typeof updaterOrValue === "function" ? updaterOrValue(sorting) : updaterOrValue; - setSorting(newSorting); - setTablePagination((prev) => ({ ...prev, pageIndex: 0 })); - }, - onPaginationChange: setTablePagination, - getCoreRowModel: getCoreRowModel(), - enableSorting: true, - manualSorting: true, - manualPagination: true, - pageCount: Math.ceil(totalCount / tablePagination.pageSize), - }); - - const { pageIndex, pageSize } = table.getState().pagination; - const start = pageIndex * pageSize + 1; - const end = Math.min((pageIndex + 1) * pageSize, totalCount); - const rangeLabel = `${start} - ${end}`; - return ( -
- {selectedKey ? ( + if (selectedKey) { + return ( +
setSelectedKey(null)} keyData={selectedKey} teams={allTeams} + onDelete={refetch} /> - ) : ( -
-
- + ); + } + + return ( +
+ } + title="Virtual Keys" + subtitle="Every key that authenticates requests to the gateway." + actions={headerActions} + /> + row.token} + defaultColumnVisibility={KEY_TABLE_HIDDEN_COLUMNS} + sortingMode="server" + sorting={sorting} + onSortingChange={handleSortingChange} + paginationMode="server" + pagination={tablePagination} + onPaginationChange={setTablePagination} + rowCount={rowCount} + filterMode="server" + columnFilters={columnFilters} + onColumnFiltersChange={handleColumnFiltersChange} + enableColumnResizing + columnResizeMode="onChange" + isLoading={isLoading} + loadingMessage="Loading keys..." + noDataMessage="No keys found" + maxBodyHeight="calc(75vh - 210px)" + size="compact" + toolbar={(table) => ( + <> + refetch?.()} + isRefreshing={isFetching} + onOpenFilters={() => setFiltersOpen(true)} + filterLabels={FILTER_LABELS} + formatFilterValue={formatFilterValue} /> -
- -
-
- {isLoading ? ( - - ) : ( - - Showing {rangeLabel} of {totalCount} results - + + {({ get, set }) => ( + <> + + set("team_id", value)} + placeholder="Select a team…" + emptyText="No teams found" + /> + + + set("org_id", value)} + placeholder="Select an organization…" + emptyText="No organizations found" + /> + + + set("user_id", event.target.value)} + placeholder="Enter User ID…" + /> + + + set("key_hash", event.target.value)} + placeholder="Enter Key ID…" + /> + + )} - - } - onClick={handleRefresh} - disabled={isButtonLoading} - title="Fetch data" - > - {isButtonLoading ? "Fetching" : "Fetch"} - -
- -
- {isLoading ? ( - - ) : ( - - Page {pageIndex + 1} of {table.getPageCount()} - - )} - - {isLoading ? ( - - ) : ( - - )} - - {isLoading ? ( - - ) : ( - - )} -
-
-
-
-
- - - {table.getHeaderGroups().map((headerGroup) => ( - - {headerGroup.headers.map((header) => ( - { - const resizer = document.querySelector(`[data-header-id="${header.id}"] .resizer`); - if (resizer) { - (resizer as HTMLElement).style.opacity = "0.5"; - } - }} - onMouseLeave={() => { - const resizer = document.querySelector(`[data-header-id="${header.id}"] .resizer`); - if (resizer && !header.column.getIsResizing()) { - (resizer as HTMLElement).style.opacity = "0"; - } - }} - onClick={header.column.getCanSort() ? header.column.getToggleSortingHandler() : undefined} - > -
-
- {header.isPlaceholder - ? null - : flexRender(header.column.columnDef.header, header.getContext())} -
- {header.id !== "actions" && header.column.getCanSort() && ( -
- {header.column.getIsSorted() ? ( - { - asc: , - desc: , - }[header.column.getIsSorted() as string] - ) : ( - - )} -
- )} -
header.column.resetSize()} - onMouseDown={header.getResizeHandler()} - onTouchStart={header.getResizeHandler()} - className={`resizer ${table.options.columnResizeDirection} ${header.column.getIsResizing() ? "isResizing" : ""}`} - style={{ - position: "absolute", - right: 0, - top: 0, - height: "100%", - width: "5px", - background: header.column.getIsResizing() ? "#3b82f6" : "transparent", - cursor: "col-resize", - userSelect: "none", - touchAction: "none", - opacity: header.column.getIsResizing() ? 1 : 0, - }} - /> -
- - ))} - - ))} - - - {isLoading ? ( - - -
-

🚅 Loading keys...

-
-
-
- ) : keyList.length > 0 ? ( - table.getRowModel().rows.map((row) => ( - - {row.getVisibleCells().map((cell) => ( - 3 ? "px-0" : ""}`} - > - {flexRender(cell.column.columnDef.cell, cell.getContext())} - - ))} - - )) - ) : ( - - -
-

No keys found

-
-
-
- )} -
-
-
-
-
-
- )} + + + )} + />
); } diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx new file mode 100644 index 00000000000..e2dc48fed9d --- /dev/null +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx @@ -0,0 +1,353 @@ +"use client"; + +import { InfoCircleOutlined } from "@ant-design/icons"; +import { ColumnDef } from "@tanstack/react-table"; +import { Popover, Typography } from "antd"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { Skeleton } from "@/components/ui/skeleton"; +import { + DateCell, + IdCell, + IdentityCell, + ModelsCell, + SpendBudgetCell, + StatusBadge, + type StatusTone, +} from "@/components/shared/table_cells"; + +import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; +import { KeyResponse, Team } from "../key_team_helpers/key_list"; +import { Organization } from "../networking"; + +interface KeyStatus { + tone: StatusTone; + label: string; + tooltip?: string; +} + +const getKeyStatus = (key: KeyResponse): KeyStatus => { + if (key.blocked === true) { + const isScimBlocked = (key.metadata as Record | null | undefined)?.scim_blocked === true; + return { + tone: "error", + label: "Blocked", + tooltip: isScimBlocked + ? "Blocked by SCIM (external identity provider deactivated or deleted the owning user)." + : "Blocked. Requests using this key will be rejected with 401.", + }; + } + const expiresAt = key.expires ? Date.parse(key.expires) : Number.NaN; + if (!Number.isNaN(expiresAt) && expiresAt < Date.now()) { + return { tone: "warning", label: "Expired", tooltip: "This key has passed its expiry date." }; + } + return { tone: "success", label: "Active" }; +}; + +const UserPopoverCell = ({ + userAlias, + userEmail, + userId, + width, +}: { + userAlias: string | null; + userEmail: string | null; + userId: string | null; + width: number; +}) => { + const displayValue = userAlias || userEmail || userId; + const isDefaultAdmin = userId === "default_user_id"; + + const popoverContent = ( +
+ {[ + { label: "User Alias", value: userAlias }, + { label: "User Email", value: userEmail }, + { label: "User ID", value: userId }, + ].map(({ label, value }) => ( +
+ {label} + {value ? ( + + {value} + + ) : ( + - + )} +
+ ))} +
+ ); + + if (isDefaultAdmin && !userAlias && !userEmail) { + return ( + + + + + + ); + } + + return ( + + + {displayValue || "-"} + + + ); +}; + +const InfoHeader = ({ label, tooltip }: { label: string; tooltip: string }) => ( + + {label} + + + + +); + +interface KeyTableColumnsDeps { + allTeams: Team[]; + organizations: Organization[]; + onSelectKey: (key: KeyResponse) => void; +} + +export const getKeyTableColumns = ({ + allTeams, + organizations, + onSelectKey, +}: KeyTableColumnsDeps): ColumnDef[] => [ + { + id: "key_alias", + accessorKey: "key_alias", + meta: { + title: "Key", + renderSkeleton: () => ( +
+ +
+ + +
+
+ ), + }, + header: ({ column }) => , + size: 260, + enableSorting: true, + cell: ({ row }) => { + const status = getKeyStatus(row.original); + return ( + + } + onClick={() => onSelectKey(row.original)} + /> + ); + }, + }, + { + id: "token", + accessorKey: "token", + meta: { title: "Key ID" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: (info) => onSelectKey(info.row.original)} />, + }, + { + id: "team_alias", + accessorKey: "team_id", + meta: { title: "Team" }, + header: "Team", + size: 120, + enableSorting: false, + cell: (info) => { + const teamId = info.getValue() as string | null; + if (!teamId) return "-"; + const team = allTeams.find((t) => t.team_id === teamId); + const displayValue = team?.team_alias || teamId; + const width = info.cell.column.getSize(); + return ( + + {displayValue} + + ); + }, + }, + { + id: "organization_alias", + accessorKey: "org_id", + meta: { title: "Organization" }, + header: "Organization", + size: 140, + enableSorting: false, + cell: (info) => { + const orgId = info.getValue() as string | null; + if (!orgId) return "-"; + const org = organizations.find((o) => o.organization_id === orgId); + const displayValue = org?.organization_alias || orgId; + const width = info.cell.column.getSize(); + return ( + + {displayValue} + + ); + }, + }, + { + id: "user", + accessorKey: "user", + meta: { title: "User" }, + header: () => ( + + ), + size: 160, + enableSorting: false, + cell: ({ row }) => { + const key = row.original; + return ( + + ); + }, + }, + { + id: "created_at", + accessorKey: "created_at", + meta: { title: "Created At" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: (info) => , + }, + { + id: "created_by", + accessorKey: "created_by", + meta: { title: "Created By" }, + header: "Created By", + size: 160, + enableSorting: false, + cell: (info) => { + const userId = info.getValue() as string | null; + if (!userId) return "-"; + const createdByUser = info.row.original.created_by_user; + return ( + + ); + }, + }, + { + id: "updated_at", + accessorKey: "updated_at", + meta: { title: "Updated At" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: (info) => , + }, + { + id: "last_active", + accessorKey: "last_active", + meta: { title: "Last Active" }, + header: () => ( + + ), + size: 130, + enableSorting: false, + cell: (info) => , + }, + { + id: "expires", + accessorKey: "expires", + meta: { title: "Expires" }, + header: "Expires", + size: 120, + enableSorting: false, + cell: (info) => , + }, + { + id: "spend", + accessorKey: "spend", + meta: { title: "Spend / Budget", skeleton: "meter" }, + header: ({ column }) => , + size: 180, + enableSorting: true, + cell: ({ row }) => { + const teamId = row.original.team_id; + const team = allTeams.find((t) => t.team_id === teamId); + return ( + + ); + }, + }, + { + id: "budget_reset_at", + accessorKey: "budget_reset_at", + meta: { title: "Budget Reset" }, + header: "Budget Reset", + size: 130, + enableSorting: false, + cell: (info) => , + }, + { + id: "models", + accessorKey: "models", + meta: { title: "Models", skeleton: "chips" }, + header: "Models", + size: 220, + enableSorting: false, + cell: (info) => , + }, + { + id: "rate_limits", + meta: { title: "Rate Limits" }, + header: "Rate Limits", + size: 140, + enableSorting: false, + cell: ({ row }) => { + const key = row.original; + return ( +
+
TPM: {key.tpm_limit !== null ? key.tpm_limit : "Unlimited"}
+
RPM: {key.rpm_limit !== null ? key.rpm_limit : "Unlimited"}
+
+ ); + }, + }, +]; + +export const KEY_TABLE_HIDDEN_COLUMNS: Record = { + token: false, + organization_alias: false, + created_by: false, + updated_at: false, + expires: false, + budget_reset_at: false, + rate_limits: false, +}; diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx index ef0d842ad3e..d8d4c9392dc 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx @@ -303,6 +303,38 @@ describe("DataTable loading", () => { // per-column widths differ instead of every cell sharing one fixed width expect(new Set(bars.map((bar) => bar.className)).size).toBeGreaterThan(1); }); + + it("renders shape-specific skeletons for badge, chips, and meter columns", () => { + const columns: ColumnDef[] = [ + { id: "badge", header: "Badge", meta: { skeleton: "badge" }, cell: () => null }, + { id: "chips", header: "Chips", meta: { skeleton: "chips" }, cell: () => null }, + { id: "meter", header: "Meter", meta: { skeleton: "meter" }, cell: () => null }, + ]; + render(); + + const firstRow = screen.getAllByTestId("skeleton-row").at(0); + const cells = Array.from(firstRow?.querySelectorAll("td") ?? []); + const barsIn = (cell: Element | undefined) => cell?.querySelectorAll('[data-slot="skeleton"]').length ?? 0; + + // badge = a single pill, chips = three pills, meter = value bar + track bar + expect(barsIn(cells[0])).toBe(1); + expect(cells[0]?.querySelector('[data-slot="skeleton"]')?.className).toContain("rounded-full"); + expect(barsIn(cells[1])).toBe(3); + expect(barsIn(cells[2])).toBe(2); + }); + + it("uses a column's renderSkeleton override when provided", () => { + const columns: ColumnDef[] = [ + { + id: "custom", + header: "Custom", + meta: { renderSkeleton: () =>
loading
}, + cell: () => null, + }, + ]; + render(); + expect(screen.getAllByTestId("custom-skeleton").length).toBeGreaterThan(0); + }); }); describe("DataTable column visibility", () => { diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx index 758ca5a597b..bc13318dd58 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx @@ -338,7 +338,11 @@ const SKELETON_WIDTHS = ["w-[58%]", "w-[44%]", "w-[70%]", "w-[50%]", "w-[64%]", function SkeletonCell({ column, index }: { column: Column | undefined; index: number }) { const meta = column?.columnDef.meta; const width = SKELETON_WIDTHS[index % SKELETON_WIDTHS.length]; - if (meta?.skeleton === "twoLine") { + const shape = meta?.skeleton; + if (meta?.renderSkeleton !== undefined) { + return <>{meta.renderSkeleton()}; + } + if (shape === "twoLine") { return (
@@ -346,6 +350,26 @@ function SkeletonCell({ column, index }: { column: Column
); } + if (shape === "badge") { + return ; + } + if (shape === "chips") { + return ( +
+ + + +
+ ); + } + if (shape === "meter") { + return ( +
+ + +
+ ); + } return ; } diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts b/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts index 0f14c277c6f..eff4e0cb7db 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/columnMeta.ts @@ -1,4 +1,5 @@ import type { RowData } from "@tanstack/react-table"; +import type * as React from "react"; import type { ColumnPinnedSide, DataTableSkeletonShape } from "./types"; @@ -10,5 +11,7 @@ declare module "@tanstack/react-table" { title?: string; pinned?: ColumnPinnedSide; skeleton?: DataTableSkeletonShape; + /** Full control over this column's loading skeleton, for cells the built-in shapes can't mirror. */ + renderSkeleton?: () => React.ReactNode; } } diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts index f5130b4c823..672ab512ef4 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts @@ -18,7 +18,7 @@ export type FilterMode = "none" | "client" | "server"; export type ColumnResizeMode = "onEnd" | "onChange"; export type DataTableSize = "compact" | "default"; export type ColumnPinnedSide = "left" | "right"; -export type DataTableSkeletonShape = "text" | "twoLine"; +export type DataTableSkeletonShape = "text" | "twoLine" | "badge" | "chips" | "meter"; export interface DataTableProps { data: TData[]; diff --git a/ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx b/ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx new file mode 100644 index 00000000000..f7a313271da --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx @@ -0,0 +1,31 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; + +import { PageHeader } from "./PageHeader"; + +describe("PageHeader", () => { + it("renders the title as a heading", () => { + render(); + expect(screen.getByRole("heading", { name: "Virtual Keys" })).toBeInTheDocument(); + }); + + it("renders the subtitle, icon, and actions when provided", () => { + render( + } + actions={} + />, + ); + expect(screen.getByText("Every key that authenticates requests")).toBeInTheDocument(); + expect(screen.getByTestId("icon")).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Create New Key" })).toBeInTheDocument(); + }); + + it("omits the optional slots when not provided", () => { + render(); + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + expect(document.querySelector("p")).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/PageHeader.tsx b/ui/litellm-dashboard/src/components/shared/PageHeader.tsx new file mode 100644 index 00000000000..34d478e1cd9 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/PageHeader.tsx @@ -0,0 +1,29 @@ +"use client"; + +import * as React from "react"; + +interface PageHeaderProps { + title: React.ReactNode; + subtitle?: React.ReactNode; + icon?: React.ReactNode; + actions?: React.ReactNode; +} + +export function PageHeader({ title, subtitle, icon, actions }: PageHeaderProps) { + return ( +
+
+ {icon != null && ( + + {icon} + + )} +
+

{title}

+ {subtitle != null &&

{subtitle}

} +
+
+ {actions != null &&
{actions}
} +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/shared/SearchSelect.test.tsx b/ui/litellm-dashboard/src/components/shared/SearchSelect.test.tsx new file mode 100644 index 00000000000..acf50d282b4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/SearchSelect.test.tsx @@ -0,0 +1,64 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +import { SearchSelect } from "./SearchSelect"; + +const OPTIONS = [ + { label: "Acme Prod", value: "team-1" }, + { label: "Growth", value: "team-2" }, + { label: "Data Team", value: "team-3" }, +]; + +describe("SearchSelect", () => { + it("renders the placeholder when nothing is selected", () => { + render(); + expect(screen.getByPlaceholderText("Select Team…")).toBeInTheDocument(); + }); + + it("shows the selected option's label in the field", () => { + render(); + expect(screen.getByRole("combobox")).toHaveValue("Growth"); + }); + + it("shows a clear control only when a value is selected", () => { + const { rerender } = render(); + expect(document.querySelector('[data-slot="combobox-clear"]')).toBeNull(); + rerender(); + expect(document.querySelector('[data-slot="combobox-clear"]')).not.toBeNull(); + }); + + it("filters the options client-side as you type", async () => { + const user = userEvent.setup(); + render(); + const input = screen.getByRole("combobox"); + await user.click(input); + await user.type(input, "grow"); + expect(await screen.findByText("Growth")).toBeInTheDocument(); + expect(screen.queryByText("Acme Prod")).not.toBeInTheDocument(); + }); + + it("renders a muted sublabel and matches it when searching", async () => { + const user = userEvent.setup(); + render( + , + ); + const input = screen.getByRole("combobox"); + await user.click(input); + expect(await screen.findByText("team-abc-123")).toBeInTheDocument(); + await user.type(input, "abc-123"); + expect(await screen.findByText("Acme Prod")).toBeInTheDocument(); + }); + + it("selects an option and reports its value", async () => { + const onValueChange = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByText("Growth")); + expect(onValueChange).toHaveBeenCalledWith("team-2"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/SearchSelect.tsx b/ui/litellm-dashboard/src/components/shared/SearchSelect.tsx new file mode 100644 index 00000000000..c29e099a1c6 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/SearchSelect.tsx @@ -0,0 +1,76 @@ +"use client"; + +import { + Combobox, + ComboboxContent, + ComboboxEmpty, + ComboboxInput, + ComboboxItem, + ComboboxList, +} from "@/components/ui/combobox"; + +export interface SearchSelectOption { + label: string; + value: string; + /** Optional muted second line (e.g. an id); also matched when searching. */ + sublabel?: string; +} + +interface SearchSelectProps { + options: SearchSelectOption[]; + value?: string; + onValueChange: (value: string) => void; + placeholder?: string; + emptyText?: string; + disabled?: boolean; + className?: string; +} + +export function SearchSelect({ + options, + value, + onValueChange, + placeholder = "Select…", + emptyText = "No results", + disabled = false, + className, +}: SearchSelectProps) { + const selected = options.find((option) => option.value === value) ?? null; + + return ( + onValueChange(item?.value ?? "")} + isItemEqualToValue={(a: SearchSelectOption, b: SearchSelectOption) => a.value === b.value} + itemToStringLabel={(item: SearchSelectOption) => item.label} + filter={(item: SearchSelectOption, query: string) => { + const q = query.trim().toLowerCase(); + if (!q) return true; + return item.label.toLowerCase().includes(q) || (item.sublabel?.toLowerCase().includes(q) ?? false); + }} + disabled={disabled} + > + + + {emptyText} + + {(item: SearchSelectOption) => ( + + + {item.label} + {item.sublabel != null && item.sublabel !== "" && ( + {item.sublabel} + )} + + + )} + + + + ); +} diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.test.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.test.tsx new file mode 100644 index 00000000000..db4e93c7cb2 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.test.tsx @@ -0,0 +1,38 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +import { IdentityCell } from "./identity_cell"; + +describe("IdentityCell", () => { + it("renders the title", () => { + render(); + expect(screen.getByText("prod-gateway")).toBeInTheDocument(); + }); + + it("renders the subtitle and an inline badge together", () => { + render(Active} />); + expect(screen.getByText("sk-...v0Pw")).toBeInTheDocument(); + expect(screen.getByText("Active")).toBeInTheDocument(); + }); + + it("omits the subtitle row when there is no subtitle or badge", () => { + render(); + expect(document.querySelector("span.font-mono")).toBeNull(); + }); + + it("renders a static div (no button) when not clickable", () => { + render(); + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); + + it("renders a clickable button and fires onClick", async () => { + const onClick = vi.fn(); + const user = userEvent.setup(); + render(); + const button = screen.getByRole("button"); + expect(button.querySelector(".lucide-chevron-right")).not.toBeNull(); + await user.click(button); + expect(onClick).toHaveBeenCalledTimes(1); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.tsx new file mode 100644 index 00000000000..4d3e3d8e4dd --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/table_cells/identity_cell.tsx @@ -0,0 +1,48 @@ +"use client"; + +import { ChevronRight } from "lucide-react"; +import * as React from "react"; + +import { cn } from "@/lib/cva.config"; + +interface IdentityCellProps { + title: React.ReactNode; + subtitle?: React.ReactNode; + badge?: React.ReactNode; + onClick?: () => void; + className?: string; + titleClassName?: string; +} + +export function IdentityCell({ title, subtitle, badge, onClick, className, titleClassName }: IdentityCellProps) { + const hasSubtitleRow = (subtitle != null && subtitle !== "") || badge != null; + + const body = ( +
+ {title} + {hasSubtitleRow && ( + + {subtitle != null && subtitle !== "" && ( + {subtitle} + )} + {badge} + + )} +
+ ); + + if (onClick != null) { + return ( + + ); + } + + return
{body}
; +} diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/index.ts b/ui/litellm-dashboard/src/components/shared/table_cells/index.ts index e189413d43d..9fdd04d169c 100644 --- a/ui/litellm-dashboard/src/components/shared/table_cells/index.ts +++ b/ui/litellm-dashboard/src/components/shared/table_cells/index.ts @@ -1,5 +1,8 @@ export { CellTooltip } from "./cell_tooltip"; export { DateCell, formatCellDate, formatFullTimestamp, type DatePrecision } from "./date_cell"; export { IdCell, type IdCellVariant } from "./id_cell"; +export { IdentityCell } from "./identity_cell"; +export { ModelsCell } from "./models_cell"; export { MoneyCell } from "./money_cell"; +export { SpendBudgetCell } from "./spend_budget_cell"; export { StatusBadge, type StatusTone } from "./status_badge"; diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/models_cell.test.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/models_cell.test.tsx new file mode 100644 index 00000000000..d3fad1d3244 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/table_cells/models_cell.test.tsx @@ -0,0 +1,45 @@ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it } from "vitest"; + +import { ModelsCell } from "./models_cell"; + +describe("ModelsCell", () => { + it("shows 'All Proxy Models' when the list is empty, null, or undefined", () => { + const { rerender } = render(); + expect(screen.getByText("All Proxy Models")).toBeInTheDocument(); + rerender(); + expect(screen.getByText("All Proxy Models")).toBeInTheDocument(); + rerender(); + expect(screen.getByText("All Proxy Models")).toBeInTheDocument(); + }); + + it("renders every model with no overflow badge when at or below the limit", () => { + render(); + expect(screen.getByText("gpt-4o")).toBeInTheDocument(); + expect(screen.getByText("claude-sonnet-4-5")).toBeInTheDocument(); + expect(screen.getByText("o3-mini")).toBeInTheDocument(); + expect(screen.queryByText(/more$/)).not.toBeInTheDocument(); + }); + + it("collapses models beyond the limit into a '+N more' badge", () => { + render(); + expect(screen.getByText("a")).toBeInTheDocument(); + expect(screen.getByText("b")).toBeInTheDocument(); + expect(screen.queryByText("c")).not.toBeInTheDocument(); + expect(screen.getByText("+3 more")).toBeInTheDocument(); + }); + + it("reveals the hidden models in a tooltip on hover", async () => { + const user = userEvent.setup(); + render(); + await user.hover(screen.getByText("+2 more")); + expect(await screen.findByText("c")).toBeInTheDocument(); + expect(await screen.findByText("d")).toBeInTheDocument(); + }); + + it("labels the all-proxy-models wildcard", () => { + render(); + expect(screen.getByText("All Proxy Models")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/models_cell.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/models_cell.tsx new file mode 100644 index 00000000000..712d8511c78 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/table_cells/models_cell.tsx @@ -0,0 +1,56 @@ +"use client"; + +import { getModelDisplayName } from "@/components/key_team_helpers/fetch_available_models_team_key"; +import { Badge } from "@/components/ui/badge"; + +import { CellTooltip } from "./cell_tooltip"; + +interface ModelsCellProps { + models: string[] | null | undefined; + maxVisible?: number; +} + +const WILDCARD_MODEL = "all-proxy-models"; + +const formatModel = (model: string): string => { + if (model === WILDCARD_MODEL) { + return "All Proxy Models"; + } + const name = getModelDisplayName(model); + return name.length > 30 ? `${name.slice(0, 30)}...` : name; +}; + +export function ModelsCell({ models, maxVisible = 3 }: ModelsCellProps) { + if (!Array.isArray(models) || models.length === 0) { + return All Proxy Models; + } + + const visible = models.slice(0, maxVisible); + const overflow = models.slice(maxVisible); + + return ( +
+ {visible.map((model, index) => ( + + {formatModel(model)} + + ))} + {overflow.length > 0 && ( + + {overflow.map((model, index) => ( + {formatModel(model)} + ))} +
+ } + trigger={ + + +{overflow.length} more + + } + /> + )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.test.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.test.tsx new file mode 100644 index 00000000000..707441aef1d --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.test.tsx @@ -0,0 +1,53 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; + +import { SpendBudgetCell } from "./spend_budget_cell"; + +const indicator = (container: HTMLElement) => container.querySelector('[data-slot="meter-indicator"]'); + +describe("SpendBudgetCell", () => { + it("shows Unlimited and renders no meter when there is no budget", () => { + const { container } = render(); + expect(screen.getByText("· Unlimited")).toBeInTheDocument(); + expect(screen.queryByRole("meter")).not.toBeInTheDocument(); + expect(indicator(container)).toBeNull(); + }); + + it("shows $0.00 for zero or undefined spend, never a hyphen", () => { + const { rerender } = render(); + expect(screen.getByText("$0.00")).toBeInTheDocument(); + expect(screen.queryByText("-")).not.toBeInTheDocument(); + rerender(); + expect(screen.getByText("$0.00")).toBeInTheDocument(); + expect(screen.queryByText("-")).not.toBeInTheDocument(); + }); + + it("renders a meter carrying the spend and budget when a budget exists", () => { + render(); + const meter = screen.getByRole("meter"); + expect(meter).toHaveAttribute("aria-valuenow", "25"); + expect(meter).toHaveAttribute("aria-valuemax", "100"); + expect(screen.getByText("of $100")).toBeInTheDocument(); + }); + + it("keeps the default tone below 80% usage", () => { + const { container } = render(); + expect(indicator(container)?.className).toContain("bg-primary"); + }); + + it("switches to the warning tone at 80% usage", () => { + const { container } = render(); + expect(indicator(container)?.className).toContain("bg-amber-500"); + }); + + it("switches to the over tone above 100% usage", () => { + const { container } = render(); + expect(indicator(container)?.className).toContain("bg-destructive"); + }); + + it("falls back to the team budget and labels it", () => { + render(); + expect(screen.getByText("of $200 (Team)")).toBeInTheDocument(); + expect(screen.getByRole("meter")).toHaveAttribute("aria-valuemax", "200"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.tsx new file mode 100644 index 00000000000..10956f23b1c --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/table_cells/spend_budget_cell.tsx @@ -0,0 +1,44 @@ +"use client"; + +import { Meter, MeterIndicator, MeterTrack } from "@/components/ui/meter"; +import { formatNumberWithCommas, getSpendString } from "@/utils/dataUtils"; + +interface SpendBudgetCellProps { + spend: number | null | undefined; + maxBudget: number | null | undefined; + teamMaxBudget?: number | null; +} + +const meterTone = (pct: number): "default" | "warning" | "over" => { + if (pct > 100) return "over"; + if (pct >= 80) return "warning"; + return "default"; +}; + +export function SpendBudgetCell({ spend, maxBudget, teamMaxBudget }: SpendBudgetCellProps) { + const spendValue = typeof spend === "number" && !Number.isNaN(spend) ? spend : 0; + const budget = maxBudget ?? teamMaxBudget ?? null; + const isTeamBudget = maxBudget == null && teamMaxBudget != null; + const hasBudget = typeof budget === "number" && budget > 0; + const pct = hasBudget ? (spendValue / budget) * 100 : 0; + + const spendText = spendValue > 0 ? getSpendString(spendValue, 4) : "$0.00"; + const budgetLabel = + budget === null ? "· Unlimited" : `of $${formatNumberWithCommas(budget)}${isTeamBudget ? " (Team)" : ""}`; + + return ( +
+
+ {spendText}{" "} + {budgetLabel} +
+ {hasBudget && ( + + + + + + )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/ui/combobox.tsx b/ui/litellm-dashboard/src/components/ui/combobox.tsx new file mode 100644 index 00000000000..2854928140e --- /dev/null +++ b/ui/litellm-dashboard/src/components/ui/combobox.tsx @@ -0,0 +1,266 @@ +"use client"; + +import * as React from "react"; +import { Combobox as ComboboxPrimitive } from "@base-ui/react"; + +import { cn } from "@/lib/cva.config"; +import { Button } from "@/components/ui/button"; +import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; +import { ChevronDownIcon, XIcon, CheckIcon } from "lucide-react"; + +const Combobox = ComboboxPrimitive.Root; + +function ComboboxValue({ ...props }: ComboboxPrimitive.Value.Props) { + return ; +} + +const ComboboxTrigger = React.forwardRef< + React.ComponentRef, + ComboboxPrimitive.Trigger.Props +>(({ className, children, ...props }, ref) => { + return ( + + {children} + + + ); +}); +ComboboxTrigger.displayName = "ComboboxTrigger"; + +function ComboboxClear({ className, ...props }: ComboboxPrimitive.Clear.Props) { + return ( + } + className={cn(className)} + {...props} + > + + + ); +} + +function ComboboxInput({ + className, + children, + disabled = false, + showTrigger = true, + showClear = false, + ...props +}: ComboboxPrimitive.Input.Props & { + showTrigger?: boolean; + showClear?: boolean; +}) { + return ( + + } {...props} /> + + {showTrigger && ( + } + data-slot="input-group-button" + className="group-has-data-[slot=combobox-clear]/input-group:hidden data-pressed:bg-transparent" + disabled={disabled} + /> + )} + {showClear && } + + {children} + + ); +} + +function ComboboxContent({ + className, + side = "bottom", + sideOffset = 6, + align = "start", + alignOffset = 0, + anchor, + ...props +}: ComboboxPrimitive.Popup.Props & + Pick) { + return ( + + + + + + ); +} + +function ComboboxList({ className, ...props }: ComboboxPrimitive.List.Props) { + return ( + + ); +} + +function ComboboxItem({ className, children, ...props }: ComboboxPrimitive.Item.Props) { + return ( + + {children} + } + > + + + + ); +} + +function ComboboxGroup({ className, ...props }: ComboboxPrimitive.Group.Props) { + return ; +} + +function ComboboxLabel({ className, ...props }: ComboboxPrimitive.GroupLabel.Props) { + return ( + + ); +} + +function ComboboxCollection({ ...props }: ComboboxPrimitive.Collection.Props) { + return ; +} + +function ComboboxEmpty({ className, ...props }: ComboboxPrimitive.Empty.Props) { + return ( + + ); +} + +function ComboboxSeparator({ className, ...props }: ComboboxPrimitive.Separator.Props) { + return ( + + ); +} + +function ComboboxChips({ + className, + ...props +}: React.ComponentPropsWithRef & ComboboxPrimitive.Chips.Props) { + return ( + + ); +} + +function ComboboxChip({ + className, + children, + showRemove = true, + ...props +}: ComboboxPrimitive.Chip.Props & { + showRemove?: boolean; +}) { + return ( + + {children} + {showRemove && ( + } + className="-ml-1 opacity-50 hover:opacity-100" + data-slot="combobox-chip-remove" + > + + + )} + + ); +} + +function ComboboxChipsInput({ className, ...props }: ComboboxPrimitive.Input.Props) { + return ( + + ); +} + +function useComboboxAnchor() { + return React.useRef(null); +} + +export { + Combobox, + ComboboxInput, + ComboboxContent, + ComboboxList, + ComboboxItem, + ComboboxGroup, + ComboboxLabel, + ComboboxCollection, + ComboboxEmpty, + ComboboxSeparator, + ComboboxChips, + ComboboxChip, + ComboboxChipsInput, + ComboboxTrigger, + ComboboxValue, + useComboboxAnchor, +}; diff --git a/ui/litellm-dashboard/src/components/ui/input-group.tsx b/ui/litellm-dashboard/src/components/ui/input-group.tsx new file mode 100644 index 00000000000..8ee9b7f17bd --- /dev/null +++ b/ui/litellm-dashboard/src/components/ui/input-group.tsx @@ -0,0 +1,140 @@ +"use client"; + +import * as React from "react"; +import { type VariantProps } from "cva"; + +import { cn, cva } from "@/lib/cva.config"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Textarea } from "@/components/ui/textarea"; + +function InputGroup({ className, ...props }: React.ComponentProps<"div">) { + return ( +
[data-align=block-end]]:h-auto has-[>[data-align=block-end]]:flex-col has-[>[data-align=block-start]]:h-auto has-[>[data-align=block-start]]:flex-col has-[>textarea]:h-auto dark:bg-input/30 dark:has-[[data-slot][aria-invalid=true]]:ring-destructive/40 has-[>[data-align=block-end]]:[&>input]:pt-3 has-[>[data-align=block-start]]:[&>input]:pb-3 has-[>[data-align=inline-end]]:[&>input]:pr-1.5 has-[>[data-align=inline-start]]:[&>input]:pl-1.5", + className, + )} + {...props} + /> + ); +} + +const inputGroupAddonVariants = cva({ + base: "flex h-auto cursor-text items-center justify-center gap-2 py-1.5 text-sm font-medium text-muted-foreground select-none group-data-[disabled=true]/input-group:opacity-50 [&>kbd]:rounded-[calc(var(--radius)-5px)] [&>svg:not([class*='size-'])]:size-4", + variants: { + align: { + "inline-start": "order-first pl-2 has-[>button]:-ml-1 has-[>kbd]:ml-[-0.15rem]", + "inline-end": "order-last pr-2 has-[>button]:-mr-1 has-[>kbd]:mr-[-0.15rem]", + "block-start": + "order-first w-full justify-start px-2.5 pt-2 group-has-[>input]/input-group:pt-2 [.border-b]:pb-2", + "block-end": "order-last w-full justify-start px-2.5 pb-2 group-has-[>input]/input-group:pb-2 [.border-t]:pt-2", + }, + }, + defaultVariants: { + align: "inline-start", + }, +}); + +function InputGroupAddon({ + className, + align = "inline-start", + ...props +}: React.ComponentProps<"div"> & VariantProps) { + return ( +
{ + if ((e.target as HTMLElement).closest("button")) { + return; + } + e.currentTarget.parentElement?.querySelector("input")?.focus(); + }} + {...props} + /> + ); +} + +const inputGroupButtonVariants = cva({ + base: "flex items-center gap-2 text-sm shadow-none", + variants: { + size: { + xs: "h-6 gap-1 rounded-[calc(var(--radius)-5px)] px-1.5 [&>svg:not([class*='size-'])]:size-3.5", + sm: "", + "icon-xs": "size-6 rounded-[calc(var(--radius)-5px)] p-0 has-[>svg]:p-0", + "icon-sm": "size-8 p-0 has-[>svg]:p-0", + }, + }, + defaultVariants: { + size: "xs", + }, +}); + +const InputGroupButton = React.forwardRef< + React.ComponentRef, + Omit, "size" | "type"> & + VariantProps & { + type?: "button" | "submit" | "reset"; + } +>(({ className, type = "button", variant = "ghost", size = "xs", ...props }, ref) => { + return ( +