Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_anthropic_output_format_remaining_keywords

This commit is contained in:
mateo-berri 2026-07-22 19:59:09 -07:00
commit 16dad256d4
57 changed files with 6734 additions and 1306 deletions

View file

@ -19,7 +19,7 @@ concurrency:
jobs:
ui-unit-tests:
runs-on: ubuntu-latest
runs-on: ubuntu-latest-16-cores
timeout-minutes: 20
defaults:
run:
@ -50,8 +50,8 @@ jobs:
if [ -n "$BASE_SHA" ]; then
echo "Pull request: running only tests related to changes since $BASE_SHA"
npm run test -- --run --changed "$BASE_SHA" --passWithNoTests \
--pool forks --poolOptions.forks.maxForks=4
--pool forks --poolOptions.forks.maxForks=14
else
echo "Push to $GITHUB_REF_NAME: running the full suite"
npm run test -- --run --pool forks --poolOptions.forks.maxForks=4
npm run test -- --run --pool forks --poolOptions.forks.maxForks=14
fi

View file

@ -9,9 +9,22 @@ duration_in_seconds is used in diff parts of the code base, example
import re
import time as time_module
from datetime import datetime, time, timedelta, timezone, tzinfo
from typing import Optional, Tuple
from typing import Final, Optional, Tuple
from zoneinfo import ZoneInfo
from litellm._logging import verbose_logger
_BUDGET_DURATION_WORD_ALIASES: Final[dict[str, str]] = {
"hourly": "1h",
"daily": "24h",
"weekly": "7d",
"monthly": "30d",
}
def _normalize_duration(duration: str) -> str:
return _BUDGET_DURATION_WORD_ALIASES.get(duration.strip().lower(), duration)
def _extract_from_regex(duration: str) -> Tuple[int, str]:
match = re.match(r"(\d+)(mo|[smhdw]?)", duration)
@ -48,7 +61,7 @@ def duration_in_seconds(duration: str) -> int:
Returns time in seconds till when budget needs to be reset
"""
value, unit = _extract_from_regex(duration=duration)
value, unit = _extract_from_regex(duration=_normalize_duration(duration))
if unit == "s":
return value
@ -124,9 +137,13 @@ def get_next_standardized_reset_time(
current_time, _ = _setup_timezone(current_time, timezone_str)
# Parse duration
value, unit = _parse_duration(duration)
value, unit = _parse_duration(_normalize_duration(duration))
if value is None:
# Fall back to default if format is invalid
verbose_logger.warning(
"Unrecognized budget_duration %r; falling back to a next-midnight reset. "
"Use the <int><unit> format (e.g. '1h', '7d', '30d', '1mo').",
duration,
)
return current_time.replace(hour=0, minute=0, second=0, microsecond=0) + timedelta(days=1)
# Midnight of the current day in the specified timezone

View file

@ -3,6 +3,7 @@ import html as _html
import json
import secrets
import time
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
@ -13,6 +14,7 @@ from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Resp
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -20,7 +22,9 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
TokenEndpointAuthConfigError,
build_token_endpoint_client_auth,
normalize_token_endpoint_auth_method,
)
from litellm.types.mcp_server.mcp_server_manager import MCPTokenEndpointAuthMethod
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_bridge_mint_error_response,
_BridgeMintReady,
@ -111,6 +115,9 @@ def encode_state_with_base_url(
client_redirect_uri: Optional[str] = None,
litellm_user_id: str | None = None,
mcp_server_id: str | None = None,
dcr_client_id: str | None = None,
dcr_client_secret: str | None = None,
dcr_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
) -> str:
"""
Encode the base_url, original state, and PKCE parameters using encryption.
@ -124,8 +131,18 @@ def encode_state_with_base_url(
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
mcp_server_id: The server the flow targets, sealed alongside litellm_user_id (bridge) or
dcr_client_id (ephemeral mint) so the gateway code cannot be replayed against another
server
dcr_client_id: The ephemeral DCR client the gateway minted at authorize for a
client-forwarded-token server with no caller-supplied client; the callback seals it
into the forwarded authorization code so the token exchange can authenticate with it
while the gateway stores nothing
dcr_client_secret: The minted client's secret, when the upstream issued one
dcr_token_endpoint_auth_method: The token-endpoint auth method the upstream's registration
response granted the minted client, sealed alongside the credentials so the exchange
authenticates the way the upstream expects instead of falling back to the server row's
configured method
Returns:
An encrypted string that encodes all values
@ -138,6 +155,9 @@ def encode_state_with_base_url(
"client_redirect_uri": client_redirect_uri,
"litellm_user_id": litellm_user_id,
"mcp_server_id": mcp_server_id,
"dcr_client_id": dcr_client_id,
"dcr_client_secret": dcr_client_secret,
"dcr_token_endpoint_auth_method": dcr_token_endpoint_auth_method,
}
state_json = json.dumps(state_data, sort_keys=True)
encrypted_state = encrypt_value_helper(state_json)
@ -217,6 +237,93 @@ def open_bridge_authorization_code(code: str) -> _BridgeAuthorizationCode | None
return None
_PASSTHROUGH_AUTH_CODE_PREFIX = "llm_ptcode_"
class PassthroughAuthorizationCode(BaseModel):
"""The ephemeral DCR client and upstream code the gateway seals into the authorization code it
forwards for a client-forwarded-token server (``true_passthrough`` / ``oauth_delegate``) whose
authorize fell through to gateway-side registration. These modes forbid the gateway from storing
an OAuth client identity, so the minted client survives only inside this sealed value: the
client echoes it back at the token endpoint, where the gateway recovers the client to
authenticate the upstream exchange. ``mcp_server_id`` binds the code to the server it was minted
for so it cannot be spent at another server's token endpoint."""
model_config = ConfigDict(frozen=True)
upstream_code: str = Field(min_length=1)
client_id: str = Field(min_length=1)
client_secret: str | None = None
token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None
mcp_server_id: str = Field(min_length=1)
def seal_passthrough_authorization_code(
upstream_code: str,
client_id: str,
client_secret: str | None,
mcp_server_id: str,
token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
) -> str:
"""Seal the upstream authorization code together with the ephemeral DCR client that authorized
it. Encrypted with the same authenticated symmetric helper as the OAuth state and bridge codes,
so the client can neither read the (possibly confidential) client credentials nor forge a
code."""
payload = json.dumps(
{
"upstream_code": upstream_code,
"client_id": client_id,
"client_secret": client_secret,
"token_endpoint_auth_method": token_endpoint_auth_method,
"mcp_server_id": mcp_server_id,
},
sort_keys=True,
)
return _PASSTHROUGH_AUTH_CODE_PREFIX + encrypt_value_helper(payload)
def open_passthrough_authorization_code(code: str) -> PassthroughAuthorizationCode | None:
"""Recover the sealed ephemeral client and upstream code, or ``None`` when ``code`` is not a
gateway passthrough code or does not decrypt / validate, so a raw upstream code falls through to
the existing caller-supplied-client behavior."""
if not code.startswith(_PASSTHROUGH_AUTH_CODE_PREFIX):
return None
decrypted = decrypt_value_helper(
code[len(_PASSTHROUGH_AUTH_CODE_PREFIX) :], "passthrough_authorization_code", return_original_value=False
)
if not isinstance(decrypted, str):
return None
try:
return PassthroughAuthorizationCode.model_validate_json(decrypted)
except ValidationError:
return None
def redeem_passthrough_authorization_code(
code: str | None, mcp_server: MCPServer, code_verifier: str | None
) -> PassthroughAuthorizationCode | None:
"""The single redemption gate for sealed passthrough codes: a raw or foreign code returns
``None`` so the caller keeps its existing behavior, while a genuine sealed code must be spent
at the server it was minted for and must carry the PKCE verifier of the S256 flow that minted
it (the mint refuses downgraded flows, so a verifier-less redemption is an interception
attempt, not a legitimate client)."""
if not code:
return None
sealed = open_passthrough_authorization_code(code)
if sealed is None:
return None
if sealed.mcp_server_id != mcp_server.server_id:
raise HTTPException(
status_code=400,
detail="Authorization code was issued for a different MCP server",
)
if not code_verifier:
raise HTTPException(
status_code=400,
detail="code_verifier is required to redeem this authorization code",
)
return sealed
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,
@ -594,6 +701,7 @@ async def authorize_with_server(
code_challenge_method: Optional[str] = None,
response_type: Optional[str] = None,
scope: Optional[str] = None,
ephemeral_dcr_client: "EphemeralDcrClient | None" = None,
):
_raise_if_not_oauth2(mcp_server)
if mcp_server.authorization_url is None:
@ -612,7 +720,10 @@ async def authorize_with_server(
# calling this for its enforcement side effect, then falls through to the gateway
# /callback flow below, which reads the original code_challenge names.
bridge_challenge, bridge_method = _require_s256_pkce(code_challenge, code_challenge_method)
if _dcr_bridge_relays_client_registration(mcp_server):
# A gateway-minted ephemeral client is registered against {base}/callback, so its
# flow must run the short-circuit arm; the relay arm is only for clients that
# registered themselves through the front door and hold their own redirect binding.
if _dcr_bridge_relays_client_registration(mcp_server) and ephemeral_dcr_client is None:
return _redirect_to_upstream_authorize(
mcp_server=mcp_server,
client_id=client_id,
@ -656,7 +767,12 @@ async def authorize_with_server(
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,
mcp_server_id=mcp_server.server_id if (litellm_user_id or ephemeral_dcr_client) else None,
dcr_client_id=ephemeral_dcr_client.client_id if ephemeral_dcr_client else None,
dcr_client_secret=ephemeral_dcr_client.client_secret if ephemeral_dcr_client else None,
dcr_token_endpoint_auth_method=ephemeral_dcr_client.token_endpoint_auth_method
if ephemeral_dcr_client
else None,
)
relay_state = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES)
@ -703,6 +819,7 @@ async def exchange_token_with_server(
code_verifier: Optional[str],
refresh_token: Optional[str] = None,
scope: Optional[str] = None,
client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None,
):
_raise_if_not_oauth2(mcp_server)
if grant_type not in ("authorization_code", "refresh_token"):
@ -718,15 +835,24 @@ async def exchange_token_with_server(
),
)
# The id and secret must come from the same source. When the server-side client_id wins,
# falling back to the caller's secret pairs the persisted client with a foreign secret; the
# register short-circuit hands clients a placeholder secret ("dummy"), so a re-auth against a
# persisted public PKCE client (no stored secret) would send that placeholder and the IdP 401s.
# The id, secret, and token-endpoint auth method must come from the same source. When the
# server-side client_id wins, falling back to the caller's secret pairs the persisted client
# with a foreign secret; the register short-circuit hands clients a placeholder secret
# ("dummy"), so a re-auth against a persisted public PKCE client (no stored secret) would send
# that placeholder and the IdP 401s. Symmetrically, a caller-side client (an ephemeral mint
# recovered from a sealed code) must authenticate the way its own registration was granted,
# not the way the server row is configured; callers that carry no method keep the row's method
# as before.
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
resolved_client_secret = mcp_server.client_secret if mcp_server.client_id else client_secret
resolved_auth_method = (
mcp_server.token_endpoint_auth_method
if mcp_server.client_id
else (client_token_endpoint_auth_method or mcp_server.token_endpoint_auth_method)
)
try:
client_auth = build_token_endpoint_client_auth(
auth_method=mcp_server.token_endpoint_auth_method,
auth_method=resolved_auth_method,
client_id=resolved_client_id,
client_secret=resolved_client_secret,
)
@ -1229,7 +1355,7 @@ async def _persist_dcr_client_registration(
return "failed"
def _client_supplied_redirect_uris(value: object) -> list[str] | None:
def client_supplied_redirect_uris(value: object) -> list[str] | None:
"""RFC 7591 redirect_uris must be a non-empty array of URI strings. Any other shape (not a list,
an empty list, or a list holding a non-string or empty-string element) yields None so every
register arm falls back to the gateway callback instead of echoing a malformed value back to the
@ -1241,6 +1367,142 @@ def _client_supplied_redirect_uris(value: object) -> list[str] | None:
return uris if len(uris) == len(value) else None
async def _post_dcr_registration(
registration_url: str,
register_data: Mapping[str, object],
server_id: str,
) -> httpx.Response:
"""POST an RFC 7591 registration to the upstream and return its response, relaying a classified
upstream rejection instead of a generic 500 and failing loud on an absent response."""
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Register)
try:
response = await async_client.post(
registration_url,
headers=headers,
json=register_data,
)
if response is not None:
response.raise_for_status()
except httpx.HTTPStatusError as exc:
status_code, detail = dcr_fault_detail(classify_upstream_dcr_rejection(exc.response, log_context=server_id))
raise HTTPException(status_code=status_code, detail=detail) from exc
if response is None:
raise HTTPException(
status_code=502,
detail="MCP upstream registration endpoint returned no response",
)
return response
class EphemeralDcrClient(BaseModel):
"""A DCR client minted for a single authorize round trip and never stored by the gateway."""
model_config = ConfigDict(frozen=True)
client_id: str = Field(min_length=1)
client_secret: str | None = None
token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None
_EPHEMERAL_DCR_CLIENT_CACHE = InMemoryCache(default_ttl=_OAUTH_STATE_COOKIE_TTL_SECONDS)
_EPHEMERAL_DCR_MINT_LOCKS: dict[str, asyncio.Lock] = {}
async def mint_ephemeral_dcr_client(request: Request, mcp_server: MCPServer) -> EphemeralDcrClient | None:
"""Mint a throwaway OAuth client via the upstream's RFC 7591 registration endpoint for a
client-forwarded-token server whose authorize arrived with no client_id. Returns ``None`` when
the upstream exposes no registration endpoint, so the caller keeps its existing failure path.
The minted client is deliberately not persisted anywhere: ``true_passthrough`` /
``oauth_delegate`` require the gateway to hold no OAuth client identity, so it survives only in
the encrypted OAuth state and the sealed authorization code the callback forwards.
Reloading the authorize page or retrying a flow must not register a fresh upstream client every
time (an OAuth client identifies the application, not the user, so reuse is semantically
correct). A per-process TTL cache bounded to the OAuth state cookie's lifetime dedupes the mint
per (server, gateway origin), and a per-server lock single-flights concurrent mints (the
``_OAUTH_METADATA_FETCH_LOCKS`` pattern; keyed by server_id alone so the lock registry stays
bounded by the server count even when the request origin varies) so parallel authorize requests
cannot each register an upstream client; the cache stamps nothing onto the server record and
correctness never depends on it because the sealed state carries the client through the flow."""
if mcp_server.registration_url is None:
return None
request_base_url = get_request_base_url(request)
cache_key = f"mcp_ephemeral_dcr_client:{mcp_server.server_id}:{request_base_url}"
cached = _EPHEMERAL_DCR_CLIENT_CACHE.get_cache(cache_key)
if isinstance(cached, EphemeralDcrClient):
return cached
lock = _EPHEMERAL_DCR_MINT_LOCKS.setdefault(mcp_server.server_id, asyncio.Lock())
async with lock:
cached_after_wait = _EPHEMERAL_DCR_CLIENT_CACHE.get_cache(cache_key)
if isinstance(cached_after_wait, EphemeralDcrClient):
return cached_after_wait
register_data: dict[str, object] = {
"client_name": mcp_server.server_name or mcp_server.server_id,
"redirect_uris": [f"{request_base_url}/callback"],
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
"token_endpoint_auth_method": "none",
}
response = await _post_dcr_registration(
registration_url=mcp_server.registration_url,
register_data=register_data,
server_id=mcp_server.server_id,
)
try:
registration = _DcrClientRegistration.model_validate_json(response.text)
except ValidationError as exc:
raise HTTPException(
status_code=502,
detail="MCP upstream registration endpoint returned no usable client_id",
) from exc
if not registration.client_id:
raise HTTPException(
status_code=502,
detail="MCP upstream registration endpoint returned no usable client_id",
)
minted = EphemeralDcrClient(
client_id=registration.client_id,
client_secret=registration.client_secret,
token_endpoint_auth_method=normalize_token_endpoint_auth_method(registration.token_endpoint_auth_method),
)
_EPHEMERAL_DCR_CLIENT_CACHE.set_cache(cache_key, minted)
return minted
async def resolve_ephemeral_dcr_client(
request: Request,
mcp_server: MCPServer,
code_challenge: str | None,
code_challenge_method: str | None,
redirect_uri: str,
) -> EphemeralDcrClient | None:
"""The single owner of the gateway-side mint policy for a clientless authorize. Returns
``None`` for servers whose mode does not permit gateway minting and for upstreams without a
registration endpoint, so those callers keep their existing failure paths: plain ``oauth2``
keeps its persisted-client contract, and the interactive ``oauth_delegate`` dcr_bridge
sign-in has its own sealed-identity flow. ``true_passthrough`` mints regardless of the
``dcr_bridge`` flag (the UI creates passthrough servers with the flag on by default): a
minted flow runs the bridge short-circuit arm, while the relay front door remains for
external clients that registered themselves. Flows that could never succeed fail loud
before any upstream registration: a missing ``authorization_url``, a downgraded PKCE pair
(without S256 the sealed code would be bearer-redeemable by any authenticated caller who
intercepts the redirect), or an untrusted ``redirect_uri`` (a rejected redirect must not be
usable to generate orphan IdP clients)."""
if not (mcp_server.is_true_passthrough or (mcp_server.is_oauth_delegate and not mcp_server.is_dcr_bridge)):
return None
if mcp_server.authorization_url is None:
raise HTTPException(
status_code=400,
detail="MCP server authorization url is not set",
)
_require_s256_pkce(code_challenge, code_challenge_method)
validate_trusted_redirect_uri(request, redirect_uri)
return await mint_ephemeral_dcr_client(request, mcp_server)
async def register_client_with_server(
request: Request,
mcp_server: MCPServer,
@ -1302,30 +1564,11 @@ async def register_client_with_server(
"response_types": response_types or (["code"] if bridge_relay else []),
"token_endpoint_auth_method": token_endpoint_auth_method or ("none" if bridge_relay else ""),
}
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Register)
try:
response = await async_client.post(
mcp_server.registration_url,
headers=headers,
json=register_data,
)
if response is not None:
response.raise_for_status()
except httpx.HTTPStatusError as exc:
status_code, detail = dcr_fault_detail(
classify_upstream_dcr_rejection(exc.response, log_context=mcp_server.server_id)
)
raise HTTPException(status_code=status_code, detail=detail) from exc
if response is None:
raise HTTPException(
status_code=502,
detail="MCP upstream registration endpoint returned no response",
)
response = await _post_dcr_registration(
registration_url=mcp_server.registration_url,
register_data=register_data,
server_id=mcp_server.server_id,
)
token_response = response.json()
@ -1563,11 +1806,23 @@ async def callback(
# 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")
dcr_client_id = state_data.get("dcr_client_id")
dcr_client_secret = state_data.get("dcr_client_secret")
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
)
elif isinstance(dcr_client_id, str) and dcr_client_id and isinstance(mcp_server_id, str) and mcp_server_id:
forwarded_code = seal_passthrough_authorization_code(
upstream_code=code,
client_id=dcr_client_id,
client_secret=dcr_client_secret if isinstance(dcr_client_secret, str) and dcr_client_secret else None,
mcp_server_id=mcp_server_id,
token_endpoint_auth_method=normalize_token_endpoint_auth_method(
state_data.get("dcr_token_endpoint_auth_method")
),
)
params = {"code": forwarded_code, "state": original_state}
complete_returned_url = _append_query_params(redirect_uri, params)
@ -2158,7 +2413,7 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
request_data = await _read_request_body(request=request)
data: dict = {**request_data}
client_redirect_uris = _client_supplied_redirect_uris(data.get("redirect_uris"))
client_redirect_uris = client_supplied_redirect_uris(data.get("redirect_uris"))
dummy_return = {
"client_id": mcp_server_name or "dummy_client",

View file

@ -1856,6 +1856,17 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
default_team_member_models: Optional[List[str]] = None # default allowed_models seeded onto new team members
class PatchTeamRequest(UpdateTeamRequest):
"""
Body of PATCH /team/{team_id}.
Identical to UpdateTeamRequest except team_id is optional, because PATCH takes it
from the path. A team_id in the body is still accepted when it matches the path.
"""
team_id: str | None = None
class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase):
"""
internal type used to reset the budget on a team

View file

@ -136,9 +136,12 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_raise_if_not_oauth2,
authorize_with_server,
client_supplied_redirect_uris,
exchange_token_with_server,
get_request_base_url,
redeem_passthrough_authorization_code,
register_client_with_server,
resolve_ephemeral_dcr_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
@ -1661,7 +1664,21 @@ if MCP_AVAILABLE:
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
_raise_if_not_oauth2(mcp_server)
# Use the server's stored client_id when the caller doesn't supply one
resolved_client_id = mcp_server.client_id or client_id or ""
stored_or_supplied_client_id = mcp_server.client_id or client_id or ""
ephemeral_dcr_client = (
await resolve_ephemeral_dcr_client(
request=request,
mcp_server=mcp_server,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
redirect_uri=redirect_uri,
)
if not stored_or_supplied_client_id
else None
)
resolved_client_id = stored_or_supplied_client_id or (
ephemeral_dcr_client.client_id if ephemeral_dcr_client else ""
)
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -1683,6 +1700,7 @@ if MCP_AVAILABLE:
code_challenge_method=code_challenge_method,
response_type=response_type,
scope=scope,
ephemeral_dcr_client=ephemeral_dcr_client,
)
@router.post(
@ -1705,7 +1723,21 @@ if MCP_AVAILABLE:
):
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
_raise_if_not_oauth2(mcp_server)
resolved_client_id = mcp_server.client_id or client_id or ""
# Sealed passthrough codes exist only for the authorization_code grant. A refresh_token
# grant must never open one: the minted client is unrecoverable after the single flow by
# contract, so an expired browser-held token re-runs authorize instead.
sealed_code = (
redeem_passthrough_authorization_code(code=code, mcp_server=mcp_server, code_verifier=code_verifier)
if grant_type == "authorization_code"
else None
)
resolved_code = sealed_code.upstream_code if sealed_code else code
# A sealed flow ran the gateway /callback as its upstream redirect (bridge short-circuit
# or plain flow alike), so the exchange must present that binding, not the browser page.
resolved_redirect_uri = f"{get_request_base_url(request)}/callback" if sealed_code else redirect_uri
caller_client_id = sealed_code.client_id if sealed_code else client_id
caller_client_secret = sealed_code.client_secret if sealed_code else client_secret
resolved_client_id = mcp_server.client_id or caller_client_id or ""
if not resolved_client_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@ -1721,13 +1753,14 @@ if MCP_AVAILABLE:
request=request,
mcp_server=mcp_server,
grant_type=grant_type,
code=code,
redirect_uri=redirect_uri,
code=resolved_code,
redirect_uri=resolved_redirect_uri,
client_id=resolved_client_id,
client_secret=client_secret,
client_secret=caller_client_secret,
code_verifier=code_verifier,
refresh_token=refresh_token,
scope=scope,
client_token_endpoint_auth_method=sealed_code.token_endpoint_auth_method if sealed_code else None,
)
@router.post(
@ -1743,6 +1776,7 @@ if MCP_AVAILABLE:
mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request)
request_data = await _read_request_body(request=request)
data: dict = {**request_data}
client_redirect_uris = client_supplied_redirect_uris(data.get("redirect_uris"))
return await register_client_with_server(
request=request,
@ -1753,6 +1787,7 @@ if MCP_AVAILABLE:
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=server_id,
persist_credentials=_user_is_full_admin(user_api_key_dict),
client_redirect_uris=client_redirect_uris,
)
@router.delete(

View file

@ -47,6 +47,7 @@ from litellm.proxy._types import (
Member,
NewTeamRequest,
OrgMember,
PatchTeamRequest,
ProxyErrorTypes,
ProxyException,
SpecialManagementEndpointEnums,
@ -1956,6 +1957,7 @@ async def update_team(
)
async def patch_team(
team_id: str,
data: PatchTeamRequest,
http_request: Request,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
litellm_changed_by: Annotated[
@ -1968,11 +1970,12 @@ async def patch_team(
"""
Partially update a team using RFC 7386 JSON Merge Patch semantics.
`team_id` is taken from the path. `metadata` is merged with the team's stored
metadata rather than replacing it: an omitted key is preserved, `key: null`
deletes it, and any other value overwrites (recursing into nested objects).
Every other field behaves exactly like `POST /team/update` (omitted preserves,
a value overwrites). Returns the full updated team.
`team_id` is taken from the path; a `team_id` in the body is accepted only when it
matches. `metadata` is merged with the team's stored metadata rather than replacing
it: an omitted key is preserved, `key: null` deletes it, and any other value
overwrites (recursing into nested objects). Every other field behaves exactly like
`POST /team/update` (omitted preserves, a value overwrites). Returns the full
updated team.
```
curl --location --request PATCH 'http://0.0.0.0:4000/team/8d916b1c-510d-4894-a334-1c16a93344f5' \
@ -1992,21 +1995,15 @@ async def patch_team(
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
try:
body = await http_request.json()
except (json.JSONDecodeError, ValueError):
raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"})
if not isinstance(body, dict):
raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"})
body_team_id = body.pop("team_id", None)
if body_team_id is not None and body_team_id != team_id:
if data.team_id is not None and data.team_id != team_id:
raise HTTPException(
status_code=400,
detail={"error": f"team_id in body ({body_team_id}) does not match team_id in path ({team_id})"},
detail={"error": f"team_id in body ({data.team_id}) does not match team_id in path ({team_id})"},
)
if "metadata" in body:
patch_fields = data.model_dump(exclude_unset=True, exclude={"team_id"})
if "metadata" in patch_fields:
existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id})
if existing_team_row is None:
raise HTTPException(
@ -2014,9 +2011,9 @@ async def patch_team(
detail={"error": f"Team not found, passed team_id={team_id}"},
)
existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {}
body["metadata"] = apply_json_merge_patch(existing_metadata, body["metadata"])
patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"])
update_request = UpdateTeamRequest(team_id=team_id, **body)
update_request = UpdateTeamRequest(team_id=team_id, **patch_fields)
result = await update_team(
data=update_request,

View file

@ -198,13 +198,9 @@
"icon_url": "https://cdn.simpleicons.org/googledrive",
"category": "Productivity",
"registry_url": null,
"transport": "stdio",
"command": "npx",
"args": ["-y", "@modelcontextprotocol/server-gdrive"],
"env_vars": [
{"name": "GOOGLE_CLIENT_ID", "description": "Google OAuth Client ID", "secret": false},
{"name": "GOOGLE_CLIENT_SECRET", "description": "Google OAuth Client Secret", "secret": true}
]
"transport": "http",
"url": "https://drivemcp.googleapis.com/mcp/v1",
"env_vars": []
},
{
"name": "google_calendar",

View file

@ -494,7 +494,14 @@ async def aresponses(
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
input = cast(Union[str, ResponseInputParam], merged_input)
input = cast(
Union[str, ResponseInputParam],
ResponsesAPIRequestUtils.merge_prompt_management_input(
original_input=input,
client_input=client_input,
merged_input=merged_input,
),
)
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
kwargs.pop("prompt_id", None)
@ -609,7 +616,14 @@ def _apply_prompt_management_to_responses_call(
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
input = cast(Union[str, ResponseInputParam], merged_input)
input = cast(
Union[str, ResponseInputParam],
ResponsesAPIRequestUtils.merge_prompt_management_input(
original_input=input,
client_input=client_input,
merged_input=merged_input,
),
)
local_vars["input"] = input
local_vars["model"] = model
if model != original_model:

View file

@ -19,7 +19,9 @@ import litellm
from litellm._logging import verbose_logger
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.types.llms.openai import (
AllMessageValues,
ResponseAPIUsage,
ResponseInputParam,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponseText,
@ -36,6 +38,57 @@ from litellm.types.utils import (
class ResponsesAPIRequestUtils:
"""Helper utils for constructing ResponseAPI requests"""
@staticmethod
def merge_prompt_management_input(
original_input: str | ResponseInputParam,
client_input: list[AllMessageValues],
merged_input: list[AllMessageValues],
) -> list[object]:
if isinstance(original_input, str):
return [*merged_input]
original_items = tuple(original_input)
client_item_ids = frozenset(id(item) for item in client_input)
message_positions = tuple(index for index, item in enumerate(original_items) if id(item) in client_item_ids)
if len(message_positions) == len(original_items):
return [*merged_input]
if not message_positions:
verbose_logger.warning(
"Prompt management hook returned messages without Responses API input messages; merged messages were ignored"
)
return [*original_items]
corresponding_messages = len(client_input) == len(merged_input) and all(
original.get("role") == merged.get("role")
and (not isinstance(original.get("id"), str) or original.get("id") == merged.get("id"))
for original, merged in zip(client_input, merged_input)
)
if corresponding_messages:
merged_by_position = dict(zip(message_positions, merged_input))
return [
merged_by_position[index] if index in merged_by_position else item
for index, item in enumerate(original_items)
]
all_messages_preserved = all(any(original is merged for merged in merged_input) for original in client_input)
if all_messages_preserved:
prefixes = {
id(original_items[position]): original_items[
message_positions[index - 1] + 1 if index else 0 : position
]
for index, position in enumerate(message_positions)
}
trailing_items = original_items[message_positions[-1] + 1 :]
return [item for merged in merged_input for item in (*prefixes.get(id(merged), ()), merged)] + list(
trailing_items
)
verbose_logger.warning(
"Prompt management hook replaced Responses API messages; non-message input items were dropped"
)
return [*merged_input]
@staticmethod
def _check_valid_arg(
supported_params: Optional[List[str]],

View file

@ -5,9 +5,9 @@ answers or when credentials/env are missing; they never skip. Pure unit coverage
of the harness itself carries no `e2e` marker and runs regardless of whether a
proxy is up.
Lifecycle: the `resources` fixture maps the init -> run -> teardown contract
(lifecycle.E2ECase) onto pytest - setup is init(), the test body is run(), and
teardown deletes every resource the test created on the long-lived proxy.
Lifecycle: the `resources` fixture hands each test a lifecycle.ResourceManager -
the test registers a cleanup for every resource it creates, and the fixture's
teardown deletes them all on the long-lived proxy, even when the test fails.
Each suite provides its own `client` fixture (a lifecycle.ResourceClient); these
shared fixtures build on it.

View file

@ -1,13 +1,11 @@
"""Lifecycle contract and resource cleanup for stateful e2e tests.
"""Resource cleanup for stateful e2e tests.
Shared by every e2e suite under tests/e2e/. The proxy under test is
long-lived and never reset between tests, so anything a test creates (keys,
customers, teams, orgs, users, guardrails, budgets, ...) persists unless
explicitly deleted. Every check follows an init -> run -> teardown lifecycle;
teardown releases each resource init() created, even when run() raises.
In pytest terms (see conftest.py): the `resources` fixture's setup is init(),
the test body is run(), and the fixture's teardown is teardown().
explicitly deleted. The `resources` fixture (see conftest.py) hands each test a
ResourceManager; the test registers a cleanup for every resource it creates, and
the fixture's teardown releases them all even when the test body raises.
"""
from dataclasses import dataclass, field
@ -17,38 +15,6 @@ from proxy_client import ProxyClient
from models import KeyGenerateBody
@runtime_checkable
class E2ECase(Protocol):
"""A stateful e2e check run against a long-lived proxy.
init() acquires resources, run() exercises behaviour and asserts, teardown()
releases everything init() created. teardown() must run even if init() fails
partway or run() raises.
"""
def init(self) -> None: ...
def run(self) -> None: ...
def teardown(self) -> None: ...
def run_case(case: E2ECase) -> None:
"""Drive a case through its lifecycle: init -> run -> teardown.
teardown always runs - even when init() fails partway or run() raises (or
skips) - so resources the case already registered on the long-lived proxy are
released. init() is inside the try because cases register cleanups
progressively (e.g. create team, then user, then key), and a failure after
the first creation must still release what came before.
"""
try:
case.init()
case.run()
finally:
case.teardown()
@runtime_checkable
class ResourceClient(Protocol):
"""Proxy operations the convenience creators use. Resource types without a

View file

@ -1,28 +1,36 @@
"""Live e2e: a tiny max_budget on an entity actually blocks requests.
Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
teardown() deletes everything init() created (always runs, even on failure/skip).
Covers the entities with no prior live coverage - internal user, end-user,
organization, team member - plus key and team. See BUDGET_TEST_COVERAGE_MATRIX.md.
One test per budget level (key, team, internal user, end-user, organization,
team member): put the tiny cap on that level, drive spend until a
`budget_exceeded` block, and where a cap could be confused with a neighbor,
prove isolation with an uncapped control key that must keep serving. The
capped-key sweep proves the key's own max_budget blocks across mint shapes
(personal, team, team-member) with roomy surroundings, so the key-level cap is
provably the blocker no matter who the key was minted to.
A non-budget error fails hard (never a skip); if calls never get blocked, budget
enforcement is broken -> fail.
"""
import time
from dataclasses import dataclass, field
from typing import Callable, List, Type
import pytest
from budget_client import BudgetClient, is_budget_block
from e2e_config import unique_marker
from e2e_http import StreamingResponse, require_successful_call
from lifecycle import run_case
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
TINY_CAP = 3e-6
ROOMY_CAP = 100.0
def _chat(client: BudgetClient, key: str, *, user: str | None = None) -> StreamingResponse:
return client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16, user=user)
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> StreamingResponse:
"""Send paid calls until the entity's budget blocks one; return the blocked
response so callers can assert on its shape. Key/user/org/member block within
@ -30,13 +38,7 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
enforces off table spend that lands on the batch write, so it takes a few
more. A non-budget error fails hard (never a skip)."""
for _ in range(40):
result = client.chat(
key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
user=user or None,
)
result = _chat(client, key, user=user or None)
if is_budget_block(result):
return result
require_successful_call(result)
@ -44,225 +46,154 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
pytest.fail("budget never enforced within the call budget")
@dataclass
class _BudgetCase:
"""Base E2ECase: a key under some budgeted entity must get blocked.
Subclasses set up the budgeted entity in init() and register every created id
in `_undo` (run LIFO in teardown so a key is deleted before its team/org).
"""
client: BudgetClient
key: str = ""
_undo: List[Callable[[], None]] = field(
default_factory=list
) # mutable-ok: per-case teardown registry
def init(self) -> None:
raise NotImplementedError
def run(self) -> None:
_assert_budget_blocks(self.client, self.key)
def teardown(self) -> None:
for undo in reversed(self._undo):
undo()
def _assert_blocked_429(client: BudgetClient, key: str) -> StreamingResponse:
blocked = _assert_budget_blocks(client, key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
return blocked
class KeyBudgetCase(_BudgetCase):
"""A bare key (no team_id / user_id) carrying its own max_budget, so only the
key-level budget can be the thing that blocks. The refusal must be a 429
budget_exceeded; any other error already fails via _assert_budget_blocks."""
class TestBudgetBlocksPerLevel:
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_bare_key_blocks_over_its_own_budget(self, client: BudgetClient, resources: ResourceManager) -> None:
key = client.generate_key(max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(key))
def init(self) -> None:
self.key = self.client.generate_key(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_key(self.key))
_assert_blocked_429(client, key)
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
@pytest.mark.covers("quota_management.budget.team.blocks_over_limit")
def test_team_budget_blocks_every_team_key(self, client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", max_budget=TINY_CAP)
resources.defer(lambda: client.delete_team(team_id))
spender_key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(spender_key))
sibling_key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(sibling_key))
class TeamBudgetCase(_BudgetCase):
"""An admin caps a whole team: two keys under a tiny-budget team, neither with
a key-level budget. Key A is driven until the team cap blocks it; key B's very
first call must then be refused too, proving the cap sits on the team, not the
key that spent. Both refusals must be 429 budget_exceeded."""
def init(self) -> None:
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=3e-6
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
self._sibling_key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self._sibling_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
sibling = self.client.chat(
self._sibling_key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
)
_assert_blocked_429(client, spender_key)
sibling = _chat(client, sibling_key)
assert is_budget_block(sibling) and sibling.status_code == 429, (
f"a sibling key on the capped team must get the same 429 budget_exceeded, "
f"got {sibling.status_code}: {sibling.body[:200]}"
)
@pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit")
def test_user_budget_enforced_across_all_their_keys(
self, client: BudgetClient, resources: ResourceManager
) -> None:
user_id = client.create_user(max_budget=TINY_CAP)
resources.defer(lambda: client.delete_user(user_id))
first_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(first_key))
second_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(second_key))
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}")
resources.defer(lambda: client.delete_team(team_id))
client.add_team_member(team_id, user_id)
team_key = client.generate_key(team_id=team_id, user_id=user_id)
resources.defer(lambda: client.delete_key(team_key))
class InternalUserBudgetCase(_BudgetCase):
"""A user's max_budget follows the person, not the key. The capped user holds
two personal keys (no team, no key budgets) plus a team-member key on an
uncapped team; once the first personal key is refused, the other two must be
refused as well - a second key is not a fresh allowance, and since #32005 the
user budget draws down team keys too. All refusals must be 429 budget_exceeded."""
def init(self) -> None:
user_id = self.client.create_user(max_budget=3e-6)
self._undo.append(lambda: self.client.delete_user(user_id))
self.key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self.key))
self._second_key = self.client.generate_key(user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self._second_key))
team_id = self.client.create_team(alias=f"e2e-budget-team-{unique_marker()}")
self._undo.append(lambda: self.client.delete_team(team_id))
self.client.add_team_member(team_id, user_id)
self._team_key = self.client.generate_key(team_id=team_id, user_id=user_id)
self._undo.append(lambda: self.client.delete_key(self._team_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
for label, key in (("second personal key", self._second_key), ("team-member key", self._team_key)):
result = self.client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16)
_assert_blocked_429(client, first_key)
for label, key in (("second personal key", second_key), ("team-member key", team_key)):
result = _chat(client, key)
assert is_budget_block(result) and result.status_code == 429, (
f"the {label} of a user over budget must get the same 429 budget_exceeded, "
f"got {result.status_code}: {result.body[:200]}"
)
class EndUserBudgetCase(_BudgetCase):
def init(self) -> None:
@pytest.mark.covers("quota_management.budget.end_user.blocks_over_limit")
def test_end_user_budget_blocks_attributed_calls(
self, client: BudgetClient, resources: ResourceManager
) -> None:
customer = f"e2e-budget-cust-{unique_marker()}"
self.client.create_customer(customer, max_budget=3e-6)
self._undo.append(lambda: self.client.delete_customers([customer]))
self.key = self.client.generate_key(models=["claude-haiku-4-5"])
self._undo.append(lambda: self.client.delete_key(self.key))
self._customer = customer
client.create_customer(customer, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_customers([customer]))
key = client.generate_key(models=["claude-haiku-4-5"])
resources.defer(lambda: client.delete_key(key))
def run(self) -> None:
_assert_budget_blocks(self.client, self.key, user=self._customer)
_assert_budget_blocks(client, key, user=customer)
@pytest.mark.covers("quota_management.budget.organization.blocks_over_limit")
def test_org_budget_blocks_keys_under_it(self, client: BudgetClient, resources: ResourceManager) -> None:
org_id = client.create_org(max_budget=TINY_CAP, alias=f"e2e-budget-org-{unique_marker()}")
resources.defer(lambda: client.delete_org(org_id))
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id)
resources.defer(lambda: client.delete_team(team_id))
key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(key))
class OrganizationBudgetCase(_BudgetCase):
"""Org carries the tiny budget; the team under it and the key carry none, so
the org is the only entity that can block (the historically weak link). The
refusal must be a 429 budget_exceeded that names the org as the blocker."""
def init(self) -> None:
self._org_id = self.client.create_org(
max_budget=3e-6, alias=f"e2e-budget-org-{unique_marker()}"
)
self._undo.append(lambda: self.client.delete_org(self._org_id))
team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", organization_id=self._org_id
)
self._undo.append(lambda: self.client.delete_team(team_id))
self.key = self.client.generate_key(team_id=team_id)
self._undo.append(lambda: self.client.delete_key(self.key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
assert f"Organization={self._org_id}" in blocked.body, (
blocked = _assert_blocked_429(client, key)
assert f"Organization={org_id}" in blocked.body, (
f"refusal must name the org as the blocker, got: {blocked.body[:200]}"
)
@pytest.mark.covers("quota_management.budget.team_member.blocks_over_limit")
def test_member_budget_blocks_without_touching_teammates(
self, client: BudgetClient, resources: ResourceManager
) -> None:
team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
member_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(member_id))
client.add_team_member(team_id, member_id, max_budget_in_team=TINY_CAP)
member_key = client.generate_key(team_id=team_id, user_id=member_id)
resources.defer(lambda: client.delete_key(member_key))
teammate_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(teammate_id))
client.add_team_member(team_id, teammate_id)
teammate_key = client.generate_key(team_id=team_id, user_id=teammate_id)
resources.defer(lambda: client.delete_key(teammate_key))
class TeamMemberBudgetCase(_BudgetCase):
"""Member A's per-team budget is tiny while the team and both members' user
budgets are roomy (100.0), so the only cap that can trip is A's: a block
proves member-level enforcement and must be a 429 budget_exceeded. Teammate
B, uncapped on the same team, must keep serving after A is cut off, proving
the member cap does not leak onto the team or its members."""
def init(self) -> None:
self._team_id = self.client.create_team(
alias=f"e2e-budget-team-{unique_marker()}", max_budget=100.0
)
self._undo.append(lambda: self.client.delete_team(self._team_id))
self._member_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(self._member_id))
self.client.add_team_member(self._team_id, self._member_id, max_budget_in_team=3e-6)
self.key = self.client.generate_key(team_id=self._team_id, user_id=self._member_id)
self._undo.append(lambda: self.client.delete_key(self.key))
teammate_id = self.client.create_user(max_budget=100.0)
self._undo.append(lambda: self.client.delete_user(teammate_id))
self.client.add_team_member(self._team_id, teammate_id)
self._teammate_key = self.client.generate_key(team_id=self._team_id, user_id=teammate_id)
self._undo.append(lambda: self.client.delete_key(self._teammate_key))
def run(self) -> None:
blocked = _assert_budget_blocks(self.client, self.key)
assert blocked.status_code == 429, (
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
)
teammate = self.client.chat(
self._teammate_key,
"claude-haiku-4-5",
f"spend {unique_marker()}",
max_tokens=16,
)
require_successful_call(teammate)
_assert_blocked_429(client, member_key)
require_successful_call(_chat(client, teammate_key))
def _case_id(case_cls: Type[_BudgetCase]) -> str:
return case_cls.__name__
class TestKeyBudgetBlocksAcrossKeyKinds:
"""The tiny max_budget sits on the key itself while every budget around it
(user / team / membership) is roomy, so only the key-level cap can block; the
uncapped control key minted to the same surroundings must keep serving after
the capped key is refused, proving nothing around the key was the blocker."""
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_personal_key_blocks_over_its_own_budget(
self, client: BudgetClient, resources: ResourceManager
) -> None:
user_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(user_id))
capped_key = client.generate_key(user_id=user_id, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(capped_key))
control_key = client.generate_key(user_id=user_id)
resources.defer(lambda: client.delete_key(control_key))
@pytest.mark.parametrize(
"case_cls",
[
pytest.param(
KeyBudgetCase,
marks=pytest.mark.covers("quota_management.budget.key.blocks_over_limit"),
),
pytest.param(
TeamBudgetCase,
marks=pytest.mark.covers("quota_management.budget.team.blocks_over_limit"),
),
pytest.param(
InternalUserBudgetCase,
marks=pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit"),
),
pytest.param(
EndUserBudgetCase,
marks=pytest.mark.covers("quota_management.budget.end_user.blocks_over_limit"),
),
pytest.param(
OrganizationBudgetCase,
marks=pytest.mark.covers("quota_management.budget.organization.blocks_over_limit"),
),
pytest.param(
TeamMemberBudgetCase,
marks=pytest.mark.covers("quota_management.budget.team_member.blocks_over_limit"),
),
],
ids=_case_id,
)
def test_budget_enforcement(
client: BudgetClient, case_cls: Type[_BudgetCase]
) -> None:
run_case(case_cls(client))
_assert_blocked_429(client, capped_key)
require_successful_call(_chat(client, control_key))
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_team_key_blocks_over_its_own_budget(self, client: BudgetClient, resources: ResourceManager) -> None:
team_id = client.create_team(alias=f"e2e-key-cap-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
capped_key = client.generate_key(team_id=team_id, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(capped_key))
control_key = client.generate_key(team_id=team_id)
resources.defer(lambda: client.delete_key(control_key))
_assert_blocked_429(client, capped_key)
require_successful_call(_chat(client, control_key))
@pytest.mark.covers("quota_management.budget.key.blocks_over_limit")
def test_team_member_key_blocks_over_its_own_budget(
self, client: BudgetClient, resources: ResourceManager
) -> None:
team_id = client.create_team(alias=f"e2e-key-cap-team-{unique_marker()}", max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_team(team_id))
member_id = client.create_user(max_budget=ROOMY_CAP)
resources.defer(lambda: client.delete_user(member_id))
client.add_team_member(team_id, member_id, max_budget_in_team=ROOMY_CAP)
capped_key = client.generate_key(team_id=team_id, user_id=member_id, max_budget=TINY_CAP)
resources.defer(lambda: client.delete_key(capped_key))
control_key = client.generate_key(team_id=team_id, user_id=member_id)
resources.defer(lambda: client.delete_key(control_key))
_assert_blocked_429(client, capped_key)
require_successful_call(_chat(client, control_key))

View file

@ -23,7 +23,7 @@ import pytest
from e2e_http import Result, Success
from lifecycle import ResourceManager
from models import ChatResponse, SpendLogs, SpendLogsParams
from models import ChatResponse, LiteLLMParamsBody, SpendLogs, SpendLogsParams
from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap
pytestmark = pytest.mark.e2e
@ -232,14 +232,12 @@ def test_cache_hit_is_zero_cost_and_suffixed(
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.cache_hit == "True" for r in rs)
)
cache_rows = [r for r in rows if r.cache_hit == "True"]
if not cache_rows:
pytest.skip(
"no cache-hit row observed; caching may be disabled on this proxy. "
f"rows seen: {_summarize(rows)}"
)
cache_row = cache_rows[0]
cache_row = _require_row(
rows,
lambda r: r.cache_hit == "True",
"with cache_hit=True (caching is enabled on the e2e proxy, so an identical "
"repeat call must hit the cache)",
)
assert (
cache_row.spend or 0
) == 0.0, f"cache hit was charged (double-charge regression): {_summarize(rows)}"
@ -504,22 +502,27 @@ def test_each_model_on_a_shared_key_gets_its_own_row(
@pytest.mark.covers("quota_management.spend_tracking.failure.writes_failure_row")
def test_failure_call_writes_failure_status_row(
client: SpendClient, scoped_key: str
client: SpendClient, resources: ResourceManager, scoped_key: str
) -> None:
result = client.chat(scoped_key, "gemini-2.5-flash", "", max_tokens=1)
if is_ok(result):
pytest.skip("call unexpectedly succeeded; could not induce a failure row")
model = f"e2e-spend-failure-{unique_marker()}"
model_id = client.proxy.create_model(
model,
LiteLLMParamsBody(model="openai/gpt-5.5", api_key="sk-invalid-e2e-failure-row"),
)
resources.defer(lambda: client.proxy.delete_model(model_id))
result = client.chat(scoped_key, model, f"trigger failure {unique_marker()}", max_tokens=1)
assert not is_ok(result), (
f"a call to a deployment with an invalid upstream key must fail, not succeed: {result}"
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.status == "failure" for r in rs)
)
failure_rows = [r for r in rows if r.status == "failure"]
if not failure_rows:
pytest.skip(
"no failure-status row was logged for the rejected call; "
"failure logging is environment-specific"
)
assert (failure_rows[0].spend or 0) == 0.0, "failed call must not be charged"
failure_row = _require_row(
rows, lambda r: r.status == "failure", "with status=failure for the rejected call"
)
assert (failure_row.spend or 0) == 0.0, "failed call must not be charged"
@pytest.mark.covers("quota_management.spend_tracking.spend_calculate.returns_cost")

View file

@ -1,8 +1,13 @@
import unittest
from datetime import datetime, time, timezone
from unittest.mock import patch
from zoneinfo import ZoneInfo
from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time
import litellm.litellm_core_utils.duration_parser as duration_parser
from litellm.litellm_core_utils.duration_parser import (
duration_in_seconds,
get_next_standardized_reset_time,
)
class TestStandardizedResetTime(unittest.TestCase):
@ -316,5 +321,69 @@ class TestResetTimeOfDay(unittest.TestCase):
)
class TestWordFormBudgetDurations(unittest.TestCase):
"""The Admin UI historically persisted word-form budget durations
(hourly/daily/weekly/monthly). They must resolve to their real interval
instead of silently collapsing to a next-midnight (daily) reset.
"""
def test_word_forms_map_to_correct_reset_times(self):
base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc)
self.assertEqual(
get_next_standardized_reset_time("hourly", base_time, "UTC"),
datetime(2023, 5, 17, 16, 0, 0, tzinfo=timezone.utc),
)
self.assertEqual(
get_next_standardized_reset_time("daily", base_time, "UTC"),
datetime(2023, 5, 18, 0, 0, 0, tzinfo=timezone.utc),
)
self.assertEqual(
get_next_standardized_reset_time("weekly", base_time, "UTC"),
datetime(2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc),
)
self.assertEqual(
get_next_standardized_reset_time("monthly", base_time, "UTC"),
datetime(2023, 6, 1, 0, 0, 0, tzinfo=timezone.utc),
)
def test_word_forms_are_not_all_collapsed_to_daily(self):
base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc)
results = {
word: get_next_standardized_reset_time(word, base_time, "UTC")
for word in ("hourly", "daily", "weekly", "monthly")
}
self.assertEqual(len(set(results.values())), len(results))
def test_word_forms_match_canonical_int_unit_forms(self):
base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc)
for word, canonical in (("hourly", "1h"), ("daily", "24h"), ("weekly", "7d"), ("monthly", "30d")):
self.assertEqual(
get_next_standardized_reset_time(word, base_time, "UTC"),
get_next_standardized_reset_time(canonical, base_time, "UTC"),
)
def test_word_forms_are_case_and_whitespace_insensitive(self):
base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc)
self.assertEqual(
get_next_standardized_reset_time(" Monthly ", base_time, "UTC"),
datetime(2023, 6, 1, 0, 0, 0, tzinfo=timezone.utc),
)
def test_duration_in_seconds_accepts_word_forms(self):
self.assertEqual(duration_in_seconds("hourly"), 3600)
self.assertEqual(duration_in_seconds("daily"), 86400)
self.assertEqual(duration_in_seconds("weekly"), 604800)
self.assertEqual(duration_in_seconds("monthly"), 2592000)
def test_invalid_duration_logs_warning_and_falls_back(self):
base_time = datetime(2023, 5, 15, 15, 0, 0, tzinfo=timezone.utc)
with patch.object(duration_parser.verbose_logger, "warning") as mock_warning:
result = get_next_standardized_reset_time("garbage", base_time, "UTC")
self.assertEqual(result, datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc))
mock_warning.assert_called_once()
self.assertIn("garbage", mock_warning.call_args.args)
if __name__ == "__main__":
unittest.main()

View file

@ -70,25 +70,33 @@ class TestAgentCoreAcceptHeader:
"""
End-to-end test: verify Accept header appears in the final HTTP request
when using JWT auth through litellm.completion().
No exception swallowing: if completion() raises (for example because the
injected client was silently ignored and a real network call was made),
the test must fail with that error, not a misleading mock assertion.
"""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime",
messages=[{"role": "user", "content": "test"}],
api_key="test-jwt-token",
client=client,
)
except Exception:
pass
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.json.return_value = {
"result": {"role": "assistant", "content": [{"text": "agent reply"}]}
}
mock_post.assert_called_once()
headers = mock_post.call_args.kwargs["headers"]
assert "Accept" in headers
assert headers["Accept"] == "application/json, text/event-stream"
with patch.object(client, "post", return_value=mock_response) as mock_post:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime",
messages=[{"role": "user", "content": "test"}],
api_key="test-jwt-token",
client=client,
)
mock_post.assert_called_once()
headers = mock_post.call_args.kwargs["headers"]
assert headers["Accept"] == "application/json, text/event-stream"
assert response.choices[0].message.content == "agent reply"
class TestAgentCoreJsonResponseParsing:

View file

@ -1,4 +1,3 @@
import importlib
import json
import os
import sys
@ -16,22 +15,7 @@ MOCK_EMBEDDING_RESPONSE = [[0.1, 0.2, 0.3, 0.4, 0.5]]
@pytest.fixture
def reload_huggingface_modules():
"""
Reload modules to ensure fresh references after conftest reloads litellm.
This ensures the HTTPHandler class being patched is the same one used by
the embedding handler during parallel test execution.
"""
import litellm.llms.custom_httpx.http_handler as http_handler_module
import litellm.llms.huggingface.embedding.handler as hf_embedding_handler_module
importlib.reload(http_handler_module)
importlib.reload(hf_embedding_handler_module)
yield
@pytest.fixture
def mock_embedding_http_handler(reload_huggingface_modules):
def mock_embedding_http_handler():
"""Fixture to mock the HTTP handler for embedding tests"""
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
mock_response = MagicMock()
@ -43,7 +27,7 @@ def mock_embedding_http_handler(reload_huggingface_modules):
@pytest.fixture
def mock_embedding_async_http_handler(reload_huggingface_modules):
def mock_embedding_async_http_handler():
"""Fixture to mock the async HTTP handler for embedding tests"""
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",

View file

@ -3,26 +3,16 @@ Integration tests for Vertex AI rerank functionality.
These tests demonstrate end-to-end usage of the Vertex AI rerank feature.
"""
import importlib
from unittest.mock import MagicMock
import httpx
from litellm.llms.vertex_ai.rerank.transformation import VertexAIRerankConfig
class TestVertexAIRerankIntegration:
def setup_method(self):
# Reload modules to ensure fresh references after conftest reloads litellm.
# This ensures the class being patched is the same one used by the tests.
import litellm.llms.vertex_ai.rerank.transformation as rerank_transformation_module
importlib.reload(rerank_transformation_module)
# Re-import after reload to get the fresh class
from litellm.llms.vertex_ai.rerank.transformation import (
VertexAIRerankConfig as FreshConfig,
)
self.config = FreshConfig()
self.config = VertexAIRerankConfig()
self.model = "semantic-ranker-default@latest"
def test_end_to_end_rerank_flow(self):

View file

@ -8148,3 +8148,489 @@ async def test_register_wall_names_the_fix_for_urlless_servers():
detail_text = str(exc_info.value.detail)
assert "set Authorization URL and Token URL" in detail_text
assert "Issuer" in detail_text
def test_passthrough_authorization_code_round_trips_and_rejects_hostile_input():
"""The passthrough gateway code seals and recovers the ephemeral DCR client and upstream code,
and is total over hostile input: a raw upstream code opens to None, and a tampered or
non-gateway value opens to None rather than raising, so every existing caller-supplied-client
flow is untouched."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
open_passthrough_authorization_code,
seal_passthrough_authorization_code,
)
with patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY):
sealed = seal_passthrough_authorization_code(
upstream_code="up-code",
client_id="minted-77",
client_secret="mint-secret",
mcp_server_id="srv-1",
token_endpoint_auth_method="client_secret_basic",
)
opened = open_passthrough_authorization_code(sealed)
assert opened is not None
assert opened.upstream_code == "up-code"
assert opened.client_id == "minted-77"
assert opened.client_secret == "mint-secret"
assert opened.mcp_server_id == "srv-1"
assert opened.token_endpoint_auth_method == "client_secret_basic"
assert open_passthrough_authorization_code("raw-upstream-code") is None
assert open_passthrough_authorization_code(sealed[:-4] + "aaaa") is None
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_BRIDGE_AUTH_CODE_PREFIX,
_PASSTHROUGH_AUTH_CODE_PREFIX,
open_bridge_authorization_code,
seal_bridge_authorization_code,
)
bridge_sealed = seal_bridge_authorization_code(
upstream_code="up-code", litellm_user_id="sso-user-9", mcp_server_id="srv-1"
)
reprefixed_as_passthrough = _PASSTHROUGH_AUTH_CODE_PREFIX + bridge_sealed[len(_BRIDGE_AUTH_CODE_PREFIX) :]
reprefixed_as_bridge = _BRIDGE_AUTH_CODE_PREFIX + sealed[len(_PASSTHROUGH_AUTH_CODE_PREFIX) :]
assert open_passthrough_authorization_code(reprefixed_as_passthrough) is None
assert open_bridge_authorization_code(reprefixed_as_bridge) is None
@pytest.mark.asyncio
async def test_authorize_with_ephemeral_dcr_client_seals_client_into_state():
"""When mcp_authorize fell through to a gateway-side DCR mint, authorize_with_server seals the
minted client and the target server into the encrypted OAuth state, so the callback can bind
them into the forwarded authorization code while the gateway stores nothing."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
EphemeralDcrClient,
authorize_with_server,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None)
captured: dict = {}
def _capture(**kwargs):
captured.update(kwargs)
return "mocked_encrypted_state"
with 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="minted-77",
redirect_uri="http://127.0.0.1:60108/callback",
state="s",
code_challenge="chal",
code_challenge_method="S256",
ephemeral_dcr_client=EphemeralDcrClient(
client_id="minted-77", client_secret="mint-secret", token_endpoint_auth_method="client_secret_basic"
),
)
assert captured["dcr_client_id"] == "minted-77"
assert captured["dcr_client_secret"] == "mint-secret"
assert captured["dcr_token_endpoint_auth_method"] == "client_secret_basic"
assert captured["mcp_server_id"] == server.server_id
assert "client_id=minted-77" in response.headers["location"]
@pytest.mark.asyncio
async def test_callback_wraps_code_into_passthrough_code_for_ephemeral_dcr_state():
"""When the OAuth state carries an ephemeral DCR client, the callback forwards a sealed
passthrough code (binding the client and the upstream code to the server) instead of the raw
upstream code, so the client's later token call can authenticate the exchange with a client the
gateway never stored."""
from urllib.parse import parse_qs, urlparse
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
callback,
open_passthrough_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",
"mcp_server_id": "srv-1",
"dcr_client_id": "minted-77",
"dcr_client_secret": "mint-secret",
"dcr_token_endpoint_auth_method": "client_secret_basic",
}
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]
opened = open_passthrough_authorization_code(forwarded_code)
assert opened is not None
assert opened.upstream_code == "REAL-UPSTREAM-CODE"
assert opened.client_id == "minted-77"
assert opened.client_secret == "mint-secret"
assert opened.mcp_server_id == "srv-1"
assert opened.token_endpoint_auth_method == "client_secret_basic"
@pytest.mark.asyncio
async def test_authorize_bridge_server_with_ephemeral_client_takes_short_circuit_arm():
"""A gateway-minted client is registered against {base}/callback, so a bridge server's
authorize with an ephemeral client must run the short-circuit (gateway /callback) arm with a
relay state cookie, never the verbatim relay: relaying would send the browser's redirect_uri
to an IdP that has the gateway callback registered, stranding the flow."""
from urllib.parse import parse_qs, urlparse
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
EphemeralDcrClient,
authorize_with_server,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.true_passthrough)
with patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY):
response = await authorize_with_server(
request=_bridge_mock_request(),
mcp_server=server,
client_id="minted-77",
redirect_uri="http://127.0.0.1:60108/callback",
state="client-state",
code_challenge="chal",
code_challenge_method="S256",
ephemeral_dcr_client=EphemeralDcrClient(client_id="minted-77", client_secret=None),
)
location = response.headers["location"]
params = parse_qs(urlparse(location).query)
assert params["redirect_uri"] == ["https://litellm.example.com/callback"]
assert params["client_id"] == ["minted-77"]
assert params["state"] != ["client-state"]
assert any(cookie.startswith("mcp_oauth_state_") for cookie in response.headers.get("set-cookie", "").split(";"))
@pytest.mark.asyncio
async def test_callback_forwards_raw_code_when_dcr_state_lacks_server_binding():
"""A state carrying a dcr client but no server id cannot produce a server-bound sealed code, so
the callback falls back to forwarding the raw upstream code instead of sealing an unbindable
one."""
from urllib.parse import parse_qs, urlparse
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import callback
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",
"dcr_client_id": "minted-77",
}
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 forwarded_code == "REAL-UPSTREAM-CODE"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"auth_type_value",
[
"none",
"api_key",
"bearer_token",
"basic",
"authorization",
"oauth2",
"aws_sigv4",
"token",
"oauth2_token_exchange",
"oauth2_id_jag",
"true_passthrough",
"oauth_delegate",
],
)
@pytest.mark.parametrize("dcr_bridge", [True, False])
async def test_resolve_ephemeral_dcr_client_mint_set_is_exact(auth_type_value, dcr_bridge):
"""The full authorize-time mint decision matrix, one cell per (auth_type, dcr_bridge). The gateway
mints iff true_passthrough (any bridge) or oauth_delegate-and-not-dcr_bridge; every other mode
returns None so no non-OAuth mode ever registers an upstream client, and the interactive
oauth_delegate dcr_bridge sign-in is left to its own browser-front-door flow. The UI
gatewayMintsClientFor helper mirrors this exact set; ui/.../mcp_tools/types.test.tsx pins the
frontend side against the same table, so a divergence fails on one side or the other."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
EphemeralDcrClient,
resolve_ephemeral_dcr_client,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(
auth_type=MCPAuth(auth_type_value),
dcr_bridge=dcr_bridge,
server_id=f"matrix_{auth_type_value}_{dcr_bridge}",
server_name=f"matrix_{auth_type_value}_{dcr_bridge}",
)
expected_mint = server.is_true_passthrough or (server.is_oauth_delegate and not server.is_dcr_bridge)
mint_mock = AsyncMock(return_value=EphemeralDcrClient(client_id="minted", client_secret=None))
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client",
mint_mock,
):
result = await resolve_ephemeral_dcr_client(
request=_bridge_mock_request(),
mcp_server=server,
code_challenge="chal",
code_challenge_method="S256",
redirect_uri="http://127.0.0.1:9/callback",
)
if expected_mint:
mint_mock.assert_awaited_once()
assert result is not None
else:
mint_mock.assert_not_awaited()
assert result is None
@pytest.mark.asyncio
async def test_mint_ephemeral_dcr_client_returns_none_without_registration_endpoint():
"""A server whose upstream exposes no RFC 7591 registration endpoint cannot mint, so the
fall-through reports None and the caller keeps its existing missing_client_id failure."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
mint_ephemeral_dcr_client,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None, registration_url=None)
assert await mint_ephemeral_dcr_client(_bridge_mock_request(), server) is None
@pytest.mark.asyncio
async def test_mint_ephemeral_dcr_client_posts_rfc7591_and_returns_client():
"""The mint POSTs a public-client RFC 7591 registration bound to the gateway /callback and hands
back the upstream's client without persisting it anywhere."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
mint_ephemeral_dcr_client,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(
auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id="mint_posts_srv", server_name="mint_posts_srv"
)
mock_response = MagicMock()
mock_response.text = json.dumps(
{"client_id": "minted-77", "client_secret": "mint-secret", "token_endpoint_auth_method": "client_secret_basic"}
)
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
):
minted = await mint_ephemeral_dcr_client(_bridge_mock_request(), server)
assert minted is not None
assert minted.client_id == "minted-77"
assert minted.client_secret == "mint-secret"
assert minted.token_endpoint_auth_method == "client_secret_basic"
register_data = mock_async_client.post.call_args.kwargs["json"]
assert register_data["redirect_uris"] == ["https://litellm.example.com/callback"]
assert register_data["token_endpoint_auth_method"] == "none"
assert register_data["grant_types"] == ["authorization_code", "refresh_token"]
@pytest.mark.asyncio
async def test_mint_ephemeral_dcr_client_reuses_minted_client_within_flow_ttl():
"""Reloading the authorize page must not spam the upstream registration endpoint with orphan
clients: within the OAuth state's lifetime a second mint for the same server and gateway origin
reuses the cached client and performs no second upstream POST."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
mint_ephemeral_dcr_client,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(
auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id="mint_reuse_srv", server_name="mint_reuse_srv"
)
mock_response = MagicMock()
mock_response.text = json.dumps({"client_id": "minted-77"})
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
):
first = await mint_ephemeral_dcr_client(_bridge_mock_request(), server)
second = await mint_ephemeral_dcr_client(_bridge_mock_request(), server)
assert first is not None
assert second == first
mock_async_client.post.assert_called_once()
@pytest.mark.asyncio
async def test_mint_ephemeral_dcr_client_single_flights_concurrent_mints():
"""Two in-flight authorize requests for the same server must not both register an upstream
client: the per-key lock makes the second waiter reuse the first mint, so exactly one upstream
POST happens."""
import asyncio
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
mint_ephemeral_dcr_client,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(
auth_type=MCPAuth.true_passthrough,
dcr_bridge=None,
server_id="mint_concurrent_srv",
server_name="mint_concurrent_srv",
)
mock_response = MagicMock()
mock_response.text = json.dumps({"client_id": "minted-77"})
mock_response.raise_for_status = MagicMock()
async def _slow_post(*args, **kwargs):
await asyncio.sleep(0.05)
return mock_response
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(side_effect=_slow_post)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
):
first, second = await asyncio.gather(
mint_ephemeral_dcr_client(_bridge_mock_request(), server),
mint_ephemeral_dcr_client(_bridge_mock_request(), server),
)
assert first is not None
assert second == first
mock_async_client.post.assert_called_once()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"payload, server_id",
[
({"unexpected": "shape"}, "mint_bad_shape_srv"),
({"client_id": ""}, "mint_empty_id_srv"),
],
)
async def test_mint_ephemeral_dcr_client_unusable_registration_response_is_502(payload, server_id):
"""An upstream registration response without a usable client_id, whether the field is missing or
an empty string, surfaces as a loud 502 instead of letting the authorize proceed with an empty
client and fail opaquely at the IdP."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
mint_ephemeral_dcr_client,
)
from litellm.types.mcp import MCPAuth
server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id=server_id, server_name=server_id)
mock_response = MagicMock()
mock_response.text = json.dumps(payload)
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
):
with pytest.raises(HTTPException) as exc:
await mint_ephemeral_dcr_client(_bridge_mock_request(), server)
assert exc.value.status_code == 502
@pytest.mark.asyncio
@pytest.mark.parametrize(
"sealed_auth_method, expects_basic_header",
[
("client_secret_basic", True),
(None, False),
],
)
async def test_token_exchange_authenticates_with_the_sealed_clients_own_auth_method(
sealed_auth_method, expects_basic_header
):
"""The id, secret, and token-endpoint auth method must come from the same source: a client
recovered from a sealed passthrough code authenticates the upstream exchange the way its own
registration was granted, not the way the server row is configured. A sealed
``client_secret_basic`` grant sends the Basic header and keeps the secret out of the body; a
sealed public client (no method) keeps the body-credential path."""
import base64
import httpx
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.true_passthrough, dcr_bridge=None, server_id="sealed_method_srv")
upstream_request = httpx.Request("POST", server.token_url)
upstream_response = httpx.Response(
200, json={"access_token": "up-token", "token_type": "Bearer"}, request=upstream_request
)
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=upstream_response)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
):
await exchange_token_with_server(
request=_bridge_mock_request(),
mcp_server=server,
grant_type="authorization_code",
code="up-code",
redirect_uri="https://litellm.example.com/callback",
client_id="minted-77",
client_secret="mint-secret",
code_verifier="verifier",
client_token_endpoint_auth_method=sealed_auth_method,
)
sent_headers = mock_async_client.post.call_args.kwargs["headers"]
sent_body = mock_async_client.post.call_args.kwargs["data"]
if expects_basic_header:
expected = base64.b64encode(b"minted-77:mint-secret").decode()
assert sent_headers["Authorization"] == f"Basic {expected}"
assert "client_secret" not in sent_body
else:
assert "Authorization" not in sent_headers
assert sent_body["client_id"] == "minted-77"
assert sent_body["client_secret"] == "mint-secret"

View file

@ -2,6 +2,7 @@ import os
import sys
import types
import json
from contextlib import ExitStack
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import List, Optional
@ -2025,8 +2026,469 @@ class TestTemporaryMCPSessionEndpoints:
code_challenge_method="S256",
response_type="code",
scope="scope1",
ephemeral_dcr_client=None,
)
async def _authorize_without_client_id(
self, server, mint_mock=None, code_challenge="chal", code_challenge_method="S256"
):
"""Drive mcp_authorize with no caller client_id against ``server``, returning the
(authorize_with_server mock, raised HTTPException or None) pair. Sends a valid S256 PKCE
pair by default because the ephemeral mint requires it."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_authorize,
)
request = MagicMock()
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
patches = [
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server",
AsyncMock(return_value=MagicMock()),
),
]
if mint_mock is not None:
patches.append(
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client",
mint_mock,
)
)
with ExitStack() as stack:
entered = [stack.enter_context(p) for p in patches]
authorize_mock = entered[1]
try:
await mcp_authorize(
request=request,
server_id=server.server_id,
user_api_key_dict=admin_auth,
client_id=None,
redirect_uri="http://127.0.0.1:60108/callback",
state="state123",
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
)
except HTTPException as exc:
return authorize_mock, exc
return authorize_mock, None
@pytest.mark.asyncio
@pytest.mark.parametrize(
"code_challenge, code_challenge_method",
[(None, None), ("chal", "plain"), ("chal", None)],
)
async def test_mcp_authorize_mint_requires_s256_pkce(self, code_challenge, code_challenge_method):
"""Without PKCE the sealed code would be bearer-redeemable by any authenticated caller who
intercepts the redirect, so the ephemeral mint refuses to run for a downgraded flow (no
challenge, or a non-S256 method) before any upstream registration happens."""
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
server.authorization_url = "https://idp.example.com/authorize"
server.registration_url = "https://idp.example.com/register"
mint_mock = AsyncMock()
authorize_mock, exc = await self._authorize_without_client_id(
server, mint_mock=mint_mock, code_challenge=code_challenge, code_challenge_method=code_challenge_method
)
assert exc is not None
assert exc.status_code == 400
assert "PKCE" in str(exc.detail)
mint_mock.assert_not_awaited()
authorize_mock.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
async def test_mcp_authorize_client_forwarded_modes_mint_ephemeral_dcr_client_when_none_supplied(self, auth_type):
"""LIT-4581 regression: a client-forwarded-token server created without an auth step has no
stored client_id and the tools-tab browser flow supplies none, so authorize must fall
through to a gateway-side DCR mint and proceed with the minted client instead of
dead-ending on a 400 missing_client_id. Both modes share the caller-held-client contract,
so both get the fall-through."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
EphemeralDcrClient,
)
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = auth_type
server.authorization_url = "https://idp.example.com/authorize"
server.registration_url = "https://idp.example.com/register"
minted = EphemeralDcrClient(client_id="minted-77", client_secret="mint-secret")
mint_mock = AsyncMock(return_value=minted)
authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock)
assert exc is None
mint_mock.assert_awaited_once()
assert authorize_mock.await_args.kwargs["client_id"] == "minted-77"
assert authorize_mock.await_args.kwargs["ephemeral_dcr_client"] is minted
@pytest.mark.asyncio
async def test_mcp_authorize_rejects_untrusted_redirect_before_minting(self):
"""An untrusted redirect_uri must be rejected before the gateway performs any upstream
registration, so bad-redirect requests cannot be used to generate orphan clients at the
IdP."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_authorize,
)
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
server.authorization_url = "https://idp.example.com/authorize"
server.registration_url = "https://idp.example.com/register"
mint_mock = AsyncMock()
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
request = MagicMock()
request.base_url = "https://litellm.example.com/"
request.headers = {}
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client",
mint_mock,
),
):
with pytest.raises(HTTPException) as exc:
await mcp_authorize(
request=request,
server_id="server-1",
user_api_key_dict=admin_auth,
client_id=None,
redirect_uri="https://evil.example.net/steal",
state="state123",
code_challenge="chal",
code_challenge_method="S256",
)
assert exc.value.status_code == 400
mint_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_authorize_true_passthrough_without_authorization_url_reports_the_real_fault(self):
"""A passthrough server whose discovery never yielded an authorize endpoint cannot start any
flow, minted client or not, so the error names the missing authorization url instead of the
misleading missing_client_id remedy."""
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
server.authorization_url = None
server.registration_url = "https://idp.example.com/register"
mint_mock = AsyncMock()
authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock)
assert exc is not None
assert exc.status_code == 400
assert "authorization url" in str(exc.detail)
mint_mock.assert_not_awaited()
authorize_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_authorize_true_passthrough_without_registration_endpoint_keeps_missing_client_id(self):
"""When the upstream exposes no registration endpoint the mint is impossible, so the
authorize fails closed with the existing missing_client_id 400 instead of proceeding with an
empty client."""
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
server.authorization_url = "https://idp.example.com/authorize"
server.registration_url = None
authorize_mock, exc = await self._authorize_without_client_id(server)
assert exc is not None
assert exc.status_code == 400
assert exc.detail["error"] == "missing_client_id"
authorize_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_authorize_oauth2_server_does_not_mint(self):
"""The ephemeral mint is scoped to the client-forwarded-token modes: a plain oauth2 server
keeps the gateway-held-client contract (its client is persisted by the admin register flow),
so an empty client_id stays a 400 and no upstream registration is attempted."""
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.oauth2
server.authorization_url = "https://idp.example.com/authorize"
server.registration_url = "https://idp.example.com/register"
mint_mock = AsyncMock()
authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock)
assert exc is not None
assert exc.status_code == 400
assert exc.detail["error"] == "missing_client_id"
mint_mock.assert_not_awaited()
authorize_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_authorize_true_passthrough_dcr_bridge_mints_too(self):
"""The UI creates passthrough servers with dcr_bridge enabled by default, so the default
clientless tools-page authorize is a bridge server; it must mint exactly like a non-bridge
one (the minted flow runs the bridge short-circuit arm) instead of dead-ending on
missing_client_id. The relay front door stays reserved for clients that present their own
client_id."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
EphemeralDcrClient,
)
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
server.dcr_bridge = True
server.authorization_url = "https://idp.example.com/authorize"
server.registration_url = "https://idp.example.com/register"
minted = EphemeralDcrClient(client_id="minted-77", client_secret=None)
mint_mock = AsyncMock(return_value=minted)
authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock)
assert exc is None
mint_mock.assert_awaited_once()
assert authorize_mock.await_args.kwargs["client_id"] == "minted-77"
assert authorize_mock.await_args.kwargs["ephemeral_dcr_client"] is minted
@pytest.mark.asyncio
async def test_mcp_authorize_oauth_delegate_dcr_bridge_does_not_mint(self):
"""The interactive oauth_delegate dcr_bridge sign-in has its own sealed-identity flow that
captures the SSO user at authorize; the ephemeral mint must not preempt it."""
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.oauth_delegate
server.dcr_bridge = True
server.authorization_url = "https://idp.example.com/authorize"
server.registration_url = "https://idp.example.com/register"
mint_mock = AsyncMock()
authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock)
assert exc is not None
assert exc.status_code == 400
assert exc.detail["error"] == "missing_client_id"
mint_mock.assert_not_awaited()
authorize_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_token_opens_sealed_passthrough_code_and_exchanges_with_minted_client(self):
"""LIT-4581 regression, token leg: the client echoes back the sealed passthrough code the
callback forwarded, so the token endpoint recovers the ephemeral client and the real
upstream code from it and authenticates the exchange with them, with no client_id supplied
by the caller and none stored on the server."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
seal_passthrough_authorization_code,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_token,
)
request = MagicMock()
request.base_url = "https://litellm.example.com/"
request.headers = {}
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
AsyncMock(return_value={"access_token": "token"}),
) as exchange_mock,
):
sealed = seal_passthrough_authorization_code(
upstream_code="up-code",
client_id="minted-77",
client_secret="mint-secret",
mcp_server_id="server-1",
token_endpoint_auth_method="client_secret_basic",
)
result = await mcp_token(
request=request,
server_id="server-1",
user_api_key_dict=admin_auth,
grant_type="authorization_code",
code=sealed,
redirect_uri="https://example.com/callback",
client_id=None,
client_secret=None,
code_verifier="verifier",
refresh_token=None,
scope=None,
)
assert result == {"access_token": "token"}
assert exchange_mock.await_args.kwargs["code"] == "up-code"
assert exchange_mock.await_args.kwargs["client_id"] == "minted-77"
assert exchange_mock.await_args.kwargs["client_secret"] == "mint-secret"
assert exchange_mock.await_args.kwargs["redirect_uri"] == "https://litellm.example.com/callback"
assert exchange_mock.await_args.kwargs["client_token_endpoint_auth_method"] == "client_secret_basic"
@pytest.mark.asyncio
async def test_mcp_token_refresh_grant_never_opens_sealed_code(self):
"""The minted client is unrecoverable outside the single authorization_code flow by
contract: a refresh_token grant that echoes a leftover sealed passthrough code (plus any
verifier) must not recover the minted credentials, so a clientless server answers
missing_client_id and the client re-runs authorize instead."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
seal_passthrough_authorization_code,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_token,
)
request = MagicMock()
request.base_url = "https://litellm.example.com/"
request.headers = {}
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
AsyncMock(return_value={"access_token": "token"}),
) as exchange_mock,
):
sealed = seal_passthrough_authorization_code(
upstream_code="up-code",
client_id="minted-77",
client_secret="mint-secret",
mcp_server_id="server-1",
token_endpoint_auth_method="client_secret_basic",
)
with pytest.raises(HTTPException) as exc:
await mcp_token(
request=request,
server_id="server-1",
user_api_key_dict=admin_auth,
grant_type="refresh_token",
code=sealed,
redirect_uri="https://example.com/callback",
client_id=None,
client_secret=None,
code_verifier="verifier",
refresh_token="leftover-refresh",
scope=None,
)
assert exc.value.status_code == 400
assert exc.value.detail["error"] == "missing_client_id"
exchange_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_token_sealed_code_requires_code_verifier(self):
"""A sealed code is minted only for S256 PKCE flows, so redeeming one without the
corresponding verifier is refused at the gateway rather than trusting the upstream to
enforce the binding."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
seal_passthrough_authorization_code,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_token,
)
request = MagicMock()
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
AsyncMock(),
) as exchange_mock,
):
sealed = seal_passthrough_authorization_code(
upstream_code="up-code", client_id="minted-77", client_secret=None, mcp_server_id="server-1"
)
with pytest.raises(HTTPException) as exc:
await mcp_token(
request=request,
server_id="server-1",
user_api_key_dict=admin_auth,
grant_type="authorization_code",
code=sealed,
redirect_uri="https://example.com/callback",
client_id=None,
client_secret=None,
code_verifier=None,
refresh_token=None,
scope=None,
)
assert exc.value.status_code == 400
assert "code_verifier" in str(exc.value.detail)
exchange_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_token_rejects_sealed_code_for_another_server(self):
"""A sealed passthrough code is bound to the server it was minted for: presenting it at
another server's token endpoint is a 400 before any upstream exchange, so a code cannot be
replayed across a server boundary."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
seal_passthrough_authorization_code,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_token,
)
request = MagicMock()
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.true_passthrough
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
with (
patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
AsyncMock(),
) as exchange_mock,
):
sealed = seal_passthrough_authorization_code(
upstream_code="up-code",
client_id="minted-77",
client_secret=None,
mcp_server_id="a-different-server",
)
with pytest.raises(HTTPException) as exc:
await mcp_token(
request=request,
server_id="server-1",
user_api_key_dict=admin_auth,
grant_type="authorization_code",
code=sealed,
redirect_uri="https://example.com/callback",
client_id=None,
client_secret=None,
code_verifier="verifier",
refresh_token=None,
scope=None,
)
assert exc.value.status_code == 400
exchange_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_authorize_rejects_non_oauth2_server(self):
"""mcp_authorize must reject a none-auth server with an accurate 'does not use OAuth'
@ -2163,6 +2625,7 @@ class TestTemporaryMCPSessionEndpoints:
code_verifier="verifier",
refresh_token=None,
scope=None,
client_token_endpoint_auth_method=None,
)
@pytest.mark.asyncio
@ -2216,6 +2679,7 @@ class TestTemporaryMCPSessionEndpoints:
code_verifier=None,
refresh_token="rt-123",
scope=None,
client_token_endpoint_auth_method=None,
)
@pytest.mark.asyncio
@ -2270,8 +2734,59 @@ class TestTemporaryMCPSessionEndpoints:
token_endpoint_auth_method="client_secret_basic",
fallback_client_id="server-1",
persist_credentials=True,
client_redirect_uris=None,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"raw_redirect_uris, forwarded",
[
(["https://app.example.com/ui/callback"], ["https://app.example.com/ui/callback"]),
(["https://app.example.com/ui/callback", 42, "", None], None),
("not-a-list", None),
([], None),
([123], None),
],
)
async def test_mcp_register_forwards_validated_redirect_uris(self, raw_redirect_uris, forwarded):
"""dcr_bridge servers relay the registration upstream and require the browser client's own
redirect_uris, so mcp_register must forward them; the value is caller-controlled and is
validated by the same client_supplied_redirect_uris boundary helper as the root /register
door, so a malformed list is rejected whole at both doors (RFC 7591 redirect_uris is
all-or-nothing) rather than silently forwarding the surviving entries here and rejecting
them there."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_register,
)
request = MagicMock()
server = generate_mock_mcp_server_config_record(server_id="server-1")
server.auth_type = MCPAuth.oauth2
request_body = {"client_name": "LiteLLM", "redirect_uris": raw_redirect_uris}
admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
AsyncMock(return_value=request_body),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.register_client_with_server",
AsyncMock(return_value={"client_id": "generated"}),
) as register_mock,
):
await mcp_register(
request=request,
server_id="server-1",
user_api_key_dict=admin_auth,
)
assert register_mock.await_args.kwargs["client_redirect_uris"] == forwarded
@pytest.mark.asyncio
async def test_mcp_register_does_not_persist_for_non_admin(self):
"""A non-admin caller (who may have access to a real server) must not persist the DCR

