mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_anthropic_output_format_remaining_keywords
This commit is contained in:
commit
16dad256d4
57 changed files with 6734 additions and 1306 deletions
6
.github/workflows/test-litellm-ui-unit.yml
vendored
6
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]],
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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: {
|
||||
|
|
|
|||
162
ui/litellm-dashboard/package-lock.json
generated
162
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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,
|
||||
},
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
.
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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"
|
||||
>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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?"
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -6737,6 +6737,7 @@ interface RegisterMcpOAuthClientPayload {
|
|||
grant_types?: string[];
|
||||
response_types?: string[];
|
||||
token_endpoint_auth_method?: string;
|
||||
redirect_uris?: string[];
|
||||
}
|
||||
|
||||
export const registerMcpOAuthClient = async (
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
89
ui/litellm-dashboard/src/hooks/useToolsOAuthFlow.test.tsx
Normal file
89
ui/litellm-dashboard/src/hooks/useToolsOAuthFlow.test.tsx
Normal 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" }),
|
||||
);
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
113
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
113
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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: {
|
||||
|
|
|
|||
|
|
@ -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" } },
|
||||
],
|
||||
},
|
||||
],
|
||||
});
|
||||
|
|
@ -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 } }],
|
||||
},
|
||||
],
|
||||
});
|
||||
Loading…
Add table
Reference in a new issue