View file

@ -9813,7 +9813,6 @@ async def _drive_team_write(
raw_body=None,
user=None,
find_returns_none=False,
json_side_effect=None,
):
"""Drive POST ``update_team`` or PATCH ``patch_team`` against a mocked team.
@ -9828,6 +9827,7 @@ async def _drive_team_write(
from litellm.proxy._types import (
LiteLLM_TeamTable,
LitellmUserRoles,
PatchTeamRequest,
UpdateTeamRequest,
UserAPIKeyAuth,
)
@ -9873,14 +9873,10 @@ async def _drive_team_write(
litellm_changed_by=None,
)
else:
if json_side_effect is not None:
req.json = AsyncMock(side_effect=json_side_effect)
else:
req.json = AsyncMock(
return_value=raw_body if raw_body is not None else dict(payload or {})
)
body = raw_body if raw_body is not None else dict(payload or {})
result = await patch_team(
team_id=_PATCH_TEAM_ID,
data=PatchTeamRequest.model_validate(body),
http_request=req,
user_api_key_dict=auth,
litellm_changed_by=None,
@ -10028,25 +10024,36 @@ async def test_patch_strips_system_managed_metadata_key_like_post():
assert patch_meta == {"cost_center": "9999"}
@pytest.mark.asyncio
@pytest.mark.parametrize("raw_body", [["not", "an", "object"], "a-string", 42, True])
async def test_patch_rejects_non_object_body(raw_body):
from litellm.proxy._types import ProxyException
@pytest.mark.parametrize(
"kwargs",
[
{"json": ["not", "an", "object"]},
{"json": "a-string"},
{"json": 42},
{"content": b"{not json"},
{"json": {"tpm_limit": "not-an-int"}},
],
ids=["list", "string", "number", "malformed-json", "wrong-field-type"],
)
def test_patch_rejects_a_malformed_body_with_422(kwargs):
"""The body is a declared parameter, so FastAPI rejects a malformed one before the
handler runs. This is the same 422 POST /team/update already returns; the route
previously answered 400 here and 500 for a wrongly typed field, reporting a caller
mistake as a server fault."""
from fastapi import FastAPI
from fastapi.testclient import TestClient
with pytest.raises(ProxyException) as exc:
await _drive_team_write("patch", existing_metadata={"a": 1}, raw_body=raw_body)
assert exc.value.code == "400" or exc.value.code == 400
from litellm.proxy._types import PatchTeamRequest
app = FastAPI()
@pytest.mark.asyncio
async def test_patch_rejects_invalid_json_body():
from litellm.proxy._types import ProxyException
@app.patch("/team/{team_id}")
async def _route(team_id: str, data: PatchTeamRequest): # pragma: no cover - schema only
return {}
with pytest.raises(ProxyException) as exc:
await _drive_team_write(
"patch", existing_metadata={"a": 1}, json_side_effect=ValueError("no body")
)
assert exc.value.code == "400" or exc.value.code == 400
response = TestClient(app).patch("/team/abc", **kwargs)
assert response.status_code == 422
@pytest.mark.asyncio
@ -10116,3 +10123,103 @@ async def test_patch_returns_full_team_object_not_wrapper():
)
assert isinstance(result, LiteLLM_TeamTable)
assert result.team_id == _PATCH_TEAM_ID
# ---------------------------------------------------------------------------
# PATCH body is validated through PatchTeamRequest before it is handed to
# update_team. The write below must stay byte-identical to what the untyped
# **body construction produced, or a partial update starts writing columns the
# caller never mentioned.
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_patch_writes_only_the_keys_the_caller_sent():
"""An omitted field must not reach the DB write at all. If validation ever
materialises defaults, every unmentioned column gets overwritten with null."""
_, update_mock = await _drive_team_write("patch", raw_body={"tpm_limit": 5})
written = update_mock.call_args.kwargs["data"]
assert written["tpm_limit"] == 5
for untouched in ("rpm_limit", "max_budget", "models", "blocked", "budget_duration"):
assert untouched not in written, f"{untouched} was written despite not being sent"
@pytest.mark.asyncio
async def test_patch_preserves_explicit_null_as_a_clear():
"""null is a clear, not an omission: it has to survive validation and reach the write."""
_, update_mock = await _drive_team_write("patch", raw_body={"max_budget": None})
written = update_mock.call_args.kwargs["data"]
assert "max_budget" in written
assert written["max_budget"] is None
def _patch_body_to_update_request(body: dict):
"""The exact reshaping patch_team performs between the raw body and update_team."""
from litellm.proxy._types import PatchTeamRequest, UpdateTeamRequest
parsed = PatchTeamRequest.model_validate(body)
return UpdateTeamRequest(
team_id=_PATCH_TEAM_ID,
**parsed.model_dump(exclude_unset=True, exclude={"team_id"}),
)
@pytest.mark.parametrize(
"body",
[
{"tpm_limit": 5},
{"max_budget": None},
{"object_permission": {"vector_stores": []}},
{"metadata": {"a": 1, "b": None}},
{"models": ["gpt-4"], "blocked": False},
],
ids=["scalar", "explicit-null", "partial-nested", "metadata-with-null", "list-and-false"],
)
def test_patch_body_reshaping_adds_no_keys_the_caller_did_not_send(body):
"""Validating through PatchTeamRequest must be shape-preserving. If it ever
materialises defaults, a partial update silently overwrites untouched columns,
and for the merge-only object_permission it would wipe sibling sub-keys."""
reshaped = _patch_body_to_update_request(body)
dumped = reshaped.model_dump(exclude_unset=True, exclude={"team_id"})
assert dumped == body
assert reshaped.model_fields_set == set(body) | {"team_id"}
@pytest.mark.asyncio
async def test_patch_ignores_unknown_body_keys():
"""Unknown keys were silently dropped by the previous construction; keep that."""
_, update_mock = await _drive_team_write(
"patch", raw_body={"tpm_limit": 5, "not_a_team_field": "x"}
)
written = update_mock.call_args.kwargs["data"]
assert written["tpm_limit"] == 5
assert "not_a_team_field" not in written
def test_patch_team_request_makes_team_id_optional():
"""PATCH takes team_id from the path, so the body model must not require it,
while still inheriting every UpdateTeamRequest field."""
from litellm.proxy._types import PatchTeamRequest, UpdateTeamRequest
parsed = PatchTeamRequest.model_validate({"tpm_limit": 5})
assert parsed.team_id is None
assert parsed.model_fields_set == {"tpm_limit"}
assert set(UpdateTeamRequest.model_fields).issubset(set(PatchTeamRequest.model_fields))
def test_patch_team_route_publishes_its_request_body_schema():
"""The dashboard's generated client types this call off the OpenAPI spec, which
FastAPI can only emit because the body is a declared parameter."""
from litellm.proxy.proxy_server import app
operation = app.openapi()["paths"]["/team/{team_id}"]["patch"]
schema = operation["requestBody"]["content"]["application/json"]["schema"]
assert schema == {"$ref": "#/components/schemas/PatchTeamRequest"}
properties = app.openapi()["components"]["schemas"]["PatchTeamRequest"]["properties"]
assert "tpm_limit" in properties and "metadata" in properties

View file

@ -14,13 +14,19 @@ Covers:
"""
import asyncio
from typing import List
from typing import List, cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import (
AllMessageValues,
ResponseInputParam,
)
# ---------------------------------------------------------------------------
# Helpers
@ -71,6 +77,56 @@ def _patch_responses_dispatch():
]
def _make_cache_control_case() -> tuple[
ResponseInputParam,
list[AllMessageValues],
dict[str, object],
]:
system_message = cast(
AllMessageValues,
{"role": "system", "content": "Analyze the request"},
)
assistant_message = cast(
AllMessageValues,
{
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [
{
"type": "output_text",
"text": "The code has a bug",
"annotations": [],
}
],
},
)
user_message = cast(
AllMessageValues,
{"role": "user", "content": "Check for security issues"},
)
reasoning_item = {
"type": "reasoning",
"id": "rs_1",
"summary": [],
"encrypted_content": "encrypted",
}
original_input = cast(
ResponseInputParam,
[system_message, reasoning_item, assistant_message, user_message],
)
_, merged_messages, _ = AnthropicCacheControlHook().get_chat_completion_prompt(
model="azure/gpt-5-codex",
messages=[system_message, assistant_message, user_message],
non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
return original_input, merged_messages, reasoning_item
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
@ -256,6 +312,66 @@ class TestResponsesAPIPromptManagement:
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)
assert len(passed_messages) == 1
def test_cache_control_hook_preserves_reasoning_items(self):
original_input, merged_messages, reasoning_item = _make_cache_control_case()
logging_obj = _make_logging_obj(
merged_model="azure/gpt-5-codex",
merged_messages=merged_messages,
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3] as mock_handler:
import litellm
litellm.responses(
input=original_input,
model="azure/gpt-5-codex",
litellm_logging_obj=logging_obj,
cache_control_injection_points=[{"location": "message", "role": "system"}],
)
sent_input = mock_handler.call_args.kwargs["input"]
assert [item.get("type") for item in sent_input] == [
None,
"reasoning",
"message",
None,
]
assert sent_input[0]["cache_control"] == {"type": "ephemeral"}
assert sent_input[1] == reasoning_item
assert sent_input[2]["id"] == "msg_1"
def test_all_non_message_input_items_remain_unchanged(self):
reasoning_item = {
"type": "reasoning",
"id": "rs_1",
"summary": [],
"encrypted_content": "encrypted",
}
original_input = cast(ResponseInputParam, [reasoning_item])
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=[
cast(
AllMessageValues,
{"role": "system", "content": "Analyze the request"},
)
],
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3] as mock_handler:
import litellm
litellm.responses(
input=original_input,
model="gpt-4o",
prompt_id="all-non-message",
litellm_logging_obj=logging_obj,
)
assert mock_handler.call_args.kwargs["input"] == original_input
def test_model_override_re_resolves_provider(self):
"""[G] When the prompt template overrides the model to a different provider,
custom_llm_provider is re-resolved so downstream routing uses the correct provider.
@ -393,3 +509,33 @@ class TestAsyncResponsesAPIPromptManagement:
passed_messages = call_kwargs["messages"]
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)
assert len(passed_messages) == 1
@pytest.mark.asyncio
async def test_async_cache_control_hook_preserves_reasoning_items(self):
original_input, merged_messages, reasoning_item = _make_cache_control_case()
logging_obj = _make_logging_obj(
merged_model="azure/gpt-5-codex",
merged_messages=merged_messages,
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3] as mock_handler:
import litellm
await litellm.aresponses(
input=original_input,
model="azure/gpt-5-codex",
litellm_logging_obj=logging_obj,
cache_control_injection_points=[{"location": "message", "role": "system"}],
)
sent_input = mock_handler.call_args.kwargs["input"]
assert [item.get("type") for item in sent_input] == [
None,
"reasoning",
"message",
None,
]
assert sent_input[0]["cache_control"] == {"type": "ephemeral"}
assert sent_input[1] == reasoning_item
assert sent_input[2]["id"] == "msg_1"

File diff suppressed because it is too large Load diff

View file

@ -19,12 +19,13 @@ const eslintConfig = [
"unused-imports/no-unused-imports": "error",
"local/no-large-inline-object-arg": "warn",
"local/no-long-condition-chain": "warn",
"local/no-complex-jsx-arrow": ["error", { maxStatements: 2 }],
"@typescript-eslint/no-explicit-any": "warn",
"no-console": ["warn", { allow: ["warn", "error"] }],
"@typescript-eslint/no-unused-vars": "off",
"@typescript-eslint/no-unused-expressions": "off",
"@typescript-eslint/ban-ts-comment": "off",
"prefer-const": "off",
"prefer-const": "error",
"no-empty": "off",
"no-prototype-builtins": "off",
"no-useless-catch": "off",
@ -51,13 +52,32 @@ const eslintConfig = [
patterns: [
{
group: ["@tremor/react", "@tremor/react/*"],
message: "@tremor/react is being phased out; build new UI with antd instead of adding tremor imports.",
message:
"@tremor/react is being phased out; build new UI with shadcn/ui primitives instead of adding tremor imports.",
},
{
group: ["antd", "antd/*"],
message:
"antd is being phased out; build new UI with shadcn/ui primitives instead of adding antd imports.",
},
],
},
],
},
},
{
files: ["src/**/*.tsx"],
rules: {
"local/filename-pascal-case": "error",
},
},
{
files: ["src/**/*.{ts,tsx}"],
ignores: ["src/**/*.test.{ts,tsx}", "src/**/*.spec.{ts,tsx}", "src/data/**"],
rules: {
"max-lines": ["error", { max: 800, skipBlankLines: true, skipComments: true }],
},
},
{
files: ["src/lib/http/**"],
rules: {

View file

@ -27,7 +27,7 @@
"jwt-decode": "4.0.0",
"lucide-react": "0.513.0",
"moment": "2.30.1",
"next": "16.2.6",
"next": "16.2.11",
"openai": "4.104.0",
"openapi-fetch": "^0.17.0",
"openapi-react-query": "^0.5.4",
@ -61,7 +61,7 @@
"@vitest/coverage-v8": "3.2.6",
"@vitest/ui": "3.2.6",
"eslint": "9.39.2",
"eslint-config-next": "16.2.6",
"eslint-config-next": "16.2.11",
"eslint-config-prettier": "10.1.8",
"eslint-plugin-unused-imports": "4.3.0",
"jsdom": "27.4.0",
@ -2244,15 +2244,15 @@
}
},
"node_modules/@next/env": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/env/-/env-16.2.6.tgz",
"integrity": "sha512-gd8HoHN4ufj73WmR3JmVolrpJR47ILK6LouP5xElPglaVxir6e1a7VzvTvDWkOoPXT9rkkTzyCxBu4yeZfZwcw==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/env/-/env-16.2.11.tgz",
"integrity": "sha512-0do5A3BJ2gxWr0ZCMcD6BhW+e595jyxdTl3rXTS6lOtD8ektMiW6CO+EPwt1Eca1DBnm90r/7GdiKWBKxH++DA==",
"license": "MIT"
},
"node_modules/@next/eslint-plugin-next": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-16.2.6.tgz",
"integrity": "sha512-Z8l6o4JWKUl755x4R+wogD86KPeU+Ckw4K+SYG4kHeOJtRenDeK+OSbGcqZpDtbwn9DsJVdir2UxmwXuinUbUw==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-16.2.11.tgz",
"integrity": "sha512-vMEf/aXOpzFFdtIvFYOnIDPKb0xBbrXONsz83CcKdRrekfxNdL8PNkq5qHqAHSXVlIifnX68LOMaxr3z5PkeLQ==",
"dev": true,
"license": "MIT",
"dependencies": {
@ -2260,9 +2260,9 @@
}
},
"node_modules/@next/swc-darwin-arm64": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-darwin-arm64/-/swc-darwin-arm64-16.2.6.tgz",
"integrity": "sha512-ZJGkkcNfYgrrMkqOdZ7zoLa1TOy0qpcMfk/z4Mh/FKUz40gVO+HNQWqmLxf67Z5WB64DRp0dhEbyHfel+6sJUg==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-darwin-arm64/-/swc-darwin-arm64-16.2.11.tgz",
"integrity": "sha512-wryL4pjKmDwGv2ox6+GZDFxvmtSRLqApBR8kL1j4+vhB7Z5vJC/zAnXpiR9Xkfzl0AS8WLMnsuGV/UKI67/rrw==",
"cpu": [
"arm64"
],
@ -2276,9 +2276,9 @@
}
},
"node_modules/@next/swc-darwin-x64": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-darwin-x64/-/swc-darwin-x64-16.2.6.tgz",
"integrity": "sha512-v/YLBHIY132Ced3puBJ7YJKw1lqsCrgcNo2aRJlCEyQrrCeRJlvGlnmxhPxNQI3KE3N1DN5r9TPNPvka3nq5RQ==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-darwin-x64/-/swc-darwin-x64-16.2.11.tgz",
"integrity": "sha512-aZl2j4f/fLyjQvOhv0Oe9UaMAQHolYpKhctsoYzplSumKJKPUmgjcf6545aBtysLTcu994TREd0+pSgNE4ohmg==",
"cpu": [
"x64"
],
@ -2292,9 +2292,9 @@
}
},
"node_modules/@next/swc-linux-arm64-gnu": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-gnu/-/swc-linux-arm64-gnu-16.2.6.tgz",
"integrity": "sha512-RPOvqlYBbcQjkz9VQQDZ2T2bARIjXZV1KFlt+V2Mr6SW/e4I9fcKsaA0hdyf2FHoTlsV2xnBd5Y912rP/1Ce6w==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-gnu/-/swc-linux-arm64-gnu-16.2.11.tgz",
"integrity": "sha512-5jEriyEnH/LWFy27L2ZG0XaLlyEJIjhsImEsiS9P563PKEVp2BVups/xfOucIrsvVntp11oNcZwjHvaDPYVB5g==",
"cpu": [
"arm64"
],
@ -2308,9 +2308,9 @@
}
},
"node_modules/@next/swc-linux-arm64-musl": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-musl/-/swc-linux-arm64-musl-16.2.6.tgz",
"integrity": "sha512-URUTu1+dMkxJsPFgm+OeEvq9wf5sujw0EvgYy80TDGHTSLTnIHeqb0Eu8A3sC95IRgjejQL+kC4mw+4yPxiAXA==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-musl/-/swc-linux-arm64-musl-16.2.11.tgz",
"integrity": "sha512-eIjcpx2fnnFSSkZDbTxy74KnokUXDjfoLClpWelfgHLf621aTqswhwXQ7GkD5K5rplrS6LZ/Bj+mVuvzluBOEg==",
"cpu": [
"arm64"
],
@ -2324,9 +2324,9 @@
}
},
"node_modules/@next/swc-linux-x64-gnu": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-gnu/-/swc-linux-x64-gnu-16.2.6.tgz",
"integrity": "sha512-DOj182mPV8G3UkrayLoREM5YEYI+Dk5wv7Ox9xl1fFibAELEsFD0lDPfHIeILlutMMfdyhlzYPELG3peuKaurw==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-gnu/-/swc-linux-x64-gnu-16.2.11.tgz",
"integrity": "sha512-8WgzpaWMs46qJT9kiV47cje86L0x/Mu9t8/Gwj+pnbgW3rETVfCnaScPjlYUwNScpOozdcIMHWmAvuZJUonR2w==",
"cpu": [
"x64"
],
@ -2340,9 +2340,9 @@
}
},
"node_modules/@next/swc-linux-x64-musl": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-musl/-/swc-linux-x64-musl-16.2.6.tgz",
"integrity": "sha512-HKQ5SP/V/ub73UvF7n/zeJlxk2kLmtL7Wzrg4WfmkjmNos5onJ2tKu7yZOPdL18A6Svfn3max29ym+ry7NkK4g==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-musl/-/swc-linux-x64-musl-16.2.11.tgz",
"integrity": "sha512-I3UgPds7G4ZYnTb/H+5GBGuUT2DhAk6j0mL6A4s63RjFs74wB2hOWP0vaxsK+3NJraExt3eYEPQ/UtT0x/64Nw==",
"cpu": [
"x64"
],
@ -2356,9 +2356,9 @@
}
},
"node_modules/@next/swc-win32-arm64-msvc": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-win32-arm64-msvc/-/swc-win32-arm64-msvc-16.2.6.tgz",
"integrity": "sha512-LZXpTlPyS5v7HhSmnvsLGP3iIYgYOBnc8r8ArlT55sGHV89bR2HlDdBjWQ+PY6SJMmk8TuVGFuxalnP3k/0Dwg==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-win32-arm64-msvc/-/swc-win32-arm64-msvc-16.2.11.tgz",
"integrity": "sha512-n89CjtcThnjrwgJMAiI5xbqwLY51zvwC9tSlArmVndAJLYVl9T9UAdlkXTmZvE++idoXe8KdglQlhNRdUp1c6g==",
"cpu": [
"arm64"
],
@ -2372,9 +2372,9 @@
}
},
"node_modules/@next/swc-win32-x64-msvc": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/@next/swc-win32-x64-msvc/-/swc-win32-x64-msvc-16.2.6.tgz",
"integrity": "sha512-F0+4i0h9J6C4eE3EAPWsoCk7UW/dbzOjyzxY0qnDUOYFu6FFmdZ6l97/XdV3/Nz3VYyO7UWjyEJUXkGqcoXfMA==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/@next/swc-win32-x64-msvc/-/swc-win32-x64-msvc-16.2.11.tgz",
"integrity": "sha512-md8CLNggS1Dx9pUgApzps5uAf+N8GN9xywzmNx9vHAWo94HtBwCCqkSnhIrdfQe83Dhz8Lfo/20Nb1Zxal092w==",
"cpu": [
"x64"
],
@ -3620,6 +3620,72 @@
"node": ">=14.0.0"
}
},
"node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@emnapi/core": {
"version": "1.11.1",
"dev": true,
"inBundle": true,
"license": "MIT",
"optional": true,
"dependencies": {
"@emnapi/wasi-threads": "1.2.2",
"tslib": "^2.4.0"
}
},
"node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@emnapi/runtime": {
"version": "1.11.1",
"dev": true,
"inBundle": true,
"license": "MIT",
"optional": true,
"dependencies": {
"tslib": "^2.4.0"
}
},
"node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@emnapi/wasi-threads": {
"version": "1.2.2",
"dev": true,
"inBundle": true,
"license": "MIT",
"optional": true,
"dependencies": {
"tslib": "^2.4.0"
}
},
"node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@napi-rs/wasm-runtime": {
"version": "1.1.4",
"dev": true,
"inBundle": true,
"license": "MIT",
"optional": true,
"dependencies": {
"@tybys/wasm-util": "^0.10.1"
},
"funding": {
"type": "github",
"url": "https://github.com/sponsors/Brooooooklyn"
},
"peerDependencies": {
"@emnapi/core": "^1.7.1",
"@emnapi/runtime": "^1.7.1"
}
},
"node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@tybys/wasm-util": {
"version": "0.10.2",
"dev": true,
"inBundle": true,
"license": "MIT",
"optional": true,
"dependencies": {
"tslib": "^2.4.0"
}
},
"node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/tslib": {
"version": "2.8.1",
"dev": true,
"inBundle": true,
"license": "0BSD",
"optional": true
},
"node_modules/@tailwindcss/oxide-win32-arm64-msvc": {
"version": "4.3.2",
"resolved": "https://registry.npmjs.org/@tailwindcss/oxide-win32-arm64-msvc/-/oxide-win32-arm64-msvc-4.3.2.tgz",
@ -6642,13 +6708,13 @@
}
},
"node_modules/eslint-config-next": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-16.2.6.tgz",
"integrity": "sha512-z2ELYSkyrrJ6cuunTU8vhsT/RpouPkjaSah06nVW6Rg2Hpg0Vs8s497/e5s8G8qtdp4ccsiovz5P1rv+5VSW2Q==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-16.2.11.tgz",
"integrity": "sha512-FIpbK/dUyxUExchDB7eBg3k+VU8R2iR/Cx9/kqTBUTFv2bOIR9aRrpno4rvAQ9VhiPQAyFKNA2NlZwouGWtclA==",
"dev": true,
"license": "MIT",
"dependencies": {
"@next/eslint-plugin-next": "16.2.6",
"@next/eslint-plugin-next": "16.2.11",
"eslint-import-resolver-node": "^0.3.6",
"eslint-import-resolver-typescript": "^3.5.2",
"eslint-plugin-import": "^2.32.0",
@ -10309,12 +10375,12 @@
"license": "MIT"
},
"node_modules/next": {
"version": "16.2.6",
"resolved": "https://registry.npmjs.org/next/-/next-16.2.6.tgz",
"integrity": "sha512-qOVgKJg1+At15NpeUP+eJgCHvTCgXsogweq87Ri/Ix7PkqQHg4sdaXmSFqKlgaIXE4kW0g25LE68W87UANlHtw==",
"version": "16.2.11",
"resolved": "https://registry.npmjs.org/next/-/next-16.2.11.tgz",
"integrity": "sha512-B339zaqbyK8cmxhoAvLrcwoabwCP1wz21zSzfqxqXAemTu2BXnH7tQnfcglKv1vnMUIDBc+Hth7XODQriTZiRQ==",
"license": "MIT",
"dependencies": {
"@next/env": "16.2.6",
"@next/env": "16.2.11",
"@swc/helpers": "0.5.15",
"baseline-browser-mapping": "^2.9.19",
"caniuse-lite": "^1.0.30001579",
@ -10328,14 +10394,14 @@
"node": ">=20.9.0"
},
"optionalDependencies": {
"@next/swc-darwin-arm64": "16.2.6",
"@next/swc-darwin-x64": "16.2.6",
"@next/swc-linux-arm64-gnu": "16.2.6",
"@next/swc-linux-arm64-musl": "16.2.6",
"@next/swc-linux-x64-gnu": "16.2.6",
"@next/swc-linux-x64-musl": "16.2.6",
"@next/swc-win32-arm64-msvc": "16.2.6",
"@next/swc-win32-x64-msvc": "16.2.6",
"@next/swc-darwin-arm64": "16.2.11",
"@next/swc-darwin-x64": "16.2.11",
"@next/swc-linux-arm64-gnu": "16.2.11",
"@next/swc-linux-arm64-musl": "16.2.11",
"@next/swc-linux-x64-gnu": "16.2.11",
"@next/swc-linux-x64-musl": "16.2.11",
"@next/swc-win32-arm64-msvc": "16.2.11",
"@next/swc-win32-x64-msvc": "16.2.11",
"sharp": "^0.34.5"
},
"peerDependencies": {

View file

@ -39,7 +39,7 @@
"jwt-decode": "4.0.0",
"lucide-react": "0.513.0",
"moment": "2.30.1",
"next": "16.2.6",
"next": "16.2.11",
"openai": "4.104.0",
"openapi-fetch": "^0.17.0",
"openapi-react-query": "^0.5.4",
@ -73,7 +73,7 @@
"@vitest/coverage-v8": "3.2.6",
"@vitest/ui": "3.2.6",
"eslint": "9.39.2",
"eslint-config-next": "16.2.6",
"eslint-config-next": "16.2.11",
"eslint-config-prettier": "10.1.8",
"eslint-plugin-unused-imports": "4.3.0",
"jsdom": "27.4.0",

View file

@ -0,0 +1,59 @@
import { basename } from "path";
const NEXT_RESERVED = new Set([
"page",
"layout",
"route",
"template",
"default",
"loading",
"error",
"global-error",
"not-found",
"middleware",
"instrumentation",
"sitemap",
"robots",
"manifest",
"icon",
"apple-icon",
"favicon",
"opengraph-image",
"twitter-image",
]);
const PASCAL_CASE = /^[A-Z][A-Za-z0-9]*$/;
const rule = {
meta: {
type: "suggestion",
docs: {
description: "Require PascalCase filenames for .tsx modules; exempt Next.js reserved files and test/spec files.",
},
schema: [],
messages: {
notPascalCase: "Filename '{{name}}' should be PascalCase (e.g. '{{suggestion}}.tsx').",
},
},
create(context) {
const filename = context.filename;
const stem = basename(filename).replace(/\.tsx$/, "");
const [head, ...rest] = stem.split(".");
if (rest.includes("test") || rest.includes("spec")) return {};
if (NEXT_RESERVED.has(head)) return {};
if (PASCAL_CASE.test(head)) return {};
const pascalHead = head
.split(/[-_]/)
.filter(Boolean)
.map((part) => part.charAt(0).toUpperCase() + part.slice(1))
.join("");
const suggestion = [pascalHead, ...rest].join(".");
return {
Program(node) {
context.report({ node, messageId: "notPascalCase", data: { name: `${stem}.tsx`, suggestion } });
},
};
},
};
export default rule;

View file

@ -1,10 +1,14 @@
import noLargeInlineObjectArg from "./no-large-inline-object-arg.mjs";
import noLongConditionChain from "./no-long-condition-chain.mjs";
import noComplexJsxArrow from "./no-complex-jsx-arrow.mjs";
import filenamePascalCase from "./filename-pascal-case.mjs";
const plugin = {
rules: {
"no-large-inline-object-arg": noLargeInlineObjectArg,
"no-long-condition-chain": noLongConditionChain,
"no-complex-jsx-arrow": noComplexJsxArrow,
"filename-pascal-case": filenamePascalCase,
},
};

View file

@ -0,0 +1,41 @@
const DEFAULT_MAX_STATEMENTS = 2;
const isJsxAttributeValue = (node) => {
const parent = node.parent;
if (parent == null) return false;
return parent.type === "JSXExpressionContainer" && parent.parent?.type === "JSXAttribute";
};
const rule = {
meta: {
type: "suggestion",
docs: {
description:
"Disallow arrow functions with block bodies over a few statements passed inline as JSX attributes; extract them into a named handler.",
},
schema: [
{
type: "object",
properties: { maxStatements: { type: "integer", minimum: 1 } },
additionalProperties: false,
},
],
messages: {
tooComplex: "Inline JSX arrow handler has {{count}} statements; extract it into a named function (max {{max}}).",
},
},
create(context) {
const maxStatements = context.options[0]?.maxStatements ?? DEFAULT_MAX_STATEMENTS;
return {
ArrowFunctionExpression(node) {
if (node.body.type !== "BlockStatement") return;
if (!isJsxAttributeValue(node)) return;
const count = node.body.body.length;
if (count <= maxStatements) return;
context.report({ node, messageId: "tooComplex", data: { count, max: maxStatements } });
},
};
},
};
export default rule;

View file

@ -1,4 +1,5 @@
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi } from "vitest";
import APIReferenceView from "./APIReferenceView";
@ -44,4 +45,48 @@ describe("APIReferenceView", () => {
expect(renderedCode).toContain(apiDocUrl);
expect(renderedCode).not.toContain(proxyUrl);
});
it("renders the page title, blurb and docs link", () => {
render(<APIReferenceView proxySettings={{ PROXY_BASE_URL: "https://proxy.litellm.test" }} />);
expect(screen.getByText("OpenAI Compatible Proxy: API Reference")).toBeTruthy();
expect(screen.getByText(/LiteLLM is OpenAI Compatible/)).toBeTruthy();
const docsLink = screen.getByRole("link", { name: /API Reference Docs/ });
expect(docsLink.getAttribute("href")).toBe("https://docs.litellm.ai/docs/proxy/user_keys");
expect(docsLink.getAttribute("target")).toBe("_blank");
});
it("exposes the three SDK tabs with the first selected by default", () => {
render(<APIReferenceView proxySettings={{ PROXY_BASE_URL: "https://proxy.litellm.test" }} />);
expect(screen.getAllByRole("tab").map((tab) => tab.textContent)).toEqual([
"OpenAI Python SDK",
"LlamaIndex",
"Langchain Py",
]);
expect(screen.getAllByRole("tab").map((tab) => tab.getAttribute("aria-selected"))).toEqual([
"true",
"false",
"false",
]);
});
it.each([
["OpenAI Python SDK", "import openai"],
["LlamaIndex", "from llama_index.llms import AzureOpenAI"],
["Langchain Py", "from langchain.chat_models import ChatOpenAI"],
])("selecting %s shows its snippet wired to the base url", async (tabName, marker) => {
const proxyUrl = "https://proxy.litellm.test";
const user = userEvent.setup();
render(<APIReferenceView proxySettings={{ PROXY_BASE_URL: proxyUrl }} />);
await user.click(screen.getByRole("tab", { name: tabName }));
expect(screen.getByRole("tab", { name: tabName }).getAttribute("aria-selected")).toBe("true");
const selectedPanel = screen.getByRole("tabpanel");
expect(selectedPanel.textContent).toContain(marker);
expect(selectedPanel.textContent).toContain(proxyUrl);
});
});

View file

@ -1,7 +1,7 @@
"use client";
import React from "react";
import { Text, Tab, TabGroup, TabList, TabPanel, TabPanels, Grid } from "@tremor/react";
import CodeBlock from "@/components/CodeBlock";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import DocLink from "./DocLink";
interface ApiRefProps {
@ -21,33 +21,35 @@ const APIReferenceView: React.FC<ApiRefProps> = ({ proxySettings }) => {
}
return (
<>
<Grid className="gap-2 p-8 h-[80vh] w-full mt-2">
<div className="mb-5">
{/* Header row with Docs link on the right */}
<div className="flex items-center justify-between">
<p className="text-2xl text-tremor-content-strong dark:text-dark-tremor-content-strong font-semibold">
OpenAI Compatible Proxy: API Reference
</p>
<DocLink className="ml-3 shrink-0" href="https://docs.litellm.ai/docs/proxy/user_keys" />
</div>
<div className="grid grid-cols-1 gap-2 p-8 h-[80vh] w-full mt-2">
<div className="mb-5">
{/* Header row with Docs link on the right */}
<div className="flex items-center justify-between">
<h1 className="text-2xl font-semibold text-foreground">OpenAI Compatible Proxy: API Reference</h1>
<DocLink className="ml-3 shrink-0" href="https://docs.litellm.ai/docs/proxy/user_keys" />
</div>
<Text className="mt-2 mb-2">
LiteLLM is OpenAI Compatible. This means your API Key works with the OpenAI SDK. Just replace the base_url
to point to your litellm proxy. Example Below{" "}
</Text>
<p className="mt-2 mb-2 text-sm text-muted-foreground">
LiteLLM is OpenAI Compatible. This means your API Key works with the OpenAI SDK. Just replace the base_url to
point to your litellm proxy. Example Below{" "}
</p>
<TabGroup>
<TabList>
<Tab>OpenAI Python SDK</Tab>
<Tab>LlamaIndex</Tab>
<Tab>Langchain Py</Tab>
</TabList>
<TabPanels>
<TabPanel>
<CodeBlock
language="python"
code={`import openai
<Tabs defaultValue="openai">
<TabsList variant="line" className="border-b rounded-none w-full justify-start h-auto p-0">
<TabsTrigger value="openai" className="rounded-none px-4 py-2 flex-none">
OpenAI Python SDK
</TabsTrigger>
<TabsTrigger value="llamaindex" className="rounded-none px-4 py-2 flex-none">
LlamaIndex
</TabsTrigger>
<TabsTrigger value="langchain" className="rounded-none px-4 py-2 flex-none">
Langchain Py
</TabsTrigger>
</TabsList>
<TabsContent value="openai">
<CodeBlock
language="python"
code={`import openai
client = openai.OpenAI(
api_key="your_api_key",
base_url="${base_url}" # LiteLLM Proxy is OpenAI compatible, Read More: https://docs.litellm.ai/docs/proxy/user_keys
@ -64,13 +66,13 @@ response = client.chat.completions.create(
)
print(response)`}
/>
</TabPanel>
/>
</TabsContent>
<TabPanel>
<CodeBlock
language="python"
code={`import os, dotenv
<TabsContent value="llamaindex">
<CodeBlock
language="python"
code={`import os, dotenv
from llama_index.llms import AzureOpenAI
from llama_index.embeddings import AzureOpenAIEmbedding
@ -98,13 +100,13 @@ index = VectorStoreIndex.from_documents(documents, service_context=service_conte
query_engine = index.as_query_engine()
response = query_engine.query("What did the author do growing up?")
print(response)`}
/>
</TabPanel>
/>
</TabsContent>
<TabPanel>
<CodeBlock
language="python"
code={`from langchain.chat_models import ChatOpenAI
<TabsContent value="langchain">
<CodeBlock
language="python"
code={`from langchain.chat_models import ChatOpenAI
from langchain.prompts.chat import (
ChatPromptTemplate,
HumanMessagePromptTemplate,
@ -129,13 +131,11 @@ messages = [
response = chat(messages)
print(response)`}
/>
</TabPanel>
</TabPanels>
</TabGroup>
</div>
</Grid>
</>
/>
</TabsContent>
</Tabs>
</div>
</div>
);
};

View file

@ -211,6 +211,7 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
auth_type={mcpServer.auth_type}
oauth2_flow={mcpServer.oauth2_flow}
delegate_auth_to_upstream={mcpServer.delegate_auth_to_upstream}
dcr_bridge={mcpServer.dcr_bridge}
tokenUrl={mcpServer.token_url}
userRole={userRole}
userID={userID}

View file

@ -17,8 +17,12 @@ vi.mock("@/utils/mcpTokenStore", () => ({
removeToken: vi.fn(),
}));
const { toolsOAuthFlowSpy } = vi.hoisted(() => ({
toolsOAuthFlowSpy: vi.fn(() => ({ startOAuthFlow: vi.fn(), status: "idle", error: null })),
}));
vi.mock("@/hooks/useToolsOAuthFlow", () => ({
useToolsOAuthFlow: () => ({ startOAuthFlow: vi.fn(), status: "idle", error: null }),
useToolsOAuthFlow: toolsOAuthFlowSpy,
}));
vi.mock("@/hooks/useUserMcpOAuthFlow", () => ({
@ -54,6 +58,27 @@ const credStatus = (overrides: Record<string, unknown> = {}) => ({
...overrides,
});
describe("MCPToolsViewer gatewayMintsClient wiring", () => {
// Pins the call site (not just the helper): the viewer must pass the bridge-AWARE
// gatewayMintsClientFor value to useToolsOAuthFlow, so the browser skips its own register exactly
// when the gateway mints. The oauth_delegate + dcr_bridge cell is the regression guard: with the
// old bridge-blind predicate it would have passed true here and dead-ended.
beforeEach(() => toolsOAuthFlowSpy.mockClear());
it.each([
{ auth_type: "true_passthrough", dcr_bridge: true, gatewayMintsClient: true },
{ auth_type: "true_passthrough", dcr_bridge: false, gatewayMintsClient: true },
{ auth_type: "oauth_delegate", dcr_bridge: false, gatewayMintsClient: true },
{ auth_type: "oauth_delegate", dcr_bridge: true, gatewayMintsClient: false },
])(
"passes gatewayMintsClient=$gatewayMintsClient for $auth_type dcr_bridge=$dcr_bridge",
({ auth_type, dcr_bridge, gatewayMintsClient }) => {
renderViewer({ auth_type, dcr_bridge, tokenUrl: null });
expect(toolsOAuthFlowSpy).toHaveBeenCalledWith(expect.objectContaining({ gatewayMintsClient }));
},
);
});
describe("MCPToolsViewer auth gate routing", () => {
beforeEach(() => {
vi.mocked(listMCPTools).mockReset().mockResolvedValue({ tools: [], error: null });

View file

@ -4,6 +4,7 @@ import { ToolTestPanel } from "./ToolTestPanel";
import { resolveLogoSrc } from "@/lib/assetPaths";
import {
isClientForwardedTokenMode,
gatewayMintsClientFor,
MCPTool,
MCPToolsViewerProps,
MCPContent,
@ -28,6 +29,7 @@ const MCPToolsViewer = ({
auth_type,
oauth2_flow,
delegate_auth_to_upstream,
dcr_bridge,
userRole,
userID,
serverAlias,
@ -76,6 +78,7 @@ const MCPToolsViewer = ({
serverId,
serverAlias,
userId: userID,
gatewayMintsClient: gatewayMintsClientFor({ auth_type, dcr_bridge }),
onSuccess: setOauthToken,
});

View file

@ -0,0 +1,202 @@
import React from "react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { renderWithProviders } from "../../../../../tests/test-utils";
import UsagePage from "./usage";
const networking = vi.hoisted(() => ({
adminSpendLogsCall: vi.fn(),
adminTopKeysCall: vi.fn(),
adminTopModelsCall: vi.fn(),
adminTopEndUsersCall: vi.fn(),
teamSpendLogsCall: vi.fn(),
tagsSpendLogsCall: vi.fn(),
allTagNamesCall: vi.fn(),
adminspendByProvider: vi.fn(),
adminGlobalActivity: vi.fn(),
adminGlobalActivityPerModel: vi.fn(),
getProxyUISettings: vi.fn(),
modelAvailableCall: vi.fn(),
keyInfoV1Call: vi.fn(),
}));
vi.mock("@/components/networking", () => networking);
vi.mock("../../../../components/networking", () => networking);
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
default: () => ({
accessToken: "sk-test",
token: "tok",
userRole: "Admin",
userId: "u1",
premiumUser: true,
}),
}));
const UNLIMITED_SETTINGS = { DISABLE_EXPENSIVE_DB_QUERIES: false, NUM_SPEND_LOGS_ROWS: 10 };
const renderUsage = (overrides: Partial<React.ComponentProps<typeof UsagePage>> = {}) =>
renderWithProviders(
<UsagePage
accessToken="sk-test"
token="tok"
userRole="Admin"
userID="u1"
keys={null}
premiumUser={true}
{...overrides}
/>,
);
beforeEach(() => {
vi.clearAllMocks();
networking.getProxyUISettings.mockResolvedValue(UNLIMITED_SETTINGS);
networking.adminSpendLogsCall.mockResolvedValue([{ date: "2026-07-01", spend: 12.5 }]);
networking.adminTopKeysCall.mockResolvedValue([
{ api_key: "sk-abcdefghijk", key_alias: "prod-key", total_spend: 9.5 },
]);
networking.adminTopModelsCall.mockResolvedValue([{ model: "gpt-5.1", total_spend: 7.25 }]);
networking.adminTopEndUsersCall.mockResolvedValue([
{ end_user: "customer-alpha", total_spend: 3.5, total_count: 42 },
]);
networking.teamSpendLogsCall.mockResolvedValue({
daily_spend: [{ date: "2026-07-01", "team-a": 5 }],
teams: ["team-a"],
total_spend_per_team: [{ team_id: "team-a", total_spend: 5 }],
});
networking.tagsSpendLogsCall.mockResolvedValue({ spend_per_tag: [{ name: "prod", spend: 4 }] });
networking.allTagNamesCall.mockResolvedValue({ tag_names: ["prod", "staging"] });
networking.adminspendByProvider.mockResolvedValue([{ provider: "openai", spend: 6.75 }]);
networking.adminGlobalActivity.mockResolvedValue({
sum_api_requests: 120,
sum_total_tokens: 4500,
daily_data: [{ date: "2026-07-01", api_requests: 120, total_tokens: 4500 }],
});
networking.adminGlobalActivityPerModel.mockResolvedValue([]);
networking.modelAvailableCall.mockResolvedValue({ data: [] });
networking.keyInfoV1Call.mockResolvedValue({ info: {} });
});
describe("old usage page", () => {
describe("when the proxy has disabled expensive DB queries", () => {
beforeEach(() => {
networking.getProxyUISettings.mockResolvedValue({
DISABLE_EXPENSIVE_DB_QUERIES: true,
NUM_SPEND_LOGS_ROWS: 2500000,
});
});
it("shows the database query limit warning instead of the usage dashboard", async () => {
renderUsage();
expect(await screen.findByText("Database Query Limit Reached")).toBeInTheDocument();
expect(screen.getByText(/SpendLogs in DB has/)).toHaveTextContent("2500000");
expect(screen.getByText(/Please follow our guide to view usage when SpendLogs has more than 1M rows/i));
expect(screen.queryByRole("tab", { name: "All Up" })).not.toBeInTheDocument();
});
it("links to the cost tracking guide in a new tab", async () => {
renderUsage();
const link = await screen.findByRole("link", { name: "View Usage Guide" });
expect(link).toHaveAttribute("href", "https://docs.litellm.ai/docs/proxy/cost_tracking");
expect(link).toHaveAttribute("target", "_blank");
});
it("skips every expensive usage query", async () => {
renderUsage();
await screen.findByText("Database Query Limit Reached");
await waitFor(() => expect(networking.getProxyUISettings).toHaveBeenCalled());
expect(networking.adminSpendLogsCall).not.toHaveBeenCalled();
expect(networking.adminspendByProvider).not.toHaveBeenCalled();
expect(networking.adminTopKeysCall).not.toHaveBeenCalled();
expect(networking.adminTopModelsCall).not.toHaveBeenCalled();
expect(networking.adminGlobalActivity).not.toHaveBeenCalled();
expect(networking.adminGlobalActivityPerModel).not.toHaveBeenCalled();
expect(networking.teamSpendLogsCall).not.toHaveBeenCalled();
expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled();
expect(networking.tagsSpendLogsCall).not.toHaveBeenCalled();
});
});
describe("as an admin", () => {
it("renders the admin tabs", async () => {
renderUsage();
expect(await screen.findByRole("tab", { name: "All Up" })).toBeInTheDocument();
expect(screen.getByRole("tab", { name: "Team Based Usage" })).toBeInTheDocument();
expect(screen.getByRole("tab", { name: "Customer Usage" })).toBeInTheDocument();
expect(screen.getByRole("tab", { name: "Tag Based Usage" })).toBeInTheDocument();
});
it("renders the cost panel cards", async () => {
renderUsage();
expect(await screen.findByText("Monthly Spend")).toBeInTheDocument();
expect(screen.getByText("Top Virtual Keys")).toBeInTheDocument();
expect(screen.getByText("Top Models")).toBeInTheDocument();
expect(screen.getByText("Spend by Provider")).toBeInTheDocument();
});
it("lists spend by provider in a table", async () => {
renderUsage();
const providerCell = await screen.findByText("openai");
const row = providerCell.closest("tr");
expect(row).not.toBeNull();
expect(within(row as HTMLElement).getByText("$6.75")).toBeInTheDocument();
expect(screen.getByRole("columnheader", { name: "Provider" })).toBeInTheDocument();
});
it("shows the customer usage table when its tab is selected", async () => {
const user = userEvent.setup();
renderUsage();
await user.click(await screen.findByRole("tab", { name: "Customer Usage" }));
const customerCell = await screen.findByText("customer-alpha");
const row = customerCell.closest("tr");
expect(row).not.toBeNull();
expect(within(row as HTMLElement).getByText("$3.50")).toBeInTheDocument();
expect(within(row as HTMLElement).getByText("42")).toBeInTheDocument();
expect(screen.getByRole("columnheader", { name: "Total Events" })).toBeInTheDocument();
});
it("shows the tag spend panel when its tab is selected", async () => {
const user = userEvent.setup();
renderUsage();
await user.click(await screen.findByRole("tab", { name: "Tag Based Usage" }));
expect(await screen.findByText("Spend Per Tag")).toBeInTheDocument();
});
it("shows the team spend panel when its tab is selected", async () => {
const user = userEvent.setup();
renderUsage();
await user.click(await screen.findByRole("tab", { name: "Team Based Usage" }));
expect(await screen.findByText("Total Spend Per Team")).toBeInTheDocument();
expect(screen.getByText("Daily Spend Per Team")).toBeInTheDocument();
});
});
describe("as a non-admin", () => {
it("renders only the All Up tab and skips admin-only queries", async () => {
renderUsage({ userRole: "Internal User" });
expect(await screen.findByRole("tab", { name: "All Up" })).toBeInTheDocument();
expect(screen.queryByRole("tab", { name: "Team Based Usage" })).not.toBeInTheDocument();
expect(screen.queryByRole("tab", { name: "Customer Usage" })).not.toBeInTheDocument();
expect(screen.queryByRole("tab", { name: "Tag Based Usage" })).not.toBeInTheDocument();
await waitFor(() => expect(networking.adminSpendLogsCall).toHaveBeenCalled());
expect(networking.teamSpendLogsCall).not.toHaveBeenCalled();
expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled();
});
});
});

View file

@ -1,40 +1,26 @@
import {
BarChart,
BarList,
Card,
Title,
Table,
TableHead,
TableHeaderCell,
TableRow,
TableCell,
TableBody,
Subtitle,
} from "@tremor/react";
import React, { useState, useEffect } from "react";
import ViewUserSpend from "@/components/view_user_spend";
import { ProxySettings } from "@/components/user_dashboard";
import UsageDatePicker from "@/components/shared/usage_date_picker";
import { Button } from "@/components/ui/button";
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import {
Grid,
Col,
Text,
TabPanel,
TabPanels,
TabGroup,
TabList,
Tab,
Select,
SelectItem,
DateRangePickerValue,
DonutChart,
AreaChart,
Button,
MultiSelect,
MultiSelectItem,
} from "@tremor/react";
Combobox,
ComboboxChip,
ComboboxChips,
ComboboxChipsInput,
ComboboxContent,
ComboboxEmpty,
ComboboxItem,
ComboboxList,
ComboboxValue,
} from "@/components/ui/combobox";
import { Meter, MeterIndicator, MeterTrack } from "@/components/ui/meter";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { AreaChart, BarChart, DonutChart } from "@/components/shared/charts";
import {
adminSpendLogsCall,
@ -68,69 +54,41 @@ interface GlobalActivityData {
daily_data: { date: string; api_requests: number; total_tokens: number }[];
}
type CustomTooltipTypeBar = {
payload: any;
active: boolean | undefined;
label: any;
};
type UsageDateRange = { from?: Date; to?: Date };
const customTooltip = (props: CustomTooltipTypeBar) => {
const { payload, active } = props;
if (!active || !payload) return null;
type TeamSpendTotal = { name: string; value: number };
const value = payload[0].payload;
const date = value["startTime"];
const model_values = value["models"];
const entries: [string, number][] = Object.entries(model_values).map(([key, value]) => [key, value as number]);
type TagOption = { value: string; label: string; disabled: boolean };
entries.sort((a, b) => b[1] - a[1]);
const topEntries = entries.slice(0, 5);
return (
<div className="w-56 rounded-tremor-default border border-tremor-border bg-tremor-background p-2 text-tremor-default shadow-tremor-dropdown">
{date}
{topEntries.map(([key, value]) => (
<div key={key} className="flex flex-1 space-x-10">
<div className="p-2">
<p className="text-tremor-content text-xs">
{key}
{":"}
<span className="text-xs text-tremor-content-emphasis">
{" "}
{value ? `$${formatNumberWithCommas(value, 2)}` : ""}
</span>
</p>
</div>
</div>
))}
</div>
);
};
function getTopKeys(data: Array<{ [key: string]: unknown }>): any[] {
const spendKeys: { key: string; spend: unknown }[] = [];
data.forEach((dict) => {
Object.entries(dict).forEach(([key, value]) => {
if (key !== "spend" && key !== "startTime" && key !== "models" && key !== "users") {
spendKeys.push({ key, spend: value });
}
});
});
spendKeys.sort((a, b) => Number(b.spend) - Number(a.spend));
const topKeys = spendKeys.slice(0, 5).map((k) => k.key);
return topKeys;
}
type DataDict = { [key: string]: unknown };
type UserData = { user_id: string; spend: number };
const ALL_TAGS = "all-tags";
const isAdminOrAdminViewer = (role: string | null): boolean => {
if (role === null) return false;
return role === "Admin" || role === "Admin Viewer";
};
const TeamSpendBarList: React.FC<{ data: TeamSpendTotal[] }> = ({ data }) => {
const max = Math.max(0, ...data.map((team) => team.value));
return (
<div className="flex flex-col gap-3">
{data.map((team) => (
<div key={team.name} className="flex items-center gap-4">
<p className="w-1/3 truncate text-sm text-foreground">{team.name}</p>
<Meter value={team.value} max={max === 0 ? 1 : max} className="flex-1">
<MeterTrack>
<MeterIndicator />
</MeterTrack>
</Meter>
<p className="w-24 shrink-0 text-right text-sm tabular-nums text-foreground">
{formatNumberWithCommas(team.value, 2)}
</p>
</div>
))}
</div>
);
};
const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, userID, keys, premiumUser }) => {
const currentDate = new Date();
const [keySpendData, setKeySpendData] = useState<any[]>([]);
@ -141,13 +99,13 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
const [topTagsData, setTopTagsData] = useState<any[]>([]);
const [allTagNames, setAllTagNames] = useState<string[]>([]);
const [uniqueTeamIds, setUniqueTeamIds] = useState<any[]>([]);
const [totalSpendPerTeam, setTotalSpendPerTeam] = useState<any[]>([]);
const [totalSpendPerTeam, setTotalSpendPerTeam] = useState<TeamSpendTotal[]>([]);
const [spendByProvider, setSpendByProvider] = useState<any[]>([]);
const [globalActivity, setGlobalActivity] = useState<GlobalActivityData>({} as GlobalActivityData);
const [globalActivityPerModel, setGlobalActivityPerModel] = useState<any[]>([]);
const [selectedKeyID, setSelectedKeyID] = useState<string | null>("");
const [selectedTags, setSelectedTags] = useState<string[]>(["all-tags"]);
const [dateValue, setDateValue] = useState<DateRangePickerValue>({
const [selectedKeyToken, setSelectedKeyToken] = useState<string | null>(null);
const [selectedTags, setSelectedTags] = useState<string[]>([ALL_TAGS]);
const [dateValue, setDateValue] = useState<UsageDateRange>({
from: new Date(Date.now() - 7 * 24 * 60 * 60 * 1000),
to: new Date(),
});
@ -160,6 +118,21 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
let startTime = formatDate(firstDay);
let endTime = formatDate(lastDay);
const selectableKeys: { token: string; alias: string }[] = (keys ?? [])
.filter((key: any) => key && typeof key["key_alias"] === "string" && key["key_alias"].length > 0)
.map((key: any) => ({ token: String(key["token"]), alias: String(key["key_alias"]) }));
const tagOptions: TagOption[] = [
{ value: ALL_TAGS, label: "All Tags", disabled: false },
...allTagNames
.filter((tag) => tag !== ALL_TAGS)
.map((tag) => ({
value: tag,
label: premiumUser ? tag : `✨ ${tag} (Enterprise only Feature)`,
disabled: !premiumUser,
})),
];
function valueFormatterNumbers(number: number) {
const formatter = new Intl.NumberFormat("en-US", {
maximumFractionDigits: 0,
@ -405,7 +378,7 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
setUniqueTeamIds(teamSpend.teams);
return teamSpend.total_spend_per_team.map((tspt: any) => ({
name: tspt["team_id"] || "",
value: formatNumberWithCommas(tspt["total_spend"] || 0, 2),
value: Number(tspt["total_spend"] || 0),
}));
},
setTotalSpendPerTeam,
@ -524,223 +497,252 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
if (proxySettings?.DISABLE_EXPENSIVE_DB_QUERIES) {
return (
<div style={{ width: "100%" }} className="p-8">
<div className="w-full p-8">
<Card>
<Title>Database Query Limit Reached</Title>
<Text className="mt-4">
SpendLogs in DB has {proxySettings.NUM_SPEND_LOGS_ROWS} rows.
<br></br>
Please follow our guide to view usage when SpendLogs has more than 1M rows.
</Text>
<Button className="mt-4">
<a href="https://docs.litellm.ai/docs/proxy/cost_tracking" target="_blank">
View Usage Guide
</a>
</Button>
<CardHeader>
<CardTitle>Database Query Limit Reached</CardTitle>
</CardHeader>
<CardContent className="flex flex-col items-start gap-4">
<p className="text-sm text-muted-foreground">
SpendLogs in DB has {proxySettings.NUM_SPEND_LOGS_ROWS} rows.
<br></br>
Please follow our guide to view usage when SpendLogs has more than 1M rows.
</p>
<Button
render={
<a href="https://docs.litellm.ai/docs/proxy/cost_tracking" target="_blank" rel="noreferrer">
View Usage Guide
</a>
}
/>
</CardContent>
</Card>
</div>
);
}
return (
<div style={{ width: "100%" }} className="p-8">
<TabGroup>
<TabList className="mt-2">
<Tab>All Up</Tab>
<div className="w-full p-8">
<Tabs defaultValue="all-up">
<TabsList variant="line" className="mt-2">
<TabsTrigger value="all-up">All Up</TabsTrigger>
{isAdminOrAdminViewer(userRole) ? (
{isAdminOrAdminViewer(userRole) && (
<>
<Tab>Team Based Usage</Tab>
<Tab>Customer Usage</Tab>
<Tab>Tag Based Usage</Tab>
</>
) : (
<>
<div></div>
<TabsTrigger value="team-based-usage">Team Based Usage</TabsTrigger>
<TabsTrigger value="customer-usage">Customer Usage</TabsTrigger>
<TabsTrigger value="tag-based-usage">Tag Based Usage</TabsTrigger>
</>
)}
</TabList>
<TabPanels>
<TabPanel>
<TabGroup>
<TabList variant="solid" className="mt-1">
<Tab>Cost</Tab>
<Tab>Activity</Tab>
</TabList>
<TabPanels>
<TabPanel>
<Grid numItems={2} className="gap-2 h-screen w-full">
<Col numColSpan={2}>
<Text className="text-tremor-default text-tremor-content dark:text-dark-tremor-content mb-2 mt-2 text-lg">
Project Spend {new Date().toLocaleString("default", { month: "long" })} 1 -{" "}
{new Date(new Date().getFullYear(), new Date().getMonth() + 1, 0).getDate()}
</Text>
<ViewUserSpend userSpend={totalMonthlySpend} selectedTeam={null} userMaxBudget={null} />
</Col>
<Col numColSpan={2}>
<Card>
<Title>Monthly Spend</Title>
<BarChart
data={keySpendData}
</TabsList>
<TabsContent value="all-up">
<Tabs defaultValue="cost">
<TabsList className="mt-1">
<TabsTrigger value="cost">Cost</TabsTrigger>
<TabsTrigger value="activity">Activity</TabsTrigger>
</TabsList>
<TabsContent value="cost">
<div className="grid h-screen w-full grid-cols-2 gap-2">
<div className="col-span-2">
<p className="mt-2 mb-2 text-lg text-muted-foreground">
Project Spend {new Date().toLocaleString("default", { month: "long" })} 1 -{" "}
{new Date(new Date().getFullYear(), new Date().getMonth() + 1, 0).getDate()}
</p>
<ViewUserSpend userSpend={totalMonthlySpend} selectedTeam={null} userMaxBudget={null} />
</div>
<div className="col-span-2">
<Card>
<CardHeader>
<CardTitle>Monthly Spend</CardTitle>
</CardHeader>
<CardContent>
<BarChart
data={keySpendData}
index="date"
categories={["spend"]}
colors={["cyan"]}
valueFormatter={valueFormatter}
yAxisWidth={100}
tickGap={5}
/>
</CardContent>
</Card>
</div>
<div className="col-span-1">
<Card className="h-full">
<CardHeader>
<CardTitle>Top Virtual Keys</CardTitle>
</CardHeader>
<CardContent>
<TopKeyView topKeys={topKeys} teams={null} topKeysLimit={5} setTopKeysLimit={() => {}} />
</CardContent>
</Card>
</div>
<div className="col-span-1">
<Card className="h-full">
<CardHeader>
<CardTitle>Top Models</CardTitle>
</CardHeader>
<CardContent>
<BarChart
className="mt-4 h-40"
data={topModels}
index="key"
categories={["spend"]}
colors={["cyan"]}
yAxisWidth={200}
layout="vertical"
showXAxis={false}
showLegend={false}
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
/>
</CardContent>
</Card>
</div>
<div className="col-span-1"></div>
<div className="col-span-2">
<Card className="mb-2">
<CardHeader>
<CardTitle>Spend by Provider</CardTitle>
</CardHeader>
<CardContent>
<div className="grid grid-cols-2">
<div className="col-span-1">
<DonutChart
className="mt-4 h-40"
variant="pie"
data={spendByProvider}
index="provider"
category="spend"
colors={["cyan"]}
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
/>
</div>
<div className="col-span-1">
<Table>
<TableHeader>
<TableRow>
<TableHead>Provider</TableHead>
<TableHead>Spend</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{spendByProvider.map((provider) => (
<TableRow key={provider.provider}>
<TableCell>{provider.provider}</TableCell>
<TableCell>
<MoneyCell value={provider.spend} decimals={2} />
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
</div>
</CardContent>
</Card>
</div>
</div>
</TabsContent>
<TabsContent value="activity">
<div className="grid h-[75vh] w-full grid-cols-1 gap-2">
<Card>
<CardHeader>
<CardTitle>All Up</CardTitle>
</CardHeader>
<CardContent>
<div className="grid grid-cols-2">
<div>
<p className="text-[15px] font-normal text-muted-foreground">
API Requests {valueFormatterNumbers(globalActivity.sum_api_requests)}
</p>
<AreaChart
className="h-40"
data={globalActivity.daily_data}
valueFormatter={valueFormatterNumbers}
index="date"
categories={["spend"]}
colors={["cyan"]}
valueFormatter={valueFormatter}
yAxisWidth={100}
tickGap={5}
// customTooltip={customTooltip}
categories={["api_requests"]}
/>
</Card>
</Col>
<Col numColSpan={1}>
<Card className="h-full">
<Title>Top Virtual Keys</Title>
<TopKeyView topKeys={topKeys} teams={null} topKeysLimit={5} setTopKeysLimit={() => {}} />
</Card>
</Col>
<Col numColSpan={1}>
<Card className="h-full">
<Title>Top Models</Title>
</div>
<div>
<p className="text-[15px] font-normal text-muted-foreground">
Tokens {valueFormatterNumbers(globalActivity.sum_total_tokens)}
</p>
<BarChart
className="mt-4 h-40"
data={topModels}
index="key"
categories={["spend"]}
className="h-40"
data={globalActivity.daily_data}
valueFormatter={valueFormatterNumbers}
index="date"
colors={["cyan"]}
yAxisWidth={200}
layout="vertical"
showXAxis={false}
showLegend={false}
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
categories={["total_tokens"]}
/>
</Card>
</Col>
<Col numColSpan={1}></Col>
<Col numColSpan={2}>
<Card className="mb-2">
<Title>Spend by Provider</Title>
<>
<Grid numItems={2}>
<Col numColSpan={1}>
<DonutChart
className="mt-4 h-40"
variant="pie"
data={spendByProvider}
index="provider"
category="spend"
colors={["cyan"]}
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
/>
</Col>
<Col numColSpan={1}>
<Table>
<TableHead>
<TableRow>
<TableHeaderCell>Provider</TableHeaderCell>
<TableHeaderCell>Spend</TableHeaderCell>
</TableRow>
</TableHead>
<TableBody>
{spendByProvider.map((provider) => (
<TableRow key={provider.provider}>
<TableCell>{provider.provider}</TableCell>
<TableCell>
<MoneyCell value={provider.spend} decimals={2} />
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</Col>
</Grid>
</>
</Card>
</Col>
</Grid>
</TabPanel>
<TabPanel>
<Grid numItems={1} className="gap-2 h-[75vh] w-full">
<Card>
<Title>All Up</Title>
<Grid numItems={2}>
<Col>
<Subtitle style={{ fontSize: "15px", fontWeight: "normal", color: "#535452" }}>
</div>
</div>
</CardContent>
</Card>
{globalActivityPerModel.map((globalActivity, index) => (
<Card key={index}>
<CardHeader>
<CardTitle>{globalActivity.model}</CardTitle>
</CardHeader>
<CardContent>
<div className="grid grid-cols-2">
<div>
<p className="text-[15px] font-normal text-muted-foreground">
API Requests {valueFormatterNumbers(globalActivity.sum_api_requests)}
</Subtitle>
</p>
<AreaChart
className="h-40"
data={globalActivity.daily_data}
valueFormatter={valueFormatterNumbers}
index="date"
colors={["cyan"]}
categories={["api_requests"]}
valueFormatter={valueFormatterNumbers}
/>
</Col>
<Col>
<Subtitle style={{ fontSize: "15px", fontWeight: "normal", color: "#535452" }}>
</div>
<div>
<p className="text-[15px] font-normal text-muted-foreground">
Tokens {valueFormatterNumbers(globalActivity.sum_total_tokens)}
</Subtitle>
</p>
<BarChart
className="h-40"
data={globalActivity.daily_data}
valueFormatter={valueFormatterNumbers}
index="date"
colors={["cyan"]}
categories={["total_tokens"]}
valueFormatter={valueFormatterNumbers}
/>
</Col>
</Grid>
</Card>
</div>
</div>
</CardContent>
</Card>
))}
</div>
</TabsContent>
</Tabs>
</TabsContent>
<>
{globalActivityPerModel.map((globalActivity, index) => (
<Card key={index}>
<Title>{globalActivity.model}</Title>
<Grid numItems={2}>
<Col>
<Subtitle style={{ fontSize: "15px", fontWeight: "normal", color: "#535452" }}>
API Requests {valueFormatterNumbers(globalActivity.sum_api_requests)}
</Subtitle>
<AreaChart
className="h-40"
data={globalActivity.daily_data}
index="date"
colors={["cyan"]}
categories={["api_requests"]}
valueFormatter={valueFormatterNumbers}
/>
</Col>
<Col>
<Subtitle style={{ fontSize: "15px", fontWeight: "normal", color: "#535452" }}>
Tokens {valueFormatterNumbers(globalActivity.sum_total_tokens)}
</Subtitle>
<BarChart
className="h-40"
data={globalActivity.daily_data}
index="date"
colors={["cyan"]}
categories={["total_tokens"]}
valueFormatter={valueFormatterNumbers}
/>
</Col>
</Grid>
</Card>
))}
</>
</Grid>
</TabPanel>
</TabPanels>
</TabGroup>
</TabPanel>
<TabPanel>
<Grid numItems={2} className="gap-2 h-[75vh] w-full">
<Col numColSpan={2}>
<Card className="mb-2">
<Title>Total Spend Per Team</Title>
<BarList data={totalSpendPerTeam} />
</Card>
<Card>
<Title>Daily Spend Per Team</Title>
<TabsContent value="team-based-usage">
<div className="grid h-[75vh] w-full grid-cols-2 gap-2">
<div className="col-span-2">
<Card className="mb-2">
<CardHeader>
<CardTitle>Total Spend Per Team</CardTitle>
</CardHeader>
<CardContent>
<TeamSpendBarList data={totalSpendPerTeam} />
</CardContent>
</Card>
<Card>
<CardHeader>
<CardTitle>Daily Spend Per Team</CardTitle>
</CardHeader>
<CardContent>
<BarChart
className="h-72"
data={teamSpendData}
@ -750,178 +752,161 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
yAxisWidth={80}
stack={true}
/>
</Card>
</Col>
<Col numColSpan={2}></Col>
</Grid>
</TabPanel>
<TabPanel>
<p className="mb-2 text-gray-500 italic text-[12px]">
Customers of your LLM API calls. Tracked when a `user` param is passed in your LLM calls{" "}
<a className="text-blue-500" href="https://docs.litellm.ai/docs/proxy/users" target="_blank">
docs here
</a>
</p>
<Grid numItems={2}>
<Col>
<UsageDatePicker
value={dateValue}
onValueChange={(value) => {
setDateValue(value);
updateEndUserData(value.from, value.to, null);
}}
/>
</Col>
<Col>
<Text>Select Key</Text>
<Select defaultValue="all-keys">
<SelectItem
key="all-keys"
value="all-keys"
onClick={() => {
updateEndUserData(dateValue.from, dateValue.to, null);
}}
>
All Keys
</SelectItem>
{keys?.map((key: any, index: number) => {
if (key && key["key_alias"] !== null && key["key_alias"].length > 0) {
return (
<SelectItem
key={index}
value={String(index)}
onClick={() => {
updateEndUserData(dateValue.from, dateValue.to, key["token"]);
}}
>
{key["key_alias"]}
</SelectItem>
);
}
return null; // Add this line to handle the case when the condition is not met
})}
</Select>
</Col>
</Grid>
</CardContent>
</Card>
</div>
</div>
</TabsContent>
<Card className="mt-4">
<Table className="max-h-[70vh] min-h-[500px]">
<TableHead>
<TableRow>
<TableHeaderCell>Customer</TableHeaderCell>
<TableHeaderCell>Spend</TableHeaderCell>
<TableHeaderCell>Total Events</TableHeaderCell>
</TableRow>
</TableHead>
<TableBody>
{topUsers?.map((user: any, index: number) => (
<TableRow key={index}>
<TableCell>{user.end_user}</TableCell>
<TableCell>
<MoneyCell value={user.total_spend} decimals={2} />
</TableCell>
<TableCell>{user.total_count}</TableCell>
</TableRow>
<TabsContent value="customer-usage">
<p className="mb-2 text-[12px] text-muted-foreground italic">
Customers of your LLM API calls. Tracked when a `user` param is passed in your LLM calls{" "}
<a
className="text-primary"
href="https://docs.litellm.ai/docs/proxy/users"
target="_blank"
rel="noreferrer"
>
docs here
</a>
</p>
<div className="grid grid-cols-2">
<div>
<UsageDatePicker
value={dateValue}
onValueChange={(value) => {
setDateValue(value);
updateEndUserData(value.from, value.to, null);
}}
/>
</div>
<div>
<p className="text-sm text-muted-foreground">Select Key</p>
<Select
value={selectedKeyToken}
onValueChange={(value: string | null) => {
setSelectedKeyToken(value);
updateEndUserData(dateValue.from, dateValue.to, value);
}}
>
<SelectTrigger className="w-full">
<SelectValue placeholder="All Keys">
{(token: string | null) => selectableKeys.find((key) => key.token === token)?.alias ?? "All Keys"}
</SelectValue>
</SelectTrigger>
<SelectContent>
<SelectItem value={null}>All Keys</SelectItem>
{selectableKeys.map((key) => (
<SelectItem key={key.token} value={key.token}>
{key.alias}
</SelectItem>
))}
</TableBody>
</Table>
</Card>
</TabPanel>
<TabPanel>
<Grid numItems={2}>
<Col numColSpan={1}>
<UsageDatePicker
className="mb-4"
value={dateValue}
onValueChange={(value) => {
setDateValue(value);
updateTagSpendData(value.from, value.to);
}}
/>
</Col>
</SelectContent>
</Select>
</div>
</div>
<Col>
{premiumUser ? (
<div>
<MultiSelect value={selectedTags} onValueChange={(value) => setSelectedTags(value as string[])}>
<MultiSelectItem
key={"all-tags"}
value={"all-tags"}
onClick={() => setSelectedTags(["all-tags"])}
>
All Tags
</MultiSelectItem>
{allTagNames &&
allTagNames
.filter((tag) => tag !== "all-tags")
.map((tag: any, index: number) => {
return (
<MultiSelectItem key={tag} value={String(tag)}>
{tag}
</MultiSelectItem>
);
})}
</MultiSelect>
</div>
) : (
<div>
<MultiSelect value={selectedTags} onValueChange={(value) => setSelectedTags(value as string[])}>
<MultiSelectItem
key={"all-tags"}
value={"all-tags"}
onClick={() => setSelectedTags(["all-tags"])}
>
All Tags
</MultiSelectItem>
{allTagNames &&
allTagNames
.filter((tag) => tag !== "all-tags")
.map((tag: any, index: number) => {
return (
<SelectItem
key={tag}
value={String(tag)}
// @ts-ignore
disabled={true}
>
✨ {tag} (Enterprise only Feature)
</SelectItem>
);
})}
</MultiSelect>
</div>
)}
</Col>
</Grid>
<Grid numItems={2} className="gap-2 h-[75vh] w-full mb-4">
<Col numColSpan={2}>
<Card>
<Title>Spend Per Tag</Title>
<Text>
<Card className="mt-4">
<CardContent>
<div className="max-h-[70vh] min-h-[500px] overflow-y-auto">
<Table>
<TableHeader>
<TableRow>
<TableHead>Customer</TableHead>
<TableHead>Spend</TableHead>
<TableHead>Total Events</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{topUsers?.map((user: any, index: number) => (
<TableRow key={index}>
<TableCell>{user.end_user}</TableCell>
<TableCell>
<MoneyCell value={user.total_spend} decimals={2} />
</TableCell>
<TableCell>{user.total_count}</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
</CardContent>
</Card>
</TabsContent>
<TabsContent value="tag-based-usage">
<div className="grid grid-cols-2">
<div className="col-span-1">
<UsageDatePicker
className="mb-4"
value={dateValue}
onValueChange={(value) => {
setDateValue(value);
updateTagSpendData(value.from, value.to);
}}
/>
</div>
<div>
<Combobox
multiple
items={tagOptions}
value={tagOptions.filter((option) => selectedTags.includes(option.value))}
onValueChange={(options: TagOption[]) => setSelectedTags(options.map((option) => option.value))}
isItemEqualToValue={(a: TagOption, b: TagOption) => a.value === b.value}
itemToStringLabel={(option: TagOption) => option.label}
>
<ComboboxChips>
<ComboboxValue>
{(options: TagOption[]) =>
options.map((option) => (
<ComboboxChip key={option.value} aria-label={option.label}>
{option.label}
</ComboboxChip>
))
}
</ComboboxValue>
<ComboboxChipsInput placeholder="Select tags" className="border-0 bg-transparent" />
</ComboboxChips>
<ComboboxContent>
<ComboboxEmpty>No tags found</ComboboxEmpty>
<ComboboxList>
{(option: TagOption) => (
<ComboboxItem key={option.value} value={option} disabled={option.disabled}>
{option.label}
</ComboboxItem>
)}
</ComboboxList>
</ComboboxContent>
</Combobox>
</div>
</div>
<div className="mb-4 grid h-[75vh] w-full grid-cols-2 gap-2">
<div className="col-span-2">
<Card>
<CardHeader>
<CardTitle>Spend Per Tag</CardTitle>
</CardHeader>
<CardContent className="flex flex-col gap-2">
<p className="text-sm text-muted-foreground">
Get Started by Tracking cost per tag{" "}
<a
className="text-blue-500"
className="text-primary"
href="https://docs.litellm.ai/docs/proxy/cost_tracking"
target="_blank"
rel="noreferrer"
>
here
</a>
</Text>
<BarChart
className="h-72"
data={topTagsData}
index="name"
categories={["spend"]}
colors={["cyan"]}
></BarChart>
</Card>
</Col>
<Col numColSpan={2}></Col>
</Grid>
</TabPanel>
</TabPanels>
</TabGroup>
</p>
<BarChart className="h-72" data={topTagsData} index="name" categories={["spend"]} colors={["cyan"]} />
</CardContent>
</Card>
</div>
</div>
</TabsContent>
</Tabs>
</div>
);
};

View file

@ -1,7 +1,8 @@
import { render, screen } from "@testing-library/react";
import { render, screen, waitFor } from "@testing-library/react";
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { getPromptsList } from "@/components/networking";
import { deletePromptCall, getPromptsList } from "@/components/networking";
import PromptsPanel from "./index";
@ -12,20 +13,39 @@ vi.mock("@/components/networking", () => ({
vi.mock("./PromptTable", () => ({
__esModule: true,
default: ({ isLoading }: { isLoading: boolean }) => (
<div data-testid="prompt-table">{isLoading ? "table-loading" : "table-loaded"}</div>
default: ({
isLoading,
onDeleteClick,
}: {
isLoading: boolean;
onDeleteClick: (id: string, name: string) => void;
}) => (
<div data-testid="prompt-table">
{isLoading ? "table-loading" : "table-loaded"}
<button type="button" onClick={() => onDeleteClick("prompt-1", "my-prompt")}>
row-delete
</button>
</div>
),
}));
vi.mock("./prompt_info", () => ({ __esModule: true, default: () => null }));
vi.mock("./add_prompt_form", () => ({ __esModule: true, default: () => null }));
vi.mock("./prompt_editor_view", () => ({ __esModule: true, default: () => null }));
vi.mock("./prompt_info", () => ({ __esModule: true, default: () => <div>prompt-info-view</div> }));
vi.mock("./add_prompt_form", () => ({
__esModule: true,
default: ({ visible }: { visible: boolean }) => (visible ? <div>add-prompt-form</div> : null),
}));
vi.mock("./prompt_editor_view", () => ({ __esModule: true, default: () => <div>prompt-editor-view</div> }));
const mockGetPromptsList = vi.mocked(getPromptsList);
const mockDeletePromptCall = vi.mocked(deletePromptCall);
const renderPanel = (userRole?: string) =>
render(<PromptsPanel accessToken="sk-test" userRole={userRole ?? "Admin"} />);
describe("PromptsPanel loading state", () => {
beforeEach(() => {
vi.clearAllMocks();
mockGetPromptsList.mockResolvedValue({ prompts: [] } as never);
});
it("should resolve the loading state when accessToken is null instead of showing the skeleton forever", async () => {
@ -39,7 +59,7 @@ describe("PromptsPanel loading state", () => {
mockGetPromptsList.mockReturnValue(
new Promise((resolve) => {
resolveFetch = resolve;
}),
}) as never,
);
render(<PromptsPanel accessToken="sk-test" userRole="Admin" />);
expect(screen.getByText("table-loading")).toBeInTheDocument();
@ -49,3 +69,134 @@ describe("PromptsPanel loading state", () => {
expect(mockGetPromptsList).toHaveBeenCalledWith("sk-test", undefined);
});
});
describe("PromptsPanel toolbar", () => {
beforeEach(() => {
vi.clearAllMocks();
mockGetPromptsList.mockResolvedValue({ prompts: [] } as never);
});
it("should offer both create actions to a proxy admin", async () => {
renderPanel("Admin");
expect(await screen.findByRole("button", { name: /add new prompt/i })).toBeEnabled();
expect(screen.getByRole("button", { name: /upload \.prompt file/i })).toBeEnabled();
});
it("should hide both create actions from a read-only viewer", async () => {
renderPanel("Admin Viewer");
expect(await screen.findByText("table-loaded")).toBeInTheDocument();
expect(screen.queryByRole("button", { name: /add new prompt/i })).not.toBeInTheDocument();
expect(screen.queryByRole("button", { name: /upload \.prompt file/i })).not.toBeInTheDocument();
});
it("should open the editor view when the add action is used", async () => {
const user = userEvent.setup();
renderPanel("Admin");
await user.click(await screen.findByRole("button", { name: /add new prompt/i }));
expect(screen.getByText("prompt-editor-view")).toBeInTheDocument();
expect(screen.queryByTestId("prompt-table")).not.toBeInTheDocument();
});
it("should open the upload form when the upload action is used", async () => {
const user = userEvent.setup();
renderPanel("Admin");
expect(screen.queryByText("add-prompt-form")).not.toBeInTheDocument();
await user.click(await screen.findByRole("button", { name: /upload \.prompt file/i }));
expect(screen.getByText("add-prompt-form")).toBeInTheDocument();
});
it("should refetch scoped to the environment picked in the filter", async () => {
const user = userEvent.setup();
renderPanel("Admin");
await screen.findByText("table-loaded");
expect(screen.getByText("All Environments")).toBeInTheDocument();
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("Production"));
await waitFor(() => expect(mockGetPromptsList).toHaveBeenLastCalledWith("sk-test", "production"));
});
it("should show the picked environment by label and clear back to the unfiltered list", async () => {
// Base UI's exit animation never completes in jsdom, so the closing popup keeps
// pointer-events: none and blocks the second open. The clicks still dispatch.
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
renderPanel("Admin");
await screen.findByText("table-loaded");
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("Production"));
await waitFor(() => expect(screen.getByRole("combobox")).toHaveTextContent("Production"));
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("All Environments"));
await waitFor(() => expect(screen.getByRole("combobox")).toHaveTextContent("All Environments"));
await waitFor(() => expect(mockGetPromptsList).toHaveBeenLastCalledWith("sk-test", undefined));
});
});
describe("PromptsPanel delete confirmation", () => {
beforeEach(() => {
vi.clearAllMocks();
mockGetPromptsList.mockResolvedValue({ prompts: [] } as never);
mockDeletePromptCall.mockResolvedValue(undefined as never);
});
it("should not delete until the confirmation is accepted", async () => {
const user = userEvent.setup();
renderPanel("Admin");
await user.click(await screen.findByRole("button", { name: "row-delete" }));
expect(await screen.findByText(/delete prompt: my-prompt/i)).toBeInTheDocument();
expect(screen.getByText(/cannot be undone/i)).toBeInTheDocument();
expect(mockDeletePromptCall).not.toHaveBeenCalled();
await user.click(screen.getByRole("button", { name: /^delete$/i }));
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1"));
});
it("should abandon the delete when the confirmation is dismissed", async () => {
const user = userEvent.setup();
renderPanel("Admin");
await user.click(await screen.findByRole("button", { name: "row-delete" }));
await screen.findByText(/delete prompt: my-prompt/i);
await user.click(screen.getByRole("button", { name: /cancel/i }));
await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument());
expect(mockDeletePromptCall).not.toHaveBeenCalled();
});
it("should keep the confirmation up while the delete request is still in flight", async () => {
const user = userEvent.setup();
let finishDelete: () => void = () => {};
mockDeletePromptCall.mockReturnValue(
new Promise<void>((resolve) => {
finishDelete = () => resolve();
}) as never,
);
renderPanel("Admin");
await user.click(await screen.findByRole("button", { name: "row-delete" }));
await screen.findByText(/delete prompt: my-prompt/i);
await user.click(screen.getByRole("button", { name: /^delete$/i }));
await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1"));
await user.keyboard("{Escape}");
expect(screen.getByText(/delete prompt: my-prompt/i)).toBeInTheDocument();
finishDelete();
await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument());
});
});

View file

@ -1,7 +1,6 @@
import React, { useState, useEffect } from "react";
import { Button } from "@tremor/react";
import { Modal, Select } from "antd";
import { Plus, Upload } from "lucide-react";
import { getPromptsList, PromptSpec, ListPromptsResponse, deletePromptCall } from "@/components/networking";
import PromptTable from "./PromptTable";
import PromptInfoView from "./prompt_info";
@ -9,6 +8,28 @@ import AddPromptForm from "./add_prompt_form";
import PromptEditorView from "./prompt_editor_view";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { isAdminRole, isProxyAdminRole } from "@/utils/roles";
import { Button } from "@/components/ui/button";
import {
AlertDialog,
AlertDialogCancel,
AlertDialogContent,
AlertDialogDescription,
AlertDialogFooter,
AlertDialogHeader,
AlertDialogTitle,
} from "@/components/ui/alert-dialog";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
const ALL_ENVIRONMENTS_LABEL = "All Environments";
const ENVIRONMENT_OPTIONS = [
{ label: "Development", value: "development" },
{ label: "Staging", value: "staging" },
{ label: "Production", value: "production" },
];
// SelectValue falls back to the raw value unless the root can map it to a label.
const ENVIRONMENT_ITEMS = [{ label: ALL_ENVIRONMENTS_LABEL, value: null }, ...ENVIRONMENT_OPTIONS];
interface PromptsProps {
accessToken: string | null;
@ -141,26 +162,33 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
{canModify && (
<>
<Button onClick={handleAddPrompt} disabled={!accessToken}>
+ Add New Prompt
<Plus />
Add New Prompt
</Button>
<Button onClick={handleAddPromptFromFile} disabled={!accessToken} variant="secondary">
<Upload />
Upload .prompt File
</Button>
</>
)}
</div>
<Select
placeholder="All Environments"
allowClear
value={selectedEnvironment}
onChange={(value) => setSelectedEnvironment(value)}
style={{ width: 180 }}
options={[
{ label: "Development", value: "development" },
{ label: "Staging", value: "staging" },
{ label: "Production", value: "production" },
]}
/>
items={ENVIRONMENT_ITEMS}
value={selectedEnvironment ?? null}
onValueChange={(value) => setSelectedEnvironment((value as string | null) ?? undefined)}
>
<SelectTrigger className="w-[180px]">
<SelectValue placeholder={ALL_ENVIRONMENTS_LABEL} />
</SelectTrigger>
<SelectContent>
<SelectItem value={null}>{ALL_ENVIRONMENTS_LABEL}</SelectItem>
{ENVIRONMENT_OPTIONS.map((option) => (
<SelectItem key={option.value} value={option.value}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
</div>
<PromptTable
@ -182,18 +210,27 @@ const PromptsPanel: React.FC<PromptsProps> = ({ accessToken, userRole }) => {
/>
{promptToDelete && (
<Modal
title="Delete Prompt"
open={promptToDelete !== null}
onOk={handleDeleteConfirm}
onCancel={handleDeleteCancel}
confirmLoading={isDeleting}
okText="Delete"
okButtonProps={{ danger: true }}
<AlertDialog
open
onOpenChange={(open) => {
if (!open && !isDeleting) handleDeleteCancel();
}}
>
<p>Are you sure you want to delete prompt: {promptToDelete.name} ?</p>
<p>This action cannot be undone.</p>
</Modal>
<AlertDialogContent>
<AlertDialogHeader>
<AlertDialogTitle>Delete Prompt</AlertDialogTitle>
<AlertDialogDescription>
Are you sure you want to delete prompt: {promptToDelete.name} ? This action cannot be undone.
</AlertDialogDescription>
</AlertDialogHeader>
<AlertDialogFooter>
<AlertDialogCancel disabled={isDeleting}>Cancel</AlertDialogCancel>
<Button variant="destructive" onClick={handleDeleteConfirm} disabled={isDeleting}>
Delete
</Button>
</AlertDialogFooter>
</AlertDialogContent>
</AlertDialog>
)}
</div>
);

View file

@ -0,0 +1,160 @@
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
import TransformRequestPanel from "./TransformRequestPanel";
import { transformRequestCall } from "@/components/networking";
import NotificationsManager from "@/components/molecules/notifications_manager";
vi.mock("@/components/networking", () => ({
transformRequestCall: vi.fn(),
}));
vi.mock("@/components/molecules/notifications_manager", () => ({
default: {
success: vi.fn(),
info: vi.fn(),
fromBackend: vi.fn(),
},
}));
const transformRequestCallMock = vi.mocked(transformRequestCall);
const notify = vi.mocked(NotificationsManager);
const ACCESS_TOKEN = "sk-test-token";
const getRequestTextarea = () => screen.getByPlaceholderText(/press cmd\/ctrl \+ enter to transform/i);
const getTransformButton = () => screen.getByRole("button", { name: /transform/i });
const getCopyButton = () => screen.getByRole("button", { name: /copy to clipboard/i });
describe("TransformRequestPanel", () => {
beforeEach(() => {
vi.clearAllMocks();
});
afterEach(() => {
vi.restoreAllMocks();
});
it("renders both panels, the prefilled request and the placeholder curl", () => {
render(<TransformRequestPanel accessToken={ACCESS_TOKEN} />);
expect(screen.getByText("Original Request")).toBeInTheDocument();
expect(screen.getByText("Transformed Request")).toBeInTheDocument();
expect(screen.getByText(/sensitive headers are not shown/i)).toBeInTheDocument();
expect((getRequestTextarea() as HTMLTextAreaElement).value).toContain('"model": "openai/gpt-4o"');
expect(screen.getByText(/https:\/\/api\.openai\.com\/v1\/chat\/completions/)).toBeInTheDocument();
expect(screen.getByRole("link", { name: /here/i })).toHaveAttribute(
"href",
"https://github.com/BerriAI/litellm/issues",
);
});
it("sends the edited request body as a completion call and renders the returned curl", async () => {
const user = userEvent.setup();
transformRequestCallMock.mockResolvedValue({
raw_request_api_base: "https://api.anthropic.com/v1/messages",
raw_request_body: { model: "claude-opus-4-8", max_tokens: 42 },
raw_request_headers: { "x-api-key": "redacted" },
});
render(<TransformRequestPanel accessToken={ACCESS_TOKEN} />);
const textarea = getRequestTextarea();
await user.clear(textarea);
await user.type(textarea, '{{"model": "claude-opus-4-8"}');
await user.click(getTransformButton());
await waitFor(() => expect(transformRequestCallMock).toHaveBeenCalledTimes(1));
expect(transformRequestCallMock).toHaveBeenCalledWith(ACCESS_TOKEN, {
call_type: "completion",
request_body: { model: "claude-opus-4-8" },
});
const output = await screen.findByText(/api\.anthropic\.com\/v1\/messages/);
expect(output.textContent).toContain("curl -X POST");
expect(output.textContent).toContain("-H 'x-api-key: redacted'");
expect(output.textContent).toContain('"model": "claude-opus-4-8"');
expect(output.textContent).toContain('"max_tokens": 42');
expect(notify.success).toHaveBeenCalledWith("Request transformed successfully");
});
it("transforms on Cmd/Ctrl + Enter without clicking the button", async () => {
const user = userEvent.setup();
transformRequestCallMock.mockResolvedValue({
raw_request_api_base: "https://api.openai.com/v1/chat/completions",
raw_request_body: { model: "gpt-4o" },
raw_request_headers: {},
});
render(<TransformRequestPanel accessToken={ACCESS_TOKEN} />);
getRequestTextarea().focus();
await user.keyboard("{Meta>}{Enter}{/Meta}");
await waitFor(() => expect(transformRequestCallMock).toHaveBeenCalledTimes(1));
});
it("rejects invalid JSON without calling the backend", async () => {
const user = userEvent.setup();
render(<TransformRequestPanel accessToken={ACCESS_TOKEN} />);
const textarea = getRequestTextarea();
await user.clear(textarea);
await user.type(textarea, "not json");
await user.click(getTransformButton());
await waitFor(() => expect(notify.fromBackend).toHaveBeenCalledWith("Invalid JSON in request body"));
expect(transformRequestCallMock).not.toHaveBeenCalled();
});
it("does not call the backend when there is no access token", async () => {
const user = userEvent.setup();
render(<TransformRequestPanel accessToken={null} />);
await user.click(getTransformButton());
await waitFor(() => expect(notify.fromBackend).toHaveBeenCalledWith("No access token found"));
expect(transformRequestCallMock).not.toHaveBeenCalled();
});
it("reports a failed transform and leaves the placeholder curl in place", async () => {
const user = userEvent.setup();
vi.spyOn(console, "error").mockImplementation(() => {});
transformRequestCallMock.mockRejectedValue(new Error("boom"));
render(<TransformRequestPanel accessToken={ACCESS_TOKEN} />);
await user.click(getTransformButton());
await waitFor(() => expect(notify.fromBackend).toHaveBeenCalledWith("Failed to transform request"));
expect(screen.getByText(/https:\/\/api\.openai\.com\/v1\/chat\/completions/)).toBeInTheDocument();
});
it("copies the transformed request to the clipboard", async () => {
const user = userEvent.setup();
const writeText = vi.spyOn(navigator.clipboard, "writeText");
transformRequestCallMock.mockResolvedValue({
raw_request_api_base: "https://api.anthropic.com/v1/messages",
raw_request_body: { model: "claude-opus-4-8" },
raw_request_headers: {},
});
render(<TransformRequestPanel accessToken={ACCESS_TOKEN} />);
await user.click(getTransformButton());
await screen.findByText(/api\.anthropic\.com\/v1\/messages/);
await user.click(getCopyButton());
expect(writeText).toHaveBeenCalledTimes(1);
expect(writeText.mock.calls[0]?.[0]).toContain("https://api.anthropic.com/v1/messages");
expect(notify.success).toHaveBeenCalledWith("Copied to clipboard");
});
});

View file

@ -1,9 +1,12 @@
import React, { useState } from "react";
import { Button } from "antd";
import { CopyOutlined } from "@ant-design/icons";
import { Title } from "@tremor/react";
import { ArrowRight, Copy } from "lucide-react";
import { Button } from "@/components/ui/button";
import { Card, CardContent, CardDescription, CardFooter, CardHeader, CardTitle } from "@/components/ui/card";
import { Textarea } from "@/components/ui/textarea";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { transformRequestCall } from "@/components/networking";
import NotificationsManager from "@/components/molecules/notifications_manager";
interface TransformRequestPanelProps {
accessToken: string | null;
}
@ -128,130 +131,50 @@ ${formattedBody}
};
return (
<div className="w-full m-2" style={{ overflow: "hidden" }}>
<Title>Playground</Title>
<p className="text-sm text-gray-500">See how LiteLLM transforms your request for the specified provider.</p>
<div
style={{
display: "flex",
gap: "16px",
width: "100%",
minWidth: 0,
overflow: "hidden",
}}
className="mt-4"
>
<div className="p-2">
<h1 className="text-lg font-medium text-foreground">Playground</h1>
<p className="text-sm text-muted-foreground">
See how LiteLLM transforms your request for the specified provider.
</p>
<div className="mt-4 grid grid-cols-1 gap-4 lg:grid-cols-2">
{/* Original Request Panel */}
<div
style={{
flex: "1 1 50%",
display: "flex",
flexDirection: "column",
border: "1px solid #e8e8e8",
borderRadius: "8px",
padding: "24px",
overflow: "hidden",
maxHeight: "600px",
minWidth: 0,
}}
>
<div style={{ marginBottom: "24px" }}>
<h2 style={{ fontSize: "24px", fontWeight: "bold", margin: "0 0 4px 0" }}>Original Request</h2>
<p style={{ color: "#666", margin: 0 }}>
The request you would send to LiteLLM /chat/completions endpoint.
</p>
</div>
<Card>
<CardHeader>
<CardTitle className="text-2xl font-bold">Original Request</CardTitle>
<CardDescription>The request you would send to LiteLLM /chat/completions endpoint.</CardDescription>
</CardHeader>
<textarea
style={{
flex: "1 1 auto",
width: "100%",
minHeight: "240px",
padding: "16px",
border: "1px solid #e8e8e8",
borderRadius: "6px",
fontFamily: "monospace",
fontSize: "14px",
resize: "none",
marginBottom: "24px",
overflow: "auto",
}}
value={originalRequestJSON}
onChange={(e) => setOriginalRequestJSON(e.target.value)}
onKeyDown={handleKeyDown}
placeholder="Press Cmd/Ctrl + Enter to transform"
/>
<CardContent>
<Textarea
className="h-72 resize-none p-4 font-mono text-sm field-sizing-fixed"
value={originalRequestJSON}
onChange={(e) => setOriginalRequestJSON(e.target.value)}
onKeyDown={handleKeyDown}
placeholder="Press Cmd/Ctrl + Enter to transform"
/>
</CardContent>
<div
style={{
display: "flex",
justifyContent: "flex-end",
marginTop: "auto",
}}
>
<Button
type="primary"
style={{
backgroundColor: "#000",
display: "flex",
alignItems: "center",
gap: "8px",
}}
onClick={handleTransform}
loading={isLoading}
>
<CardFooter className="justify-end">
<Button onClick={handleTransform} disabled={isLoading}>
<span>Transform</span>
<span>→</span>
{isLoading ? <UiLoadingSpinner className="size-4" /> : <ArrowRight />}
</Button>
</div>
</div>
</CardFooter>
</Card>
{/* Transformed Request Panel */}
<div
style={{
flex: "1 1 50%",
display: "flex",
flexDirection: "column",
border: "1px solid #e8e8e8",
borderRadius: "8px",
padding: "24px",
overflow: "hidden",
maxHeight: "800px",
minWidth: 0,
}}
>
<div style={{ marginBottom: "24px" }}>
<h2 style={{ fontSize: "24px", fontWeight: "bold", margin: "0 0 4px 0" }}>Transformed Request</h2>
<p style={{ color: "#666", margin: 0 }}>How LiteLLM transforms your request for the specified provider.</p>
<br />
<p style={{ color: "#666", margin: 0 }} className="text-xs">
Note: Sensitive headers are not shown.
</p>
</div>
<Card>
<CardHeader>
<CardTitle className="text-2xl font-bold">Transformed Request</CardTitle>
<CardDescription>How LiteLLM transforms your request for the specified provider.</CardDescription>
<p className="mt-2 text-xs text-muted-foreground">Note: Sensitive headers are not shown.</p>
</CardHeader>
<div
style={{
position: "relative",
backgroundColor: "#f5f5f5",
borderRadius: "6px",
flex: "1 1 auto",
display: "flex",
flexDirection: "column",
overflow: "hidden",
}}
>
<pre
style={{
padding: "16px",
fontFamily: "monospace",
fontSize: "14px",
margin: 0,
overflow: "auto",
flex: "1 1 auto",
}}
>
{transformedResponse ||
`curl -X POST \\
<CardContent>
<div className="relative rounded-md bg-muted">
<pre className="h-72 overflow-auto p-4 font-mono text-sm">
{transformedResponse ||
`curl -X POST \\
https://api.openai.com/v1/chat/completions \\
-H 'Authorization: Bearer sk-xxx' \\
-H 'Content-Type: application/json' \\
@ -265,29 +188,33 @@ ${formattedBody}
],
"temperature": 0.7
}'`}
</pre>
</pre>
<Button
type="text"
icon={<CopyOutlined />}
style={{
position: "absolute",
right: "8px",
top: "8px",
}}
size="small"
onClick={() => {
navigator.clipboard.writeText(transformedResponse || "");
NotificationsManager.success("Copied to clipboard");
}}
/>
</div>
</div>
<Button
variant="ghost"
size="icon-sm"
aria-label="Copy to clipboard"
className="absolute top-2 right-2"
onClick={() => {
navigator.clipboard.writeText(transformedResponse || "");
NotificationsManager.success("Copied to clipboard");
}}
>
<Copy />
</Button>
</div>
</CardContent>
</Card>
</div>
<div className="mt-4 text-right w-full">
<p className="text-sm text-gray-500">
<div className="mt-4 text-right">
<p className="text-sm text-muted-foreground">
Found an error? File an issue{" "}
<a href="https://github.com/BerriAI/litellm/issues" target="_blank" rel="noopener noreferrer">
<a
className="underline underline-offset-4"
href="https://github.com/BerriAI/litellm/issues"
target="_blank"
rel="noopener noreferrer"
>
here
</a>
.

View file

@ -0,0 +1,87 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import EditFallbacks, { Fallbacks } from "./EditFallbacks";
import * as fetchModelsModule from "@/components/llm_calls/fetch_models";
vi.mock("@/components/llm_calls/fetch_models", () => ({
fetchAvailableModels: vi.fn(),
}));
const renderWithQueryClient = (ui: React.ReactElement) => {
const queryClient = new QueryClient({
defaultOptions: { queries: { retry: false } },
});
return render(<QueryClientProvider client={queryClient}>{ui}</QueryClientProvider>);
};
describe("EditFallbacks", () => {
const accessToken = "test-token";
const fallbackEntry = { "gpt-4": ["gpt-3.5-turbo", "claude-3-opus"] };
const value: Fallbacks = [{ "gpt-4": ["gpt-3.5-turbo", "claude-3-opus"] }, { "claude-3-opus": ["gpt-4"] }];
const setup = (overrides: Partial<React.ComponentProps<typeof EditFallbacks>> = {}) => {
const onChange = overrides.onChange ?? vi.fn().mockResolvedValue(undefined);
const onClose = overrides.onClose ?? vi.fn();
renderWithQueryClient(
<EditFallbacks
accessToken={accessToken}
fallbackEntry={fallbackEntry}
value={value}
onChange={onChange}
onClose={onClose}
{...overrides}
/>,
);
return { onChange, onClose };
};
beforeEach(() => {
vi.clearAllMocks();
vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValue([
{ model_group: "gpt-4", mode: "chat" },
{ model_group: "gpt-3.5-turbo", mode: "chat" },
{ model_group: "claude-3-opus", mode: "chat" },
{ model_group: "gemini-pro", mode: "chat" },
]);
});
it("prefills the existing fallback chain for the primary model", async () => {
setup();
await waitFor(() => {
expect(screen.getByText("gpt-3.5-turbo")).toBeInTheDocument();
expect(screen.getByText("claude-3-opus")).toBeInTheDocument();
});
});
it("removes a fallback model and saves only the edited entry", async () => {
const user = userEvent.setup();
const onChange = vi.fn().mockResolvedValue(undefined);
const onClose = vi.fn();
setup({ onChange, onClose });
await screen.findByText("gpt-3.5-turbo");
await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo"));
await user.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onChange).toHaveBeenCalledWith([{ "gpt-4": ["claude-3-opus"] }, { "claude-3-opus": ["gpt-4"] }]);
});
await waitFor(() => expect(onClose).toHaveBeenCalled());
});
it("blocks saving with an empty fallback chain", async () => {
const user = userEvent.setup();
const onChange = vi.fn().mockResolvedValue(undefined);
setup({ fallbackEntry: { "gpt-4": ["gpt-3.5-turbo"] }, onChange });
await screen.findByText("gpt-3.5-turbo");
await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo"));
const saveButton = screen.getByRole("button", { name: /save changes/i });
expect(saveButton).toBeDisabled();
expect(onChange).not.toHaveBeenCalled();
});
});

View file

@ -0,0 +1,106 @@
/**
* Modal for editing an existing fallback entry
* Lets the user add/remove models from a primary model's fallback chain
* Reuses FallbackGroupConfig with the primary model locked
*/
import { Button } from "antd";
import { useQuery } from "@tanstack/react-query";
import { Pencil } from "lucide-react";
import React, { useMemo, useState } from "react";
import { fetchAvailableModels } from "@/components/llm_calls/fetch_models";
import NotificationManager from "../../../molecules/notifications_manager";
import { AddFallbacksModal } from "./AddFallbacksModal";
import { FallbackGroup, FallbackGroupConfig } from "./FallbackGroupConfig";
export type FallbackEntry = { [modelName: string]: string[] };
export type Fallbacks = FallbackEntry[];
interface EditFallbacksProps {
accessToken: string;
fallbackEntry: FallbackEntry;
value: Fallbacks;
onChange: (fallbacks: Fallbacks) => Promise<void>;
onClose: () => void;
maxFallbacks?: number;
}
const toGroup = (entry: FallbackEntry): FallbackGroup => {
const primaryModel = Object.keys(entry)[0] ?? null;
return {
id: "edit",
primaryModel,
fallbackModels: primaryModel ? [...(entry[primaryModel] ?? [])] : [],
};
};
export default function EditFallbacks({
accessToken,
fallbackEntry,
value,
onChange,
onClose,
maxFallbacks = 10,
}: EditFallbacksProps) {
const [group, setGroup] = useState<FallbackGroup>(() => toGroup(fallbackEntry));
const [isSaving, setIsSaving] = useState(false);
const { data: modelGroups = [] } = useQuery({
queryKey: ["availableModels", "fallbacks"],
queryFn: () => fetchAvailableModels(accessToken),
enabled: Boolean(accessToken),
});
const availableModels = useMemo(
() => Array.from(new Set(modelGroups.map((option) => option.model_group))).sort(),
[modelGroups],
);
const handleSave = async () => {
const primaryModel = group.primaryModel;
if (!primaryModel) {
return;
}
const updatedFallbacks = (value || []).map((entry) =>
primaryModel in entry ? { ...entry, [primaryModel]: group.fallbackModels } : entry,
);
setIsSaving(true);
try {
await onChange(updatedFallbacks);
NotificationManager.success(`Fallbacks for ${primaryModel} updated successfully!`);
onClose();
} catch (error) {
console.error("Error updating fallbacks:", error);
} finally {
setIsSaving(false);
}
};
return (
<AddFallbacksModal open onCancel={onClose}>
<FallbackGroupConfig
group={group}
onChange={setGroup}
availableModels={availableModels}
maxFallbacks={maxFallbacks}
disablePrimaryModel
/>
<div className="flex items-center justify-end space-x-3 pt-6 mt-6 border-t border-gray-100">
<Button type="default" onClick={onClose} disabled={isSaving}>
Cancel
</Button>
<Button
type="primary"
icon={<Pencil className="w-4 h-4" />}
onClick={handleSave}
disabled={isSaving || group.fallbackModels.length === 0}
loading={isSaving}
>
{isSaving ? "Saving Changes..." : "Save Changes"}
</Button>
</div>
</AddFallbacksModal>
);
}

View file

@ -18,9 +18,16 @@ interface FallbackGroupConfigProps {
onChange: (updatedGroup: FallbackGroup) => void;
availableModels: string[];
maxFallbacks: number;
disablePrimaryModel?: boolean;
}
export function FallbackGroupConfig({ group, onChange, availableModels, maxFallbacks }: FallbackGroupConfigProps) {
export function FallbackGroupConfig({
group,
onChange,
availableModels,
maxFallbacks,
disablePrimaryModel = false,
}: FallbackGroupConfigProps) {
// Filter available options for fallbacks (exclude primary only, allow already selected to be shown for deselection)
const availableFallbackOptions = availableModels.filter((m) => m !== group.primaryModel);
@ -70,12 +77,13 @@ export function FallbackGroupConfig({ group, onChange, availableModels, maxFallb
placeholder="Select primary model"
value={group.primaryModel}
onChange={handlePrimaryChange}
disabled={disablePrimaryModel}
showSearch
getPopupContainer={(trigger) => trigger.parentElement || document.body}
filterOption={(input, option) => (option?.label ?? "").toLowerCase().includes(input.toLowerCase())}
options={availableModels.map((m) => ({ label: m, value: m }))}
/>
{!group.primaryModel && (
{!disablePrimaryModel && !group.primaryModel && (
<div className="mt-2 flex items-center gap-2 text-amber-600 text-xs bg-amber-50 p-2 rounded-sm">
<AlertCircle className="w-4 h-4" />
<span>Select a model to begin configuring fallbacks</span>
@ -176,6 +184,7 @@ export function FallbackGroupConfig({ group, onChange, availableModels, maxFallb
<button
type="button"
data-testid={`remove-fallback-${modelValue}`}
onClick={() => removeFallback(index)}
className="opacity-0 group-hover:opacity-100 transition-opacity text-gray-400 hover:text-red-500 p-1"
>

View file

@ -1,3 +1,4 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
@ -94,6 +95,13 @@ describe("Fallbacks", () => {
return deleteButtons.length > 0 ? deleteButtons[0] : null;
};
const renderWithQueryClient = (ui: React.ReactElement) => {
const queryClient = new QueryClient({
defaultOptions: { queries: { retry: false } },
});
return render(<QueryClientProvider client={queryClient}>{ui}</QueryClientProvider>);
};
beforeEach(() => {
vi.clearAllMocks();
vi.mocked(networkingModule.getCallbacksCall).mockResolvedValue({
@ -108,7 +116,7 @@ describe("Fallbacks", () => {
});
it("should render the component", async () => {
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument();
@ -116,12 +124,12 @@ describe("Fallbacks", () => {
});
it("should not render when accessToken is null", () => {
const { container } = render(<Fallbacks {...defaultProps} accessToken={null} />);
const { container } = renderWithQueryClient(<Fallbacks {...defaultProps} accessToken={null} />);
expect(container.firstChild).toBeNull();
});
it("should fetch router settings on mount", async () => {
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(networkingModule.getCallbacksCall).toHaveBeenCalledWith(mockAccessToken, mockUserID, mockUserRole);
@ -129,7 +137,7 @@ describe("Fallbacks", () => {
});
it("should display fallback entries in table", async () => {
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
@ -139,7 +147,7 @@ describe("Fallbacks", () => {
});
it("should show delete button for each fallback row when fallbacks exist", async () => {
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
@ -149,9 +157,27 @@ describe("Fallbacks", () => {
expect(deleteButtons.length).toBe(2);
});
it("should show an edit button for each fallback row and open the edit modal", async () => {
const user = userEvent.setup();
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
});
const editButtons = screen.getAllByTestId("edit-fallback-button");
expect(editButtons.length).toBe(2);
await user.click(editButtons[0]);
await waitFor(() => {
expect(screen.getByText("Configure Model Fallbacks")).toBeInTheDocument();
});
});
it("should open delete modal when delete icon is clicked", async () => {
const user = userEvent.setup();
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
@ -170,7 +196,7 @@ describe("Fallbacks", () => {
it("should delete fallback when confirmed", async () => {
const user = userEvent.setup();
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
@ -198,7 +224,7 @@ describe("Fallbacks", () => {
it("should close delete modal when cancel is clicked", async () => {
const user = userEvent.setup();
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
@ -225,7 +251,7 @@ describe("Fallbacks", () => {
const user = userEvent.setup();
const error = new Error("Delete failed");
vi.mocked(networkingModule.setCallbacksCall).mockRejectedValueOnce(error);
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
@ -252,7 +278,7 @@ describe("Fallbacks", () => {
const user = userEvent.setup();
const error = new Error("Delete failed");
vi.mocked(networkingModule.setCallbacksCall).mockRejectedValueOnce(error);
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0);
@ -280,7 +306,7 @@ describe("Fallbacks", () => {
vi.mocked(networkingModule.getCallbacksCall).mockResolvedValueOnce({
router_settings: { fallbacks: [] },
});
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument();
@ -296,7 +322,7 @@ describe("Fallbacks", () => {
vi.mocked(networkingModule.getCallbacksCall).mockResolvedValueOnce({
router_settings: {},
});
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument();
@ -313,7 +339,7 @@ describe("Fallbacks", () => {
model_group_retry_policy: { some: "policy" },
},
});
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(networkingModule.getCallbacksCall).toHaveBeenCalled();
@ -322,7 +348,7 @@ describe("Fallbacks", () => {
it("should update fallbacks when AddFallbacks onChange is called", async () => {
const user = userEvent.setup();
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument();
@ -343,7 +369,7 @@ describe("Fallbacks", () => {
vi.mocked(networkingModule.getCallbacksCall).mockResolvedValue({
router_settings: mockRouterSettings,
});
render(<Fallbacks {...defaultProps} />);
renderWithQueryClient(<Fallbacks {...defaultProps} />);
await waitFor(() => {
expect(screen.getByTestId("add-fallbacks-button")).toBeInTheDocument();

View file

@ -1,5 +1,5 @@
import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap";
import { ArrowRightIcon, PlayIcon, TrashIcon } from "@heroicons/react/outline";
import { ArrowRightIcon, PencilAltIcon, PlayIcon, TrashIcon } from "@heroicons/react/outline";
import { Icon, Table, TableBody, TableCell, TableHead, TableHeaderCell, TableRow } from "@tremor/react";
import { Tooltip, Typography } from "antd";
import openai from "openai";
@ -10,6 +10,7 @@ import NotificationsManager from "../../../molecules/notifications_manager";
import { getCallbacksCall, setCallbacksCall } from "../../../networking";
import { isProxyAdminRole } from "@/utils/roles";
import AddFallbacks from "./AddFallbacks";
import EditFallbacks from "./EditFallbacks";
type FallbackEntry = { [modelName: string]: string[] };
type Fallbacks = FallbackEntry[];
@ -119,6 +120,7 @@ const Fallbacks: React.FC<FallbacksProps> = ({ accessToken, userRole, userID })
const [isDeleting, setIsDeleting] = useState(false);
const [fallbackToDelete, setFallbackToDelete] = useState<FallbackEntry | null>(null);
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false);
const [fallbackToEdit, setFallbackToEdit] = useState<FallbackEntry | null>(null);
const { data: modelCostMapData } = useModelCostMap();
const getProviderFromModel = (model: string): string => {
@ -146,6 +148,14 @@ const Fallbacks: React.FC<FallbacksProps> = ({ accessToken, userRole, userID })
setIsDeleteModalOpen(true);
};
const handleEditClick = (fallbackEntry: FallbackEntry) => {
setFallbackToEdit(fallbackEntry);
};
const handleEditClose = () => {
setFallbackToEdit(null);
};
const handleDeleteConfirm = async () => {
if (!fallbackToDelete || !accessToken) {
return;
@ -281,6 +291,18 @@ const Fallbacks: React.FC<FallbacksProps> = ({ accessToken, userRole, userID })
className="cursor-pointer hover:text-blue-600"
/>
</Tooltip>
<Tooltip title="Edit fallback">
<span
data-testid="edit-fallback-button"
role="button"
tabIndex={0}
onClick={() => handleEditClick(item)}
onKeyDown={(e) => e.key === "Enter" && handleEditClick(item)}
className="cursor-pointer inline-flex"
>
<Icon icon={PencilAltIcon} size="sm" className="hover:text-blue-600" />
</span>
</Tooltip>
<Tooltip title="Delete fallback">
<span
data-testid="delete-fallback-button"
@ -302,6 +324,16 @@ const Fallbacks: React.FC<FallbacksProps> = ({ accessToken, userRole, userID })
</TableBody>
</Table>
)}
{canModify && fallbackToEdit && (
<EditFallbacks
key={Object.keys(fallbackToEdit)[0]}
accessToken={accessToken || ""}
fallbackEntry={fallbackToEdit}
value={routerSettings.fallbacks || []}
onChange={handleFallbacksChange}
onClose={handleEditClose}
/>
)}
<DeleteResourceModal
isOpen={isDeleteModalOpen}
title="Delete Fallback?"

View file

@ -7,6 +7,7 @@ import {
handleTransport,
handleAuth,
getMcpOAuthMode,
gatewayMintsClientFor,
getOAuthAuthorizationIdentity,
isHeldOAuthTokenStale,
oauth2FlowToFormValue,
@ -101,6 +102,41 @@ describe("constants", () => {
});
});
describe("gatewayMintsClientFor", () => {
// The authoritative client-acquisition matrix: for each (auth_type, dcr_bridge) cell, does the
// gateway mint the OAuth client at /authorize (browser skips its own register) or not (browser
// registers)? This MUST equal the backend resolve_ephemeral_dcr_client mint set exactly, which
// test_discoverable_endpoints.py::test_resolve_ephemeral_dcr_client_mint_set_is_exact pins against
// the same predicate. A divergence in either direction dead-ends a mode (skip a register the
// gateway never performs) or double-registers, so both sides are enumerated against this table.
const MATRIX: Array<{ auth_type: string; dcr_bridge: boolean | null | undefined; mints: boolean }> = [
{ auth_type: AUTH_TYPE.TRUE_PASSTHROUGH, dcr_bridge: false, mints: true },
{ auth_type: AUTH_TYPE.TRUE_PASSTHROUGH, dcr_bridge: true, mints: true },
{ auth_type: AUTH_TYPE.TRUE_PASSTHROUGH, dcr_bridge: null, mints: true },
{ auth_type: AUTH_TYPE.OAUTH_DELEGATE, dcr_bridge: false, mints: true },
{ auth_type: AUTH_TYPE.OAUTH_DELEGATE, dcr_bridge: null, mints: true },
{ auth_type: AUTH_TYPE.OAUTH_DELEGATE, dcr_bridge: undefined, mints: true },
// The one interactive-sign-in cell the gateway must NOT mint (browser front-door register).
{ auth_type: AUTH_TYPE.OAUTH_DELEGATE, dcr_bridge: true, mints: false },
{ auth_type: AUTH_TYPE.OAUTH2, dcr_bridge: false, mints: false },
{ auth_type: AUTH_TYPE.OAUTH2, dcr_bridge: true, mints: false },
{ auth_type: AUTH_TYPE.OAUTH2_TOKEN_EXCHANGE, dcr_bridge: false, mints: false },
{ auth_type: AUTH_TYPE.API_KEY, dcr_bridge: false, mints: false },
{ auth_type: AUTH_TYPE.BEARER_TOKEN, dcr_bridge: false, mints: false },
{ auth_type: AUTH_TYPE.BASIC, dcr_bridge: false, mints: false },
{ auth_type: AUTH_TYPE.NONE, dcr_bridge: false, mints: false },
{ auth_type: AUTH_TYPE.TOKEN, dcr_bridge: false, mints: false },
{ auth_type: AUTH_TYPE.AWS_SIGV4, dcr_bridge: false, mints: false },
];
it.each(MATRIX)(
"mints=$mints for auth_type=$auth_type dcr_bridge=$dcr_bridge",
({ auth_type, dcr_bridge, mints }) => {
expect(gatewayMintsClientFor({ auth_type, dcr_bridge })).toBe(mints);
},
);
});
describe("getMcpOAuthMode", () => {
it("returns null for non-OAuth2 servers", () => {
expect(getMcpOAuthMode({ auth_type: AUTH_TYPE.API_KEY })).toBeNull();

View file

@ -52,6 +52,19 @@ export const AUTH_TYPE = {
export const isClientForwardedTokenMode = (authType?: string | null): boolean =>
authType === AUTH_TYPE.TRUE_PASSTHROUGH || authType === AUTH_TYPE.OAUTH_DELEGATE;
/**
* Whether the gateway acquires the OAuth client itself during /authorize (so the browser must NOT
* pre-register one). This MUST mirror the backend `resolve_ephemeral_dcr_client` mint condition
* exactly, cell for cell, or a mode where the two disagree either dead-ends (browser skips a
* register the gateway never performs) or double-registers. The gateway mints for every
* `true_passthrough` server, and for `oauth_delegate` only when it is NOT a dcr_bridge server: the
* interactive oauth_delegate dcr_bridge sign-in captures the SSO user through the browser's own
* front-door registration, which the mint must not preempt.
*/
export const gatewayMintsClientFor = (server: { auth_type?: string | null; dcr_bridge?: boolean | null }): boolean =>
server.auth_type === AUTH_TYPE.TRUE_PASSTHROUGH ||
(server.auth_type === AUTH_TYPE.OAUTH_DELEGATE && !server.dcr_bridge);
export const OAUTH_FLOW = {
INTERACTIVE: "interactive",
M2M: "m2m",
@ -316,6 +329,12 @@ export interface MCPToolsViewerProps {
oauth2_flow?: string | null;
/** When true (interactive OAuth2), the server uses PKCE passthrough. */
delegate_auth_to_upstream?: boolean | null;
/**
* Read together with auth_type by gatewayMintsClientFor: an oauth_delegate dcr_bridge server is
* NOT gateway-minted (its interactive sign-in registers through the browser front door), so the
* tools-page Authorize must know the flag to decide whether to pre-register a client.
*/
dcr_bridge?: boolean | null;
/**
* Connection field present on every OAuth2 flow (interactive and M2M alike),
* so it does not indicate the mode. Retained for callers/other uses; not read

View file

@ -6737,6 +6737,7 @@ interface RegisterMcpOAuthClientPayload {
grant_types?: string[];
response_types?: string[];
token_endpoint_auth_method?: string;
redirect_uris?: string[];
}
export const registerMcpOAuthClient = async (

View file

@ -453,6 +453,44 @@ describe("KeyInfoView handleKeyUpdate guardrails guard", () => {
});
});
describe("KeyInfoView handleKeyUpdate budget_duration", () => {
it("should send a canonical budget_duration through unchanged", async () => {
renderView(true);
fireEvent.click(screen.getByText("Settings"));
fireEvent.click(screen.getByText("Edit Settings"));
(globalThis as any).__TEST_FORM_VALUES = {
token: "tok_123",
budget_duration: "30d",
};
fireEvent.click(screen.getByText("Mock Submit"));
await waitFor(() => expect(keyUpdateCallMock).toHaveBeenCalled());
const [, sentPayload] = keyUpdateCallMock.mock.calls[0];
expect(sentPayload.budget_duration).toBe("30d");
});
it("should heal a legacy word-form budget_duration to canonical", async () => {
renderView(true);
fireEvent.click(screen.getByText("Settings"));
fireEvent.click(screen.getByText("Edit Settings"));
(globalThis as any).__TEST_FORM_VALUES = {
token: "tok_123",
budget_duration: "monthly",
};
fireEvent.click(screen.getByText("Mock Submit"));
await waitFor(() => expect(keyUpdateCallMock).toHaveBeenCalled());
const [, sentPayload] = keyUpdateCallMock.mock.calls[0];
expect(sentPayload.budget_duration).toBe("30d");
});
});
describe("KeyInfoView handleKeyUpdate empty strings", () => {
["tpm_limit", "rpm_limit", "max_parallel_requests", "max_budget"].forEach((limit) => {
it(`maps empty strings to null for ${limit}`, async () => {

View file

@ -1,4 +1,4 @@
import { fireEvent, screen, waitFor } from "@testing-library/react";
import { fireEvent, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { renderWithProviders } from "../../../tests/test-utils";
@ -662,6 +662,87 @@ describe("KeyEditView", () => {
});
});
it("should persist a canonical budget_duration value, not a word-form the backend cannot parse", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
userRole={"admin"}
premiumUser={false}
/>,
);
const resetBudgetItem = (await screen.findByText("Reset Budget")).closest(".ant-form-item");
expect(resetBudgetItem).not.toBeNull();
const combobox = within(resetBudgetItem as HTMLElement).getByRole("combobox");
await userEvent.click(combobox);
const weeklyOption = await screen.findByText("weekly");
await userEvent.click(weeklyOption);
const submitButton = screen.getByRole("button", { name: /save changes/i });
await userEvent.click(submitButton);
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalled();
const callArgs = onSubmitMock.mock.calls[0][0];
expect(callArgs.budget_duration).toBe("7d");
});
});
it("should keep an existing canonical budget_duration canonical when saved untouched", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
userRole={"admin"}
premiumUser={false}
/>,
);
const submitButton = await screen.findByRole("button", { name: /save changes/i });
await userEvent.click(submitButton);
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalled();
const callArgs = onSubmitMock.mock.calls[0][0];
expect(callArgs.budget_duration).toBe("30d");
});
});
it("should heal a legacy word-form budget_duration to canonical when saved untouched", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
const legacyKeyData = { ...MOCK_KEY_DATA, budget_duration: "monthly" };
renderWithProviders(
<KeyEditView
keyData={legacyKeyData}
onCancel={() => {}}
onSubmit={onSubmitMock}
accessToken={"test-token"}
userID={"test-user"}
userRole={"admin"}
premiumUser={false}
/>,
);
const submitButton = await screen.findByRole("button", { name: /save changes/i });
await userEvent.click(submitButton);
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalled();
const callArgs = onSubmitMock.mock.calls[0][0];
expect(callArgs.budget_duration).toBe("30d");
});
});
it("should omit budget_limits when existing windows are left untouched (issue #33246)", async () => {
// The backend treats any budget_limits in the payload as an admin-only
// budget change, so re-sending untouched windows 403s a non-admin owner.

View file

@ -10,6 +10,7 @@ import { useEffect, useState } from "react";
import { rolesWithWriteAccess } from "../../utils/roles";
import AgentSelector from "../agent_management/AgentSelector";
import AccessGroupSelector from "../common_components/AccessGroupSelector";
import BudgetDurationDropdown from "../common_components/budget_duration_dropdown";
import { mapInternalToDisplayNames } from "../callback_info_helpers";
import KeyLifecycleSettings from "../common_components/KeyLifecycleSettings";
import PassThroughRoutesSelector from "../common_components/PassThroughRoutesSelector";
@ -172,15 +173,16 @@ export function KeyEditView({
form.setFieldValue("disabled_callbacks", disabledCallbacks);
}, [form, disabledCallbacks]);
// Convert API budget duration to form format
// Normalize any legacy word-form budget duration to the canonical value the dropdown uses
const getBudgetDuration = (duration: string | null) => {
if (!duration) return null;
const durationMap: Record<string, string> = {
"24h": "daily",
"7d": "weekly",
"30d": "monthly",
const wordToCanonical: Record<string, string> = {
hourly: "1h",
daily: "24h",
weekly: "7d",
monthly: "30d",
};
return durationMap[duration] || null;
return wordToCanonical[duration] ?? duration;
};
// Set initial form values
@ -506,11 +508,7 @@ export function KeyEditView({
</Form.Item>
<Form.Item label="Reset Budget" name="budget_duration">
<Select placeholder="n/a">
<Select.Option value="daily">Daily</Select.Option>
<Select.Option value="weekly">Weekly</Select.Option>
<Select.Option value="monthly">Monthly</Select.Option>
</Select>
<BudgetDurationDropdown />
</Form.Item>
<Form.Item

View file

@ -295,14 +295,15 @@ export default function KeyInfoView({
}
delete formValues.logging_settings;
// Convert budget_duration to API format
// Normalize any legacy word-form budget_duration to the canonical API format
if (formValues.budget_duration) {
const durationMap: Record<string, string> = {
const wordToCanonical: Record<string, string> = {
hourly: "1h",
daily: "24h",
weekly: "7d",
monthly: "30d",
};
formValues.budget_duration = durationMap[formValues.budget_duration];
formValues.budget_duration = wordToCanonical[formValues.budget_duration] ?? formValues.budget_duration;
}
const newKeyValues = await keyUpdateCall(accessToken, formValues);

View file

@ -157,6 +157,10 @@ export const useMcpOAuthFlow = ({
response_types: ["code"],
token_endpoint_auth_method:
temporaryPayload.credentials && temporaryPayload.credentials.client_secret ? "client_secret_post" : "none",
// dcr_bridge servers relay this registration upstream and bind the
// minted client to the browser's own callback; without it the relay
// rejects the registration and the admin authorize dead-ends.
redirect_uris: [callbackUrl()],
});
registeredClient = {
clientId: registration?.client_id,

View file

@ -0,0 +1,89 @@
import { act, renderHook } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import * as networking from "@/components/networking";
import { useToolsOAuthFlow } from "./useToolsOAuthFlow";
vi.mock("@/components/networking", () => ({
exchangeMcpOAuthToken: vi.fn(),
registerMcpOAuthClient: vi.fn(),
buildMcpOAuthAuthorizeUrl: vi.fn(() => "https://gw.example.com/v1/mcp/server/oauth/server-1/authorize"),
getProxyBaseUrl: vi.fn(() => ""),
serverRootPath: "",
}));
vi.mock("@/components/molecules/notifications_manager", () => ({
default: { success: vi.fn(), error: vi.fn() },
}));
vi.mock("@/utils/pkce", () => ({
generateCodeVerifier: () => "verifier-1",
generateCodeChallenge: async () => "challenge-1",
}));
function renderFlow(options: { gatewayMintsClient?: boolean }) {
return renderHook(() =>
useToolsOAuthFlow({
accessToken: "user-token",
serverId: "server-1",
serverAlias: "server one",
userId: "user-1",
gatewayMintsClient: options.gatewayMintsClient,
onSuccess: vi.fn(),
}),
);
}
describe("useToolsOAuthFlow client registration", () => {
beforeEach(() => {
vi.clearAllMocks();
window.sessionStorage.clear();
window.localStorage.clear();
Object.defineProperty(window, "location", {
value: { href: "https://app.example.com/ui?page=mcp-servers" },
writable: true,
configurable: true,
});
});
it("skips browser-side client registration when the gateway mints the client", async () => {
// Client-forwarded token modes: the gateway's /authorize mints and carries the
// OAuth client itself, so a browser-side registration would only create an
// extra orphan client at the IdP. The authorize URL must go out clientless.
const { result } = renderFlow({ gatewayMintsClient: true });
await act(async () => {
await result.current.startOAuthFlow();
});
expect(networking.registerMcpOAuthClient).not.toHaveBeenCalled();
expect(networking.buildMcpOAuthAuthorizeUrl).toHaveBeenCalledWith(
expect.objectContaining({ serverId: "server-1", clientId: undefined }),
);
expect(window.location.href).toBe("https://gw.example.com/v1/mcp/server/oauth/server-1/authorize");
});
it("still registers browser-side for servers the gateway does not mint for", async () => {
// The cells where gatewayMintsClientFor is false keep the browser-held client: an
// oauth_delegate dcr_bridge server (interactive sign-in) and the legacy oauth2 passthrough both
// register first, so the minted client_id rides the authorize URL through the front-door relay.
vi.mocked(networking.registerMcpOAuthClient).mockResolvedValue({ client_id: "reg-client-1" });
const { result } = renderFlow({});
await act(async () => {
await result.current.startOAuthFlow();
});
expect(networking.registerMcpOAuthClient).toHaveBeenCalledTimes(1);
// The dcr_bridge relay rejects registrations without the client's own
// callback, so the browser must bind its minted client to it.
expect(networking.registerMcpOAuthClient).toHaveBeenCalledWith(
"user-token",
"server-1",
expect.objectContaining({ redirect_uris: [expect.stringContaining("callback")] }),
);
expect(networking.buildMcpOAuthAuthorizeUrl).toHaveBeenCalledWith(
expect.objectContaining({ serverId: "server-1", clientId: "reg-client-1" }),
);
});
});

View file

@ -31,6 +31,13 @@ interface UseToolsOAuthFlowOptions {
userId?: string | null;
scopes?: string[];
clientId?: string | null;
/**
* True for the client-forwarded token modes (true_passthrough / oauth_delegate): the gateway
* mints the OAuth client itself during /authorize and carries it in the sealed state/code, so
* the browser must not register a client of its own (each browser-side registration creates an
* extra client at the IdP that the gateway's minted one then supersedes).
*/
gatewayMintsClient?: boolean;
onSuccess: (accessToken: string) => void;
}
@ -61,6 +68,7 @@ export const useToolsOAuthFlow = ({
userId,
scopes,
clientId: preClientId,
gatewayMintsClient,
onSuccess,
}: UseToolsOAuthFlowOptions): UseToolsOAuthFlowResult => {
const [status, setStatus] = useState<ToolsOAuthStatus>("idle");
@ -77,14 +85,19 @@ export const useToolsOAuthFlow = ({
let clientId: string | undefined = preClientId ?? undefined;
let clientSecret: string | undefined;
const redirectUri = buildCallbackUrl();
if (!clientId) {
if (!clientId && !gatewayMintsClient) {
try {
const reg = await registerMcpOAuthClient(accessToken, serverId, {
client_name: serverAlias || serverId,
grant_types: ["authorization_code", "refresh_token"],
response_types: ["code"],
token_endpoint_auth_method: "none",
// dcr_bridge servers relay this registration upstream and bind the
// minted client to the browser's own callback; without it the relay
// rejects the registration and the flow dead-ends clientless.
redirect_uris: [redirectUri],
});
clientId = reg?.client_id;
clientSecret = reg?.client_secret;
@ -96,7 +109,6 @@ export const useToolsOAuthFlow = ({
const verifier = generateCodeVerifier();
const challenge = await generateCodeChallenge(verifier);
const state = crypto.randomUUID();
const redirectUri = buildCallbackUrl();
const scopeString = scopes?.filter((s) => s.trim()).join(" ");
const authorizeUrl = buildMcpOAuthAuthorizeUrl({
@ -129,7 +141,7 @@ export const useToolsOAuthFlow = ({
setStatus("error");
NotificationsManager.error(msg);
}
}, [accessToken, serverId, serverAlias, scopes, preClientId]);
}, [accessToken, serverId, serverAlias, scopes, preClientId, gatewayMintsClient]);
const resumeOAuthFlow = useCallback(async () => {
if (typeof window === "undefined" || processingRef.current) return;

View file

@ -13860,11 +13860,12 @@ export interface paths {
* Patch Team
* @description Partially update a team using RFC 7386 JSON Merge Patch semantics.
*
* `team_id` is taken from the path. `metadata` is merged with the team's stored
* metadata rather than replacing it: an omitted key is preserved, `key: null`
* deletes it, and any other value overwrites (recursing into nested objects).
* Every other field behaves exactly like `POST /team/update` (omitted preserves,
* a value overwrites). Returns the full updated team.
* `team_id` is taken from the path; a `team_id` in the body is accepted only when it
* matches. `metadata` is merged with the team's stored metadata rather than replacing
* it: an omitted key is preserved, `key: null` deletes it, and any other value
* overwrites (recursing into nested objects). Every other field behaves exactly like
* `POST /team/update` (omitted preserves, a value overwrites). Returns the full
* updated team.
*
* ```
* curl --location --request PATCH 'http://0.0.0.0:4000/team/8d916b1c-510d-4894-a334-1c16a93344f5' --header 'Authorization: Bearer sk-1234' --header 'Content-Type: application/json' --data-raw '{
@ -28807,6 +28808,102 @@ export interface components {
litellm_params?: components["schemas"]["PromptLiteLLMParams"] | null;
prompt_info?: components["schemas"]["PromptInfo"] | null;
};
/**
* PatchTeamRequest
* @description Body of PATCH /team/{team_id}.
*
* Identical to UpdateTeamRequest except team_id is optional, because PATCH takes it
* from the path. A team_id in the body is still accepted when it matches the path.
*/
PatchTeamRequest: {
/** Access Group Ids */
access_group_ids?: string[] | null;
/** Allowed Passthrough Routes */
allowed_passthrough_routes?: unknown[] | null;
/** Allowed Vector Store Indexes */
allowed_vector_store_indexes?: components["schemas"]["AllowedVectorStoreIndexItem"][] | null;
/** Blocked */
blocked?: boolean | null;
/** Budget Duration */
budget_duration?: string | null;
/** Budget Limits */
budget_limits?: components["schemas"]["BudgetLimitEntry"][] | null;
/** Default Team Member Models */
default_team_member_models?: string[] | null;
/** Disable Global Guardrails */
disable_global_guardrails?: boolean | null;
/** Enforced Batch Output Expires After */
enforced_batch_output_expires_after?: {
[key: string]: unknown;
} | null;
/** Enforced File Expires After */
enforced_file_expires_after?: {
[key: string]: unknown;
} | null;
/** Guardrails */
guardrails?: string[] | null;
/** Max Budget */
max_budget?: number | null;
/** Mcp Rpm Limit */
mcp_rpm_limit?: {
[key: string]: number;
} | null;
/** Metadata */
metadata?: {
[key: string]: unknown;
} | null;
/** Model Aliases */
model_aliases?: {
[key: string]: unknown;
} | null;
/** Model Rpm Limit */
model_rpm_limit?: {
[key: string]: number;
} | null;
/** Model Tpm Limit */
model_tpm_limit?: {
[key: string]: number;
} | null;
/** Models */
models?: unknown[] | null;
object_permission?: components["schemas"]["LiteLLM_ObjectPermissionBase"] | null;
/** Organization Id */
organization_id?: string | null;
/** Policies */
policies?: string[] | null;
/** Prompts */
prompts?: string[] | null;
/** Router Settings */
router_settings?: {
[key: string]: unknown;
} | null;
/** Rpm Limit */
rpm_limit?: number | null;
/** Secret Manager Settings */
secret_manager_settings?: {
[key: string]: unknown;
} | null;
/** Soft Budget */
soft_budget?: number | null;
/** Tags */
tags?: unknown[] | null;
/** Team Alias */
team_alias?: string | null;
/** Team Id */
team_id?: string | null;
/** Team Member Budget */
team_member_budget?: number | null;
/** Team Member Budget Duration */
team_member_budget_duration?: string | null;
/** Team Member Key Duration */
team_member_key_duration?: string | null;
/** Team Member Rpm Limit */
team_member_rpm_limit?: number | null;
/** Team Member Tpm Limit */
team_member_tpm_limit?: number | null;
/** Tpm Limit */
tpm_limit?: number | null;
};
/**
* PerTestingCriteriaResult
* @description Results for a specific testing criteria
@ -50613,7 +50710,11 @@ export interface operations {
};
cookie?: never;
};
requestBody?: never;
requestBody: {
content: {
"application/json": components["schemas"]["PatchTeamRequest"];
};
};
responses: {
/** @description Successful Response */
200: {

View file

@ -0,0 +1,47 @@
import { RuleTester } from "eslint";
import rule from "../../scripts/eslint-rules/filename-pascal-case.mjs";
const ruleTester = new RuleTester({
languageOptions: { ecmaVersion: "latest", sourceType: "module", parserOptions: { ecmaFeatures: { jsx: true } } },
});
ruleTester.run("filename-pascal-case", rule as never, {
valid: [
{ code: "export const x = 1;", filename: "src/components/TeamInfo.tsx" },
{ code: "export const x = 1;", filename: "src/components/Button.tsx" },
{ code: "export const x = 1;", filename: "src/app/teams/page.tsx" },
{ code: "export const x = 1;", filename: "src/app/teams/layout.tsx" },
{ code: "export const x = 1;", filename: "src/app/teams/not-found.tsx" },
{ code: "export const x = 1;", filename: "src/app/global-error.tsx" },
{ code: "export const x = 1;", filename: "src/app/apple-icon.tsx" },
{ code: "export const x = 1;", filename: "src/app/opengraph-image.tsx" },
{ code: "export const x = 1;", filename: "src/app/twitter-image.tsx" },
{ code: "export const x = 1;", filename: "src/components/TeamInfo.test.tsx" },
{ code: "export const x = 1;", filename: "src/components/view_users.test.tsx" },
{ code: "export const x = 1;", filename: "src/components/TeamInfo.spec.tsx" },
],
invalid: [
{
code: "export const x = 1;",
filename: "src/components/user_info_view.tsx",
errors: [{ messageId: "notPascalCase", data: { name: "user_info_view.tsx", suggestion: "UserInfoView" } }],
},
{
code: "export const x = 1;",
filename: "src/components/team-info.tsx",
errors: [{ messageId: "notPascalCase", data: { name: "team-info.tsx", suggestion: "TeamInfo" } }],
},
{
code: "export const x = 1;",
filename: "src/components/teamInfo.tsx",
errors: [{ messageId: "notPascalCase", data: { name: "teamInfo.tsx", suggestion: "TeamInfo" } }],
},
{
code: "export const x = 1;",
filename: "src/components/my-component.utils.tsx",
errors: [
{ messageId: "notPascalCase", data: { name: "my-component.utils.tsx", suggestion: "MyComponent.utils" } },
],
},
],
});

View file

@ -0,0 +1,34 @@
import { RuleTester } from "eslint";
import rule from "../../scripts/eslint-rules/no-complex-jsx-arrow.mjs";
const ruleTester = new RuleTester({
languageOptions: { ecmaVersion: "latest", sourceType: "module", parserOptions: { ecmaFeatures: { jsx: true } } },
});
ruleTester.run("no-complex-jsx-arrow", rule as never, {
valid: [
"const x = <button onClick={() => doThing()} />;",
"const x = <button onClick={() => { a(); }} />;",
"const x = <button onClick={() => { a(); b(); }} />;",
"const handler = () => { a(); b(); c(); }; const x = <button onClick={handler} />;",
"const run = () => { a(); b(); c(); };",
"foo(() => { a(); b(); c(); });",
"const x = <List renderItem={(i) => i.name} />;",
{ code: "const x = <button onClick={() => { a(); b(); c(); }} />;", options: [{ maxStatements: 3 }] },
],
invalid: [
{
code: "const x = <button onClick={() => { a(); b(); c(); }} />;",
errors: [{ messageId: "tooComplex", data: { count: 3, max: 2 } }],
},
{
code: "const x = <form onSubmit={() => { a(); b(); c(); d(); }} />;",
errors: [{ messageId: "tooComplex", data: { count: 4, max: 2 } }],
},
{
code: "const x = <button onClick={() => { a(); b(); }} />;",
options: [{ maxStatements: 1 }],
errors: [{ messageId: "tooComplex", data: { count: 2, max: 1 } }],
},
],
});