Merge branch 'litellm_internal_staging' into litellm_add_more_tests

This commit is contained in:
mubashir1osmani 2026-07-14 14:05:55 -07:00 committed by GitHub
commit 50edc2c283
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
47 changed files with 2312 additions and 1040 deletions

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.49"
version = "0.1.50"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.49"
version = "0.1.50"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.76"
version = "0.4.77"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.76"
version = "0.4.77"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -515,6 +515,22 @@ class CustomGuardrail(CustomLogger):
return True
return False
def uses_apply_guardrail_interface(self) -> bool:
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
def _deployment_pre_call_target(self) -> "CustomLogger":
if not self.uses_apply_guardrail_interface():
return self
try:
from litellm.proxy.utils import unified_guardrail
except ImportError as e:
raise ImportError(
f"Guardrail {self.guardrail_name or type(self).__name__} implements apply_guardrail, which needs "
"the litellm proxy dependencies to run at the deployment level. "
"Install them with: pip install 'litellm[proxy]'"
) from e
return unified_guardrail
async def async_pre_call_deployment_hook(
self, kwargs: Dict[str, Any], call_type: Optional[CallTypes]
) -> Optional[dict]:
@ -533,7 +549,10 @@ class CustomGuardrail(CustomLogger):
# CHECK IF GUARDRAIL REJECTS THE REQUEST
if call_type == CallTypes.completion or call_type == CallTypes.acompletion:
result = await self.async_pre_call_hook(
target = self._deployment_pre_call_target()
if target is not self:
kwargs["guardrail_to_apply"] = self
result = await target.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id=kwargs.get("user_api_key_user_id"),
team_id=kwargs.get("user_api_key_team_id"),
@ -543,7 +562,7 @@ class CustomGuardrail(CustomLogger):
),
cache=dc,
data=kwargs,
call_type=call_type.value or "acompletion", # type: ignore
call_type="completion" if call_type == CallTypes.completion else "acompletion",
)
if result is not None and isinstance(result, dict):

View file

@ -375,6 +375,36 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
else:
optional_params.pop("output_config", None)
@staticmethod
def _drop_incompatible_temperature_for_thinking(
model: str, optional_params: dict, custom_llm_provider: str
) -> None:
"""Anthropic rejects any ``temperature`` other than 1 while extended thinking
is enabled ("temperature may only be set to 1 when thinking is enabled").
Clients like Claude Code send ``thinking``/``output_config.effort`` together
with a pinned ``temperature`` (e.g. the safety classifier uses ``temperature=0``
for determinism). When the request lands on a non-adaptive model, the effort
interface is reshaped above into legacy ``thinking={type: enabled}`` (or kept
as ``output_config.effort`` on Opus 4.5), and the leftover ``temperature`` would
400. Preserving the thinking the caller asked for wins over an unhonorable
sampling value (Anthropic forces ``temperature=1`` under thinking regardless),
so drop it and let the API default apply.
Adaptive models (4.6+) own this natively and are left untouched.
"""
if AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider):
return
temperature = optional_params.get("temperature")
if temperature is None or temperature == 1:
return
thinking = optional_params.get("thinking")
output_config = optional_params.get("output_config")
thinking_enabled = isinstance(thinking, dict) and thinking.get("type") == "enabled"
effort_enabled = isinstance(output_config, dict) and output_config.get("effort") is not None
if thinking_enabled or effort_enabled:
optional_params.pop("temperature", None)
def transform_anthropic_messages_request(
self,
model: str,
@ -415,6 +445,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
custom_llm_provider=self._resolved_provider,
)
self._drop_incompatible_temperature_for_thinking(
model=model,
optional_params=anthropic_messages_optional_request_params,
custom_llm_provider=self._resolved_provider,
)
system_param = anthropic_messages_optional_request_params.get("system")
if self.should_strip_billing_metadata() and system_param is not None:
filtered_system = self._filter_billing_headers_from_system(system_param)

View file

@ -0,0 +1,694 @@
"""Bridge token flow: litellm identity resolution and the DCR-bridge oauth_delegate mint/refresh pipeline."""
import math
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Literal, Optional
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse
from pydantic import SecretStr
from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _BridgeAuthorizationCode
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
EnvelopeIdentity,
EnvelopeKeys,
RefreshCredential,
UpstreamTokenGrant,
)
from litellm.proxy._types import UserAPIKeyAuth
def _litellm_key_from_request(request: Request) -> Optional[str]:
"""Return the LiteLLM API key presented on the request, or ``None``.
Accepts the key from ``x-litellm-api-key`` (what MCP clients such as Claude Desktop/Code
send) as well as ``Authorization``; either may carry a bare token or ``Bearer <token>``.
``x-litellm-api-key`` wins when both are present, since ``Authorization`` may instead carry
an OAuth/upstream bearer.
"""
for header_value in (
request.headers.get("x-litellm-api-key"),
request.headers.get("Authorization") or request.headers.get("authorization"),
):
if not header_value:
continue
value = header_value.strip()
if value.lower().startswith("bearer "):
value = value[7:].strip()
if value:
return value
return None
def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool:
"""``True`` when the presented key is neither blocked nor past its expiry.
The OAuth token endpoint is unauthenticated, so the presented key is validated here before it is
trusted; a revoked or expired key must not mint a bridge envelope or write a stored credential.
``get_key_object`` resolves a row without these checks (the main ``user_api_key_auth`` pipeline
enforces them downstream, which this endpoint bypasses), so they are applied here. Deleted keys
are already rejected upstream, where ``get_key_object`` raises on a row that no longer exists.
This is an active-state gate only; it deliberately does not require a ``user_id``. A valid
team-scoped or service-account key has no ``user_id`` yet is a legitimate credential, so gating
on ``user_id`` presence would wrongly reject it. Callers that need the user (the per-user token
store) derive it separately via :func:`_active_key_user_id`.
Total by design: ``expires`` is typed ``str | datetime``, and an unparseable string would make
``datetime.fromisoformat`` raise. Since the callers run this outside their key-resolution
``try``, an uncaught parse error would surface as a 500 instead of the endpoint's fail-closed
behavior, so a malformed expiry is treated as inactive (return ``False``) rather than raising.
"""
if key_obj.blocked is True:
return False
expires = key_obj.expires
if expires is not None:
if isinstance(expires, datetime):
expiry = expires
else:
try:
expiry = datetime.fromisoformat(expires)
except (ValueError, TypeError):
return False
if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None:
expiry = expiry.replace(tzinfo=timezone.utc)
if expiry < datetime.now(timezone.utc):
return False
return True
def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None:
"""The active key's ``user_id``, or ``None`` when the key is blocked/expired or simply has no
``user_id`` (a team-scoped or service-account key). Used only by the per-user token store, which
needs a user to key the stored credential; the bridge mint uses the key hash and does not."""
return key_obj.user_id if _key_is_active(key_obj) else None
@dataclass(frozen=True, slots=True)
class _ResolvedKey:
"""An active litellm key resolved from the token request: its hash (the value ``get_key_object``
and the cache/DB layer key the record by) and the live record."""
key_hash: str
key: "UserAPIKeyAuth"
_KeyResolutionFailure = Literal["no_active_key", "unavailable", "unresolvable"]
"""Why a token request yielded no active litellm key, kept distinct so a caller statuses each truthfully
instead of blaming the client for a gateway problem:
- ``no_active_key``: none was presented, or the presented key is unknown / blocked / expired (the
caller's request is at fault)
- ``unavailable``: the auth database was transiently unreachable while resolving (retryable)
- ``unresolvable``: the gateway cannot resolve identity right now (no DB connection, or an unexpected
error) -- a gateway fault, not the caller's
The classification mirrors admission's ``_reload_admitted_key`` so the mint (ingress) and admission
(egress) never disagree on the status of the same outage."""
async def _resolve_active_litellm_key(request: Request) -> "_ResolvedKey | _KeyResolutionFailure":
"""Resolve the presented litellm key to an active key record, or say precisely why not.
Single resolution path the OAuth token endpoint reuses, resolving authoritatively via
``get_key_object`` (cache first, then DB). The failure is a value, not a bare ``None``, so a caller
can tell "the client sent no usable credential" (a request error) apart from "the gateway could not
check" (an infrastructure error) and status each truthfully; collapsing both to ``None`` is what let
a DB outage read as a 400. A resolved key is still gated by ``_key_is_active``, so a blocked or
expired key is ``no_active_key`` while a valid team-scoped or service-account key (no ``user_id``)
resolves. Classification mirrors admission's ``_reload_admitted_key``: no DB connection is a gateway
fault, a ``ProxyException`` / ``HTTPException`` from ``get_key_object`` is an unknown or invalid key,
a database-service-unavailable error is a retryable outage, and anything else is an unexpected
gateway fault."""
token = _litellm_key_from_request(request)
if not token:
return "no_active_key"
from litellm.proxy._types import hash_token # noqa: PLC0415 # inline import avoids a module-load circular import
return await _reload_active_key_by_hash(hash_token(token))
async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResolutionFailure":
"""Reload the live key record for ``key_hash`` (cache first, then DB) and gate it on active state,
returning the resolved key or a precise failure. Shared by the token request's presented-key
resolution (:func:`_resolve_active_litellm_key`, which hashes the presented key) and the refresh
path (which already holds the hash sealed in the refresh envelope), so both re-validate identity
through one active-key gate and one failure classification. Classification mirrors admission's
``_reload_admitted_key``: no DB connection is a gateway fault, a ``ProxyException`` / ``HTTPException``
from ``get_key_object`` is an unknown or invalid key, a database-service-unavailable error is a
retryable outage, and anything else is an unexpected gateway fault. A blocked or expired key is
``no_active_key``, so a revoked key can neither mint nor refresh a bridge envelope."""
from litellm.proxy._types import (
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
)
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_key_object,
)
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
PrismaDBExceptionHandler,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return "unresolvable"
try:
key_obj = await get_key_object(
hashed_token=key_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
except (ProxyException, HTTPException):
return "no_active_key"
except Exception as exc: # noqa: BLE001 # classify: a DB outage is retryable, anything else is an opaque gateway fault
if PrismaDBExceptionHandler.is_database_service_unavailable_error(exc):
return "unavailable"
verbose_logger.debug(
"_reload_active_key_by_hash: unexpected key-resolution error (%s)",
type(exc).__name__,
)
return "unresolvable"
if not _key_is_active(key_obj):
return "no_active_key"
return _ResolvedKey(key_hash=key_hash, key=key_obj)
async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None":
"""Re-validate a live litellm user by id, returning ``None`` when the user is active or a precise
failure otherwise. The interactive DCR client authenticates via SSO, so its refresh envelope seals a
user subject; renewing it must re-check the user is still live (present and not SCIM-deactivated) so a
deactivated user cannot keep refreshing, mirroring how admission re-validates the same user subject on
the egress side. No DB connection is a gateway fault (``unresolvable``) and a
database-service-unavailable error is a retryable outage (``unavailable``). Everything else fails
closed as ``no_active_key`` (the caller maps it to invalid_grant): a ``ProxyException`` /
``HTTPException``, a SCIM-deactivated user, and, unlike the key path, a missing user. ``get_user_object``
catches every DB failure and re-raises a bare ``ValueError`` (a deleted user and a real outage look
identical, the original error surviving only as ``__context__``), so the outage check walks the cause
chain, and a missing user falls through to ``no_active_key`` rather than an opaque gateway fault."""
from litellm.proxy._types import (
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
)
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_user_object,
)
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
PrismaDBExceptionHandler,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return "unresolvable"
try:
user_object = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except (ProxyException, HTTPException):
return "no_active_key"
except Exception as exc: # noqa: BLE001 # a DB outage is retryable; a missing user (get_user_object's wrapped ValueError) or any other resolution failure fails closed as no_active_key, never a 500
if PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(exc):
return "unavailable"
verbose_logger.debug("_reload_active_user_by_id: user-resolution error (%s)", type(exc).__name__)
return "no_active_key"
if user_object is None:
return "no_active_key"
if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False:
return "no_active_key"
return None
async def _key_owner_scim_deactivated(key: "UserAPIKeyAuth") -> bool:
"""True only when the key's owning user was explicitly SCIM-deactivated, so a refresh revokes an
offboarded owner's key exactly as admission does via ``_reject_if_admitted_owner_scim_deactivated``.
A key with no owner, a missing owner record, or a failed lookup fails OPEN (returns ``False``),
matching admission and the standard builder: a key may outlive its owner record, and a transient DB
blip must not revoke a live key. Only an explicit ``scim_active`` of ``False`` gates renewal."""
if key.user_id is None:
return False
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_user_object,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return False
try:
owner = await get_user_object(
user_id=key.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except Exception as exc: # noqa: BLE001 # fail open: a missing owner (get_user_object's wrapped ValueError) or a DB blip must not revoke a live key
verbose_logger.debug("refresh: key-owner SCIM lookup failed, not revoking (%s)", type(exc).__name__)
return False
return owner is not None and isinstance(owner.metadata, dict) and owner.metadata.get("scim_active") is False
async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResolutionFailure | None":
"""Re-validate that the subject sealed in a refresh envelope is still live, dispatching on its type:
a key_hash reloads the virtual key, a user_id reloads the user. Returns ``None`` when the subject is
active or a precise failure otherwise, so revocation gates renewal for either identity source the same
way admission gates the egress: a blocked or expired key, a SCIM-deactivated key owner (mirroring
admission's owner check, so an offboarded user cannot keep renewing a still-active key), and a
deactivated or deleted user all fail closed to ``no_active_key``."""
match identity.subject_type:
case "key_hash":
reloaded = await _reload_active_key_by_hash(identity.subject)
if not isinstance(reloaded, _ResolvedKey):
return reloaded
if await _key_owner_scim_deactivated(reloaded.key):
return "no_active_key"
return None
case "user_id":
return await _reload_active_user_by_id(identity.subject)
case _:
assert_never(identity.subject_type)
async def _extract_user_id_from_request(request: Request) -> str | None:
"""The litellm ``user_id`` for the token request, so a per-user token is stored under the same
identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome
(including a transient DB outage) collapses to ``None`` here and the caller simply skips the store;
the bridge mint, which must status those outcomes differently, consumes
:func:`_resolve_active_litellm_key` directly."""
resolved = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
return None
return _active_key_user_id(resolved.key)
_UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"]
"""Why an upstream token response cannot back a bridge envelope:
- ``no_access_token``: the response carries no usable ``access_token``
- ``expired_lifetime``: the response reports a parseable, non-positive ``expires_in``, i.e. an upstream
token that is already dead, so sealing it would forward a bearer the edge cannot use
An absent or unparseable ``expires_in`` is NOT a rejection; the lifetime is merely unknown and the
envelope caps it, the by-design behaviour for an upstream that omits the field."""
def _classify_upstream_lifetime(raw_expires_in: object) -> "int | Literal['unspecified', 'expired']":
"""Classify an upstream ``expires_in`` into a positive number of seconds, ``"unspecified"`` (absent
or unparseable, so the envelope caps it), or ``"expired"`` (a non-positive value the upstream reports
as already elapsed). Telling "we do not know the lifetime" apart from "the upstream says it is
already dead" is what stops an explicitly-expired token from silently receiving the envelope's 1h
cap. The expired decision is made on the parsed numeric value, not on ``int(...)`` of it, so a
positive sub-second lifetime in ``(0, 1)`` is not truncated to ``0`` and misread as elapsed; the
envelope works in whole seconds, so such a lifetime clamps up to its 1s floor. ``bool`` is excluded
(an ``int`` subclass but never a real lifetime), and the conversions can raise on ``NaN`` /
``Infinity`` / oversized input, which reads as unparseable rather than surfacing as a 500."""
if raw_expires_in is None or isinstance(raw_expires_in, bool) or not isinstance(raw_expires_in, (int, float, str)):
return "unspecified"
try:
numeric = float(raw_expires_in)
seconds = int(numeric)
except (ValueError, TypeError, OverflowError):
return "unspecified"
if numeric <= 0:
return "expired"
return max(1, seconds)
def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenGrant | _UpstreamGrantRejection":
"""Validate an upstream OAuth token response into a typed grant, or say why it cannot back an
envelope. Each field is isinstance-checked so nothing untyped from ``response.json()`` reaches the
grant. ``expires_in`` is read three ways (see :func:`_classify_upstream_lifetime`): an unknown
lifetime leaves the grant ``expires_in`` ``None`` for the envelope to cap, a positive value is
honoured, and an explicit already-elapsed value is a rejection rather than a silent fall-through to
the cap."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
UpstreamTokenGrant,
)
if not isinstance(token_response, dict):
return "no_access_token"
access = token_response.get("access_token")
if not isinstance(access, str) or not access:
return "no_access_token"
lifetime = _classify_upstream_lifetime(token_response.get("expires_in"))
if lifetime == "expired":
return "expired_lifetime"
token_type = token_response.get("token_type")
scope = token_response.get("scope")
return UpstreamTokenGrant(
access_token=SecretStr(access),
token_type=token_type if isinstance(token_type, str) and token_type else "Bearer",
# The upstream refresh_token is deliberately NOT sealed: the edge never consumes it (it forwards
# only token_type + access_token), so it would be dead weight embedding a long-lived upstream
# credential in the client-held bearer, and it enlarges the envelope. Refresh support is a
# follow-up (a dedicated refresh-envelope); the client re-runs authorization_code at the cap.
refresh_token=None,
scope=scope if isinstance(scope, str) and scope else None,
expires_in=lifetime if isinstance(lifetime, int) else None,
)
# ---------------------------------------------------------------------------
# DCR-bridge oauth_delegate mint: a three-phase pipeline whose failures are values.
#
# prepare (before the upstream exchange) -> validate every precondition and resolve identity+keys
# exchange (the single-use upstream code is consumed here, in exchange_token_with_server)
# finish (after the exchange) -> seal the upstream grant into the client-held envelope
#
# Every precondition lives in ``prepare``, which runs BEFORE the exchange, so no failure can burn the
# single-use code or rotate a refresh token, for either grant type -- that whole class of bug is gone
# by construction rather than guarded case by case. Failures are values mapped to an OAuth-shaped
# response in one place (``_bridge_mint_error_response``), so status codes and the RFC 6749 §5.2 body
# shape are uniform. Adding a failure mode is a new literal plus a match arm the type checker forces.
# ---------------------------------------------------------------------------
_BridgeMintError = Literal[
"no_identity",
"invalid_refresh",
"identity_unavailable",
"identity_unresolvable",
"not_configured",
"no_upstream_token",
"upstream_token_expired",
"too_large",
]
@dataclass(frozen=True, slots=True)
class _BridgeMintReady:
"""Everything the seal needs, resolved once before the exchange: the identity to bind the envelope
to and the master-key-derived envelope keys. The identity is a key_hash subject for the scripted
two-header client (resolved from the litellm key it presents) or a user_id subject for the
interactive SSO client (the user recovered from the gateway authorization code), so one phase-3 seal
serves both. Resolving identity here means ``_finish_bridge_mint`` has no preconditions left to
fail."""
identity: "EnvelopeIdentity"
keys: "EnvelopeKeys"
def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
"""Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape
(top-level ``error``, no-store headers) for every case, with a status truthful about where the
failure is. The caller's request is 400, a transient gateway outage is 503, a gateway
misconfiguration is 500, and an upstream problem is 502. The identity-resolution statuses match how
admission statuses the same conditions on the egress side, so mint and admit never disagree under
one outage."""
match error:
case "no_identity":
status, code, desc = (
400,
"invalid_request",
"this server issues a gateway-bound credential; complete the interactive sign-in, or "
"send a litellm credential (x-litellm-api-key or Authorization) on the token request",
)
case "invalid_refresh":
status, code, desc = (
400,
"invalid_grant",
"the refresh credential is not a valid, live refresh envelope for this server; "
"re-run authorization_code to obtain a new one",
)
case "identity_unavailable":
status, code, desc = (
503,
"temporarily_unavailable",
"the authentication database is temporarily unreachable; retry shortly",
)
case "identity_unresolvable":
status, code, desc = (
500,
"server_error",
"the gateway could not resolve the litellm identity for this request",
)
case "not_configured":
status, code, desc = (
500,
"server_error",
"the gateway is not configured to mint a gateway-bound credential (master_key is not set)",
)
case "no_upstream_token":
status, code, desc = (
502,
"server_error",
"the upstream token response has no usable access_token",
)
case "upstream_token_expired":
status, code, desc = (
502,
"server_error",
"the upstream token response reports an already-expired lifetime",
)
case "too_large":
status, code, desc = (
502,
"server_error",
"the upstream token is too large to seal into a gateway-bound credential",
)
case _:
assert_never(error)
return JSONResponse(
status_code=status, content={"error": code, "error_description": desc}, headers=TOKEN_NO_CACHE_HEADERS
)
def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
"""Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays
truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that
cannot resolve identity is 500."""
match failure:
case "no_active_key":
return "no_identity"
case "unavailable":
return "identity_unavailable"
case "unresolvable":
return "identity_unresolvable"
case _:
assert_never(failure)
def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _BridgeMintError:
"""Lift an upstream-response rejection into the mint taxonomy; both are upstream faults (502)."""
match rejection:
case "no_access_token":
return "no_upstream_token"
case "expired_lifetime":
return "upstream_token_expired"
case _:
assert_never(rejection)
async def _prepare_bridge_mint(
request: Request,
mcp_server: MCPServer,
bridge_identity: "_BridgeAuthorizationCode | None" = None,
) -> "_BridgeMintReady | _BridgeMintError":
"""Phase 1 for the authorization_code grant, BEFORE the upstream exchange: confirm the gateway can
mint (master_key set), resolve the litellm identity, and derive the envelope keys. Returns a ready
context or a precise failure value. Running before the exchange is what makes every failure here fail
closed without consuming the single-use code.
Two identity sources, one envelope. The interactive DCR client authenticates via SSO at the bridged
authorize, so its identity arrives as ``bridge_identity`` (the user recovered from the gateway
authorization code) and mints a user subject. The scripted two-header client presents a litellm key
on the token request instead, so its identity is the active key's hash and mints a key_hash subject.
A missing or invalid presented key keeps its resolution origin so the mapper statuses it truthfully;
neither source present is ``no_identity``. The refresh_token grant has its own phase-1
(:func:`_prepare_bridge_refresh`), which recovers identity from the presented refresh envelope."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
envelope_keys_from_master_key,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
key_hash_identity,
user_identity,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
master_key,
)
if not master_key:
return "not_configured"
keys = envelope_keys_from_master_key(master_key)
if bridge_identity is not None:
identity = user_identity(server_id=mcp_server.server_id, user_id=bridge_identity.litellm_user_id)
return _BridgeMintReady(identity=identity, keys=keys)
resolved = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
return _key_resolution_failure_to_mint_error(resolved)
identity = key_hash_identity(server_id=mcp_server.server_id, key_hash=resolved.key_hash)
return _BridgeMintReady(identity=identity, keys=keys)
@dataclass(frozen=True, slots=True)
class _BridgeRefreshReady:
"""A validated refresh request: the identity+keys to mint the renewed pair under, the upstream refresh
token (unwrapped from the client's refresh envelope) to exchange with the upstream IdP, and the scope
sealed alongside it at mint. The upstream refresh token is a ``SecretStr`` like every other credential
in this layer, so a repr or a traceback that captures this value never exposes the raw upstream refresh
token in plaintext. ``upstream_scope`` carries the originally-granted scope so the renewal re-requests
it when the client (a DCR/MCP client that typically omits scope on refresh) sends none, keeping the
renewed token's scope stable against an upstream that would otherwise narrow or drop it."""
ready: "_BridgeMintReady"
upstream_refresh_token: SecretStr
upstream_scope: str | None = None
def _refresh_key_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
"""Lift an identity-resolution failure on the refresh path into the mint taxonomy. Unlike the mint
path, a resolved-but-inactive (or unknown) key is ``invalid_grant`` rather than ``invalid_request``:
the client did present an identity (sealed in the refresh envelope), but it is no longer live, so the
refresh is invalid and the client must re-authenticate. A transient outage is still 503 and a gateway
fault still 500, matching the mint path and admission."""
match failure:
case "no_active_key":
return "invalid_refresh"
case "unavailable":
return "identity_unavailable"
case "unresolvable":
return "identity_unresolvable"
case _:
assert_never(failure)
async def _prepare_bridge_refresh(
mcp_server: MCPServer, refresh_value: str | None
) -> "_BridgeRefreshReady | _BridgeMintError":
"""Phase 1 for the refresh_token grant, BEFORE the upstream exchange: open the client's refresh
envelope, re-validate the sealed litellm identity so a revoked key cannot keep refreshing, and
recover the upstream refresh token to exchange. Identity comes entirely from the sealed envelope, not
the HTTP request, so the request object is not needed here. The client presents a refresh envelope,
never a raw upstream refresh token, so a missing value, a non-envelope, an unopenable envelope, or one
minted for another server is ``invalid_grant``. Running before the exchange means a rejected refresh
never consumes or rotates the upstream refresh token."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
BridgeRefreshOpened,
envelope_keys_from_master_key,
open_bridge_refresh_envelope,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
master_key,
)
if not master_key:
return "not_configured"
if not refresh_value:
return "invalid_refresh"
keys = envelope_keys_from_master_key(master_key)
opened = open_bridge_refresh_envelope(refresh_value, keys, datetime.now(timezone.utc), mcp_server.server_id)
if not isinstance(opened, BridgeRefreshOpened):
return "invalid_refresh"
failure = await _revalidate_active_subject(opened.identity)
if failure is not None:
return _refresh_key_failure_to_mint_error(failure)
return _BridgeRefreshReady(
ready=_BridgeMintReady(identity=opened.identity, keys=keys),
upstream_refresh_token=opened.refresh.refresh_token,
upstream_scope=opened.refresh.scope,
)
def _finish_bridge_mint(
ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime
) -> "JSONResponse | _BridgeMintError":
"""Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held access envelope
using the pre-resolved identity and keys, and, when the upstream returned a refresh token, seal a
long-lived refresh envelope alongside it so the client can renew without re-authenticating. Shared by
the authorization_code and refresh_token paths, so a renewal that the upstream rotates re-issues a
fresh refresh envelope. The only hard failures here are properties of the upstream access token (no
usable token, an already-expired lifetime, or a token too large to seal); a refresh token that cannot
be sealed degrades to an access-only response rather than failing the whole exchange."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
build_bridge_token_response,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
SealedEnvelope,
UpstreamTokenGrant,
)
grant = _bridge_grant_from_token_response(token_response)
if not isinstance(grant, UpstreamTokenGrant):
return _upstream_rejection_to_mint_error(grant)
sealed = build_bridge_token_response(ready.identity, grant, ready.keys, now)
if not isinstance(sealed, SealedEnvelope):
return "too_large"
# Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the
# client is never told the bearer lives past the point admission (which uses that exp) rejects it.
expires_in = max(0, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp()))
refresh_envelope = _mint_refresh_envelope_value(ready.identity, token_response, ready.keys, now, mcp_server)
body = {
"access_token": sealed.token.get_secret_value(),
"token_type": "Bearer",
"expires_in": expires_in,
# A refresh envelope rides along only when the upstream returned a refresh token to seal; when it
# rotates on renewal, the client receives the new one and the old envelope's upstream token dies.
**({"refresh_token": refresh_envelope} if refresh_envelope is not None else {}),
}
return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS)
def _upstream_refresh_credential(token_response: object) -> "RefreshCredential | None":
"""Extract the upstream refresh grant from a token response, or ``None`` when there is none to seal.
Each field is isinstance-checked so nothing untyped reaches the refresh envelope; ``refresh_expires_in``
(the refresh token's own lifetime, when the upstream reports it) is classified like ``expires_in`` and
bounds the refresh envelope's TTL. An upstream that reports the refresh token itself as already elapsed
(``refresh_expires_in`` non-positive) yields ``None`` rather than a refresh envelope: sealing a dead
token would hand the client a full-TTL-capped envelope the IdP will reject, so the exchange degrades to
an access-only response (the client re-authenticates at access expiry), mirroring how
:func:`_bridge_grant_from_token_response` refuses an already-elapsed access token instead of capping it."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
RefreshCredential,
)
if not isinstance(token_response, dict):
return None
refresh = token_response.get("refresh_token")
if not isinstance(refresh, str) or not refresh:
return None
lifetime = _classify_upstream_lifetime(token_response.get("refresh_expires_in"))
if lifetime == "expired":
return None
scope = token_response.get("scope")
return RefreshCredential(
refresh_token=SecretStr(refresh),
scope=scope if isinstance(scope, str) and scope else None,
expires_in=lifetime if isinstance(lifetime, int) else None,
)
def _mint_refresh_envelope_value(
identity: "EnvelopeIdentity", token_response: object, keys: "EnvelopeKeys", now: datetime, mcp_server: MCPServer
) -> str | None:
"""Seal the upstream refresh grant (if any) into a refresh envelope and return its bearer string, or
``None`` when the upstream returned no refresh token or the refresh token is too large to seal. A
too-large refresh token degrades to an access-only response (logged) rather than failing an exchange
that already succeeded upstream: the client simply re-authenticates when the access envelope expires."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
build_bridge_refresh_token_response,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
SealedEnvelope,
)
refresh_credential = _upstream_refresh_credential(token_response)
if refresh_credential is None:
return None
sealed = build_bridge_refresh_token_response(identity, refresh_credential, keys, now)
if isinstance(sealed, SealedEnvelope):
return sealed.token.get_secret_value()
verbose_logger.warning(
"bridge mint: the upstream refresh token is too large to seal into a refresh envelope for "
"server=%s; issuing an access-only response, so the client re-authenticates at access expiry",
mcp_server.server_id,
)
return None

View file

@ -1,10 +1,8 @@
import asyncio
import html as _html
import json
import math
import secrets
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
@ -13,7 +11,6 @@ import httpx
from fastapi import APIRouter, Form, HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import (
@ -24,6 +21,15 @@ from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
TokenEndpointAuthConfigError,
build_token_endpoint_client_auth,
)
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_bridge_mint_error_response,
_BridgeMintReady,
_BridgeRefreshReady,
_extract_user_id_from_request,
_finish_bridge_mint,
_prepare_bridge_mint,
_prepare_bridge_refresh,
)
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
CredentialSource,
@ -49,13 +55,7 @@ from litellm.types.mcp import MCPAuth, MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
EnvelopeIdentity,
EnvelopeKeys,
RefreshCredential,
UpstreamTokenGrant,
)
from litellm.proxy._types import LiteLLM_MCPServerTable, UserAPIKeyAuth
from litellm.proxy._types import LiteLLM_MCPServerTable
# TTL cache for upstream OAuth metadata fetched from pass-through MCP servers.
# Keeps us from hammering the upstream IdP on each discovery request.
@ -393,274 +393,6 @@ def _validate_token_response(
)
def _litellm_key_from_request(request: Request) -> Optional[str]:
"""Return the LiteLLM API key presented on the request, or ``None``.
Accepts the key from ``x-litellm-api-key`` (what MCP clients such as Claude Desktop/Code
send) as well as ``Authorization``; either may carry a bare token or ``Bearer <token>``.
``x-litellm-api-key`` wins when both are present, since ``Authorization`` may instead carry
an OAuth/upstream bearer.
"""
for header_value in (
request.headers.get("x-litellm-api-key"),
request.headers.get("Authorization") or request.headers.get("authorization"),
):
if not header_value:
continue
value = header_value.strip()
if value.lower().startswith("bearer "):
value = value[7:].strip()
if value:
return value
return None
def _key_is_active(key_obj: "UserAPIKeyAuth") -> bool:
"""``True`` when the presented key is neither blocked nor past its expiry.
The OAuth token endpoint is unauthenticated, so the presented key is validated here before it is
trusted; a revoked or expired key must not mint a bridge envelope or write a stored credential.
``get_key_object`` resolves a row without these checks (the main ``user_api_key_auth`` pipeline
enforces them downstream, which this endpoint bypasses), so they are applied here. Deleted keys
are already rejected upstream, where ``get_key_object`` raises on a row that no longer exists.
This is an active-state gate only; it deliberately does not require a ``user_id``. A valid
team-scoped or service-account key has no ``user_id`` yet is a legitimate credential, so gating
on ``user_id`` presence would wrongly reject it. Callers that need the user (the per-user token
store) derive it separately via :func:`_active_key_user_id`.
Total by design: ``expires`` is typed ``str | datetime``, and an unparseable string would make
``datetime.fromisoformat`` raise. Since the callers run this outside their key-resolution
``try``, an uncaught parse error would surface as a 500 instead of the endpoint's fail-closed
behavior, so a malformed expiry is treated as inactive (return ``False``) rather than raising.
"""
if key_obj.blocked is True:
return False
expires = key_obj.expires
if expires is not None:
if isinstance(expires, datetime):
expiry = expires
else:
try:
expiry = datetime.fromisoformat(expires)
except (ValueError, TypeError):
return False
if expiry.tzinfo is None or expiry.tzinfo.utcoffset(expiry) is None:
expiry = expiry.replace(tzinfo=timezone.utc)
if expiry < datetime.now(timezone.utc):
return False
return True
def _active_key_user_id(key_obj: "UserAPIKeyAuth") -> str | None:
"""The active key's ``user_id``, or ``None`` when the key is blocked/expired or simply has no
``user_id`` (a team-scoped or service-account key). Used only by the per-user token store, which
needs a user to key the stored credential; the bridge mint uses the key hash and does not."""
return key_obj.user_id if _key_is_active(key_obj) else None
@dataclass(frozen=True, slots=True)
class _ResolvedKey:
"""An active litellm key resolved from the token request: its hash (the value ``get_key_object``
and the cache/DB layer key the record by) and the live record."""
key_hash: str
key: "UserAPIKeyAuth"
_KeyResolutionFailure = Literal["no_active_key", "unavailable", "unresolvable"]
"""Why a token request yielded no active litellm key, kept distinct so a caller statuses each truthfully
instead of blaming the client for a gateway problem:
- ``no_active_key``: none was presented, or the presented key is unknown / blocked / expired (the
caller's request is at fault)
- ``unavailable``: the auth database was transiently unreachable while resolving (retryable)
- ``unresolvable``: the gateway cannot resolve identity right now (no DB connection, or an unexpected
error) -- a gateway fault, not the caller's
The classification mirrors admission's ``_reload_admitted_key`` so the mint (ingress) and admission
(egress) never disagree on the status of the same outage."""
async def _resolve_active_litellm_key(request: Request) -> "_ResolvedKey | _KeyResolutionFailure":
"""Resolve the presented litellm key to an active key record, or say precisely why not.
Single resolution path the OAuth token endpoint reuses, resolving authoritatively via
``get_key_object`` (cache first, then DB). The failure is a value, not a bare ``None``, so a caller
can tell "the client sent no usable credential" (a request error) apart from "the gateway could not
check" (an infrastructure error) and status each truthfully; collapsing both to ``None`` is what let
a DB outage read as a 400. A resolved key is still gated by ``_key_is_active``, so a blocked or
expired key is ``no_active_key`` while a valid team-scoped or service-account key (no ``user_id``)
resolves. Classification mirrors admission's ``_reload_admitted_key``: no DB connection is a gateway
fault, a ``ProxyException`` / ``HTTPException`` from ``get_key_object`` is an unknown or invalid key,
a database-service-unavailable error is a retryable outage, and anything else is an unexpected
gateway fault."""
token = _litellm_key_from_request(request)
if not token:
return "no_active_key"
from litellm.proxy._types import hash_token # noqa: PLC0415 # inline import avoids a module-load circular import
return await _reload_active_key_by_hash(hash_token(token))
async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResolutionFailure":
"""Reload the live key record for ``key_hash`` (cache first, then DB) and gate it on active state,
returning the resolved key or a precise failure. Shared by the token request's presented-key
resolution (:func:`_resolve_active_litellm_key`, which hashes the presented key) and the refresh
path (which already holds the hash sealed in the refresh envelope), so both re-validate identity
through one active-key gate and one failure classification. Classification mirrors admission's
``_reload_admitted_key``: no DB connection is a gateway fault, a ``ProxyException`` / ``HTTPException``
from ``get_key_object`` is an unknown or invalid key, a database-service-unavailable error is a
retryable outage, and anything else is an unexpected gateway fault. A blocked or expired key is
``no_active_key``, so a revoked key can neither mint nor refresh a bridge envelope."""
from litellm.proxy._types import (
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
)
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_key_object,
)
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
PrismaDBExceptionHandler,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return "unresolvable"
try:
key_obj = await get_key_object(
hashed_token=key_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
except (ProxyException, HTTPException):
return "no_active_key"
except Exception as exc: # noqa: BLE001 # classify: a DB outage is retryable, anything else is an opaque gateway fault
if PrismaDBExceptionHandler.is_database_service_unavailable_error(exc):
return "unavailable"
verbose_logger.debug(
"_reload_active_key_by_hash: unexpected key-resolution error (%s)",
type(exc).__name__,
)
return "unresolvable"
if not _key_is_active(key_obj):
return "no_active_key"
return _ResolvedKey(key_hash=key_hash, key=key_obj)
async def _reload_active_user_by_id(user_id: str) -> "_KeyResolutionFailure | None":
"""Re-validate a live litellm user by id, returning ``None`` when the user is active or a precise
failure otherwise. The interactive DCR client authenticates via SSO, so its refresh envelope seals a
user subject; renewing it must re-check the user is still live (present and not SCIM-deactivated) so a
deactivated user cannot keep refreshing, mirroring how admission re-validates the same user subject on
the egress side. No DB connection is a gateway fault (``unresolvable``) and a
database-service-unavailable error is a retryable outage (``unavailable``). Everything else fails
closed as ``no_active_key`` (the caller maps it to invalid_grant): a ``ProxyException`` /
``HTTPException``, a SCIM-deactivated user, and, unlike the key path, a missing user. ``get_user_object``
catches every DB failure and re-raises a bare ``ValueError`` (a deleted user and a real outage look
identical, the original error surviving only as ``__context__``), so the outage check walks the cause
chain, and a missing user falls through to ``no_active_key`` rather than an opaque gateway fault."""
from litellm.proxy._types import (
ProxyException, # noqa: PLC0415 # inline import avoids a module-load circular import
)
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_user_object,
)
from litellm.proxy.db.exception_handler import ( # noqa: PLC0415 # inline import avoids a module-load circular import
PrismaDBExceptionHandler,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return "unresolvable"
try:
user_object = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except (ProxyException, HTTPException):
return "no_active_key"
except Exception as exc: # noqa: BLE001 # a DB outage is retryable; a missing user (get_user_object's wrapped ValueError) or any other resolution failure fails closed as no_active_key, never a 500
if PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(exc):
return "unavailable"
verbose_logger.debug("_reload_active_user_by_id: user-resolution error (%s)", type(exc).__name__)
return "no_active_key"
if user_object is None:
return "no_active_key"
if isinstance(user_object.metadata, dict) and user_object.metadata.get("scim_active") is False:
return "no_active_key"
return None
async def _key_owner_scim_deactivated(key: "UserAPIKeyAuth") -> bool:
"""True only when the key's owning user was explicitly SCIM-deactivated, so a refresh revokes an
offboarded owner's key exactly as admission does via ``_reject_if_admitted_owner_scim_deactivated``.
A key with no owner, a missing owner record, or a failed lookup fails OPEN (returns ``False``),
matching admission and the standard builder: a key may outlive its owner record, and a transient DB
blip must not revoke a live key. Only an explicit ``scim_active`` of ``False`` gates renewal."""
if key.user_id is None:
return False
from litellm.proxy.auth.auth_checks import ( # noqa: PLC0415 # inline import avoids a module-load circular import
get_user_object,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
prisma_client,
user_api_key_cache,
)
if prisma_client is None:
return False
try:
owner = await get_user_object(
user_id=key.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except Exception as exc: # noqa: BLE001 # fail open: a missing owner (get_user_object's wrapped ValueError) or a DB blip must not revoke a live key
verbose_logger.debug("refresh: key-owner SCIM lookup failed, not revoking (%s)", type(exc).__name__)
return False
return owner is not None and isinstance(owner.metadata, dict) and owner.metadata.get("scim_active") is False
async def _revalidate_active_subject(identity: "EnvelopeIdentity") -> "_KeyResolutionFailure | None":
"""Re-validate that the subject sealed in a refresh envelope is still live, dispatching on its type:
a key_hash reloads the virtual key, a user_id reloads the user. Returns ``None`` when the subject is
active or a precise failure otherwise, so revocation gates renewal for either identity source the same
way admission gates the egress: a blocked or expired key, a SCIM-deactivated key owner (mirroring
admission's owner check, so an offboarded user cannot keep renewing a still-active key), and a
deactivated or deleted user all fail closed to ``no_active_key``."""
match identity.subject_type:
case "key_hash":
reloaded = await _reload_active_key_by_hash(identity.subject)
if not isinstance(reloaded, _ResolvedKey):
return reloaded
if await _key_owner_scim_deactivated(reloaded.key):
return "no_active_key"
return None
case "user_id":
return await _reload_active_user_by_id(identity.subject)
case _:
assert_never(identity.subject_type)
async def _extract_user_id_from_request(request: Request) -> str | None:
"""The litellm ``user_id`` for the token request, so a per-user token is stored under the same
identity the egress later reads it by. Storage is best-effort, so every non-resolved outcome
(including a transient DB outage) collapses to ``None`` here and the caller simply skips the store;
the bridge mint, which must status those outcomes differently, consumes
:func:`_resolve_active_litellm_key` directly."""
resolved = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
return None
return _active_key_user_id(resolved.key)
async def _store_per_user_token_server_side(
server: MCPServer,
user_id: str,
@ -946,350 +678,6 @@ async def authorize_with_server(
return response
_UpstreamGrantRejection = Literal["no_access_token", "expired_lifetime"]
"""Why an upstream token response cannot back a bridge envelope:
- ``no_access_token``: the response carries no usable ``access_token``
- ``expired_lifetime``: the response reports a parseable, non-positive ``expires_in``, i.e. an upstream
token that is already dead, so sealing it would forward a bearer the edge cannot use
An absent or unparseable ``expires_in`` is NOT a rejection; the lifetime is merely unknown and the
envelope caps it, the by-design behaviour for an upstream that omits the field."""
def _classify_upstream_lifetime(raw_expires_in: object) -> "int | Literal['unspecified', 'expired']":
"""Classify an upstream ``expires_in`` into a positive number of seconds, ``"unspecified"`` (absent
or unparseable, so the envelope caps it), or ``"expired"`` (a non-positive value the upstream reports
as already elapsed). Telling "we do not know the lifetime" apart from "the upstream says it is
already dead" is what stops an explicitly-expired token from silently receiving the envelope's 1h
cap. The expired decision is made on the parsed numeric value, not on ``int(...)`` of it, so a
positive sub-second lifetime in ``(0, 1)`` is not truncated to ``0`` and misread as elapsed; the
envelope works in whole seconds, so such a lifetime clamps up to its 1s floor. ``bool`` is excluded
(an ``int`` subclass but never a real lifetime), and the conversions can raise on ``NaN`` /
``Infinity`` / oversized input, which reads as unparseable rather than surfacing as a 500."""
if raw_expires_in is None or isinstance(raw_expires_in, bool) or not isinstance(raw_expires_in, (int, float, str)):
return "unspecified"
try:
numeric = float(raw_expires_in)
seconds = int(numeric)
except (ValueError, TypeError, OverflowError):
return "unspecified"
if numeric <= 0:
return "expired"
return max(1, seconds)
def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenGrant | _UpstreamGrantRejection":
"""Validate an upstream OAuth token response into a typed grant, or say why it cannot back an
envelope. Each field is isinstance-checked so nothing untyped from ``response.json()`` reaches the
grant. ``expires_in`` is read three ways (see :func:`_classify_upstream_lifetime`): an unknown
lifetime leaves the grant ``expires_in`` ``None`` for the envelope to cap, a positive value is
honoured, and an explicit already-elapsed value is a rejection rather than a silent fall-through to
the cap."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
UpstreamTokenGrant,
)
if not isinstance(token_response, dict):
return "no_access_token"
access = token_response.get("access_token")
if not isinstance(access, str) or not access:
return "no_access_token"
lifetime = _classify_upstream_lifetime(token_response.get("expires_in"))
if lifetime == "expired":
return "expired_lifetime"
token_type = token_response.get("token_type")
scope = token_response.get("scope")
return UpstreamTokenGrant(
access_token=SecretStr(access),
token_type=token_type if isinstance(token_type, str) and token_type else "Bearer",
# The upstream refresh_token is deliberately NOT sealed: the edge never consumes it (it forwards
# only token_type + access_token), so it would be dead weight embedding a long-lived upstream
# credential in the client-held bearer, and it enlarges the envelope. Refresh support is a
# follow-up (a dedicated refresh-envelope); the client re-runs authorization_code at the cap.
refresh_token=None,
scope=scope if isinstance(scope, str) and scope else None,
expires_in=lifetime if isinstance(lifetime, int) else None,
)
# ---------------------------------------------------------------------------
# DCR-bridge oauth_delegate mint: a three-phase pipeline whose failures are values.
#
# prepare (before the upstream exchange) -> validate every precondition and resolve identity+keys
# exchange (the single-use upstream code is consumed here, in exchange_token_with_server)
# finish (after the exchange) -> seal the upstream grant into the client-held envelope
#
# Every precondition lives in ``prepare``, which runs BEFORE the exchange, so no failure can burn the
# single-use code or rotate a refresh token, for either grant type -- that whole class of bug is gone
# by construction rather than guarded case by case. Failures are values mapped to an OAuth-shaped
# response in one place (``_bridge_mint_error_response``), so status codes and the RFC 6749 §5.2 body
# shape are uniform. Adding a failure mode is a new literal plus a match arm the type checker forces.
# ---------------------------------------------------------------------------
_BridgeMintError = Literal[
"no_identity",
"invalid_refresh",
"identity_unavailable",
"identity_unresolvable",
"not_configured",
"no_upstream_token",
"upstream_token_expired",
"too_large",
]
@dataclass(frozen=True, slots=True)
class _BridgeMintReady:
"""Everything the seal needs, resolved once before the exchange: the identity to bind the envelope
to and the master-key-derived envelope keys. The identity is a key_hash subject for the scripted
two-header client (resolved from the litellm key it presents) or a user_id subject for the
interactive SSO client (the user recovered from the gateway authorization code), so one phase-3 seal
serves both. Resolving identity here means ``_finish_bridge_mint`` has no preconditions left to
fail."""
identity: "EnvelopeIdentity"
keys: "EnvelopeKeys"
def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
"""Map a bridge-mint failure value to its token-endpoint response: one place, RFC 6749 §5.2 shape
(top-level ``error``, no-store headers) for every case, with a status truthful about where the
failure is. The caller's request is 400, a transient gateway outage is 503, a gateway
misconfiguration is 500, and an upstream problem is 502. The identity-resolution statuses match how
admission statuses the same conditions on the egress side, so mint and admit never disagree under
one outage."""
match error:
case "no_identity":
status, code, desc = (
400,
"invalid_request",
"this server issues a gateway-bound credential; complete the interactive sign-in, or "
"send a litellm credential (x-litellm-api-key or Authorization) on the token request",
)
case "invalid_refresh":
status, code, desc = (
400,
"invalid_grant",
"the refresh credential is not a valid, live refresh envelope for this server; "
"re-run authorization_code to obtain a new one",
)
case "identity_unavailable":
status, code, desc = (
503,
"temporarily_unavailable",
"the authentication database is temporarily unreachable; retry shortly",
)
case "identity_unresolvable":
status, code, desc = (
500,
"server_error",
"the gateway could not resolve the litellm identity for this request",
)
case "not_configured":
status, code, desc = (
500,
"server_error",
"the gateway is not configured to mint a gateway-bound credential (master_key is not set)",
)
case "no_upstream_token":
status, code, desc = (
502,
"server_error",
"the upstream token response has no usable access_token",
)
case "upstream_token_expired":
status, code, desc = (
502,
"server_error",
"the upstream token response reports an already-expired lifetime",
)
case "too_large":
status, code, desc = (
502,
"server_error",
"the upstream token is too large to seal into a gateway-bound credential",
)
case _:
assert_never(error)
return JSONResponse(
status_code=status, content={"error": code, "error_description": desc}, headers=TOKEN_NO_CACHE_HEADERS
)
def _key_resolution_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
"""Lift an identity-resolution failure into the mint taxonomy, preserving origin so the status stays
truthful: the caller's missing credential is 400, a transient DB outage is 503, and a gateway that
cannot resolve identity is 500."""
match failure:
case "no_active_key":
return "no_identity"
case "unavailable":
return "identity_unavailable"
case "unresolvable":
return "identity_unresolvable"
case _:
assert_never(failure)
def _upstream_rejection_to_mint_error(rejection: _UpstreamGrantRejection) -> _BridgeMintError:
"""Lift an upstream-response rejection into the mint taxonomy; both are upstream faults (502)."""
match rejection:
case "no_access_token":
return "no_upstream_token"
case "expired_lifetime":
return "upstream_token_expired"
case _:
assert_never(rejection)
async def _prepare_bridge_mint(
request: Request,
mcp_server: MCPServer,
bridge_identity: _BridgeAuthorizationCode | None = None,
) -> "_BridgeMintReady | _BridgeMintError":
"""Phase 1 for the authorization_code grant, BEFORE the upstream exchange: confirm the gateway can
mint (master_key set), resolve the litellm identity, and derive the envelope keys. Returns a ready
context or a precise failure value. Running before the exchange is what makes every failure here fail
closed without consuming the single-use code.
Two identity sources, one envelope. The interactive DCR client authenticates via SSO at the bridged
authorize, so its identity arrives as ``bridge_identity`` (the user recovered from the gateway
authorization code) and mints a user subject. The scripted two-header client presents a litellm key
on the token request instead, so its identity is the active key's hash and mints a key_hash subject.
A missing or invalid presented key keeps its resolution origin so the mapper statuses it truthfully;
neither source present is ``no_identity``. The refresh_token grant has its own phase-1
(:func:`_prepare_bridge_refresh`), which recovers identity from the presented refresh envelope."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
envelope_keys_from_master_key,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
key_hash_identity,
user_identity,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
master_key,
)
if not master_key:
return "not_configured"
keys = envelope_keys_from_master_key(master_key)
if bridge_identity is not None:
identity = user_identity(server_id=mcp_server.server_id, user_id=bridge_identity.litellm_user_id)
return _BridgeMintReady(identity=identity, keys=keys)
resolved = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
return _key_resolution_failure_to_mint_error(resolved)
identity = key_hash_identity(server_id=mcp_server.server_id, key_hash=resolved.key_hash)
return _BridgeMintReady(identity=identity, keys=keys)
@dataclass(frozen=True, slots=True)
class _BridgeRefreshReady:
"""A validated refresh request: the identity+keys to mint the renewed pair under, the upstream refresh
token (unwrapped from the client's refresh envelope) to exchange with the upstream IdP, and the scope
sealed alongside it at mint. The upstream refresh token is a ``SecretStr`` like every other credential
in this layer, so a repr or a traceback that captures this value never exposes the raw upstream refresh
token in plaintext. ``upstream_scope`` carries the originally-granted scope so the renewal re-requests
it when the client (a DCR/MCP client that typically omits scope on refresh) sends none, keeping the
renewed token's scope stable against an upstream that would otherwise narrow or drop it."""
ready: "_BridgeMintReady"
upstream_refresh_token: SecretStr
upstream_scope: str | None = None
def _refresh_key_failure_to_mint_error(failure: _KeyResolutionFailure) -> _BridgeMintError:
"""Lift an identity-resolution failure on the refresh path into the mint taxonomy. Unlike the mint
path, a resolved-but-inactive (or unknown) key is ``invalid_grant`` rather than ``invalid_request``:
the client did present an identity (sealed in the refresh envelope), but it is no longer live, so the
refresh is invalid and the client must re-authenticate. A transient outage is still 503 and a gateway
fault still 500, matching the mint path and admission."""
match failure:
case "no_active_key":
return "invalid_refresh"
case "unavailable":
return "identity_unavailable"
case "unresolvable":
return "identity_unresolvable"
case _:
assert_never(failure)
async def _prepare_bridge_refresh(
mcp_server: MCPServer, refresh_value: str | None
) -> "_BridgeRefreshReady | _BridgeMintError":
"""Phase 1 for the refresh_token grant, BEFORE the upstream exchange: open the client's refresh
envelope, re-validate the sealed litellm identity so a revoked key cannot keep refreshing, and
recover the upstream refresh token to exchange. Identity comes entirely from the sealed envelope, not
the HTTP request, so the request object is not needed here. The client presents a refresh envelope,
never a raw upstream refresh token, so a missing value, a non-envelope, an unopenable envelope, or one
minted for another server is ``invalid_grant``. Running before the exchange means a rejected refresh
never consumes or rotates the upstream refresh token."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
BridgeRefreshOpened,
envelope_keys_from_master_key,
open_bridge_refresh_envelope,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
master_key,
)
if not master_key:
return "not_configured"
if not refresh_value:
return "invalid_refresh"
keys = envelope_keys_from_master_key(master_key)
opened = open_bridge_refresh_envelope(refresh_value, keys, datetime.now(timezone.utc), mcp_server.server_id)
if not isinstance(opened, BridgeRefreshOpened):
return "invalid_refresh"
failure = await _revalidate_active_subject(opened.identity)
if failure is not None:
return _refresh_key_failure_to_mint_error(failure)
return _BridgeRefreshReady(
ready=_BridgeMintReady(identity=opened.identity, keys=keys),
upstream_refresh_token=opened.refresh.refresh_token,
upstream_scope=opened.refresh.scope,
)
def _finish_bridge_mint(
ready: "_BridgeMintReady", mcp_server: MCPServer, token_response: object, now: datetime
) -> "JSONResponse | _BridgeMintError":
"""Phase 3, AFTER the upstream exchange: seal the upstream grant into the client-held access envelope
using the pre-resolved identity and keys, and, when the upstream returned a refresh token, seal a
long-lived refresh envelope alongside it so the client can renew without re-authenticating. Shared by
the authorization_code and refresh_token paths, so a renewal that the upstream rotates re-issues a
fresh refresh envelope. The only hard failures here are properties of the upstream access token (no
usable token, an already-expired lifetime, or a token too large to seal); a refresh token that cannot
be sealed degrades to an access-only response rather than failing the whole exchange."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
build_bridge_token_response,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
SealedEnvelope,
UpstreamTokenGrant,
)
grant = _bridge_grant_from_token_response(token_response)
if not isinstance(grant, UpstreamTokenGrant):
return _upstream_rejection_to_mint_error(grant)
sealed = build_bridge_token_response(ready.identity, grant, ready.keys, now)
if not isinstance(sealed, SealedEnvelope):
return "too_large"
# Report expires_in from the JWT's own second-truncated exp, rounding the elapsed portion up, so the
# client is never told the bearer lives past the point admission (which uses that exp) rejects it.
expires_in = max(0, int(sealed.expires_at.timestamp()) - math.ceil(now.timestamp()))
refresh_envelope = _mint_refresh_envelope_value(ready.identity, token_response, ready.keys, now, mcp_server)
body = {
"access_token": sealed.token.get_secret_value(),
"token_type": "Bearer",
"expires_in": expires_in,
# A refresh envelope rides along only when the upstream returned a refresh token to seal; when it
# rotates on renewal, the client receives the new one and the old envelope's upstream token dies.
**({"refresh_token": refresh_envelope} if refresh_envelope is not None else {}),
}
return JSONResponse(body, headers=TOKEN_NO_CACHE_HEADERS)
def _token_credential_source(mcp_server: MCPServer) -> CredentialSource:
"""Mirrors the resolved-client rule in :func:`exchange_token_with_server`: when the server has a
stored client_id the gateway presents its own credentials upstream, so a credential rejection is
@ -1297,63 +685,6 @@ def _token_credential_source(mcp_server: MCPServer) -> CredentialSource:
return "gateway_stored" if mcp_server.client_id else "caller_supplied"
def _upstream_refresh_credential(token_response: object) -> "RefreshCredential | None":
"""Extract the upstream refresh grant from a token response, or ``None`` when there is none to seal.
Each field is isinstance-checked so nothing untyped reaches the refresh envelope; ``refresh_expires_in``
(the refresh token's own lifetime, when the upstream reports it) is classified like ``expires_in`` and
bounds the refresh envelope's TTL. An upstream that reports the refresh token itself as already elapsed
(``refresh_expires_in`` non-positive) yields ``None`` rather than a refresh envelope: sealing a dead
token would hand the client a full-TTL-capped envelope the IdP will reject, so the exchange degrades to
an access-only response (the client re-authenticates at access expiry), mirroring how
:func:`_bridge_grant_from_token_response` refuses an already-elapsed access token instead of capping it."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
RefreshCredential,
)
if not isinstance(token_response, dict):
return None
refresh = token_response.get("refresh_token")
if not isinstance(refresh, str) or not refresh:
return None
lifetime = _classify_upstream_lifetime(token_response.get("refresh_expires_in"))
if lifetime == "expired":
return None
scope = token_response.get("scope")
return RefreshCredential(
refresh_token=SecretStr(refresh),
scope=scope if isinstance(scope, str) and scope else None,
expires_in=lifetime if isinstance(lifetime, int) else None,
)
def _mint_refresh_envelope_value(
identity: "EnvelopeIdentity", token_response: object, keys: "EnvelopeKeys", now: datetime, mcp_server: MCPServer
) -> str | None:
"""Seal the upstream refresh grant (if any) into a refresh envelope and return its bearer string, or
``None`` when the upstream returned no refresh token or the refresh token is too large to seal. A
too-large refresh token degrades to an access-only response (logged) rather than failing an exchange
that already succeeded upstream: the client simply re-authenticates when the access envelope expires."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
build_bridge_refresh_token_response,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # inline import avoids a module-load circular import
SealedEnvelope,
)
refresh_credential = _upstream_refresh_credential(token_response)
if refresh_credential is None:
return None
sealed = build_bridge_refresh_token_response(identity, refresh_credential, keys, now)
if isinstance(sealed, SealedEnvelope):
return sealed.token.get_secret_value()
verbose_logger.warning(
"bridge mint: the upstream refresh token is too large to seal into a refresh envelope for "
"server=%s; issuing an access-only response, so the client re-authenticates at access expiry",
mcp_server.server_id,
)
return None
async def exchange_token_with_server(
request: Request,
mcp_server: MCPServer,

View file

@ -2422,6 +2422,16 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"is active as a reminder that hard enforcement is relaxed."
),
)
skip_user_budget_on_team_key: bool | None = Field(
None,
description=(
"If True, restores the legacy behavior where a user's personal "
"max_budget is NOT enforced when their key belongs to a team; only "
"the team (and team-member) budgets apply. Defaults to False, meaning "
"the user's personal max_budget is always enforced regardless of "
"whether the key belongs to a team (see GitHub issue #12905)."
),
)
user_url_validation: Optional[bool] = Field(
None,
description=(

View file

@ -626,26 +626,29 @@ async def common_checks(
)
async def _user_max_budget_check() -> None:
# 4.1 personal budget, if personal key
if (
(team_object is None or team_object.team_id is None)
and user_object is not None
and user_object.max_budget is not None
):
from litellm.proxy.proxy_server import get_current_spend
if user_object is None or user_object.max_budget is None:
return
skip_for_team = (
general_settings.get("skip_user_budget_on_team_key") is True
and team_object is not None
and team_object.team_id is not None
)
if skip_for_team:
return
from litellm.proxy.proxy_server import get_current_spend
user_budget = user_object.max_budget
user_spend = await get_current_spend(
counter_key=f"spend:user:{user_object.user_id}",
fallback_spend=user_object.spend or 0.0,
user_budget = user_object.max_budget
user_spend = await get_current_spend(
counter_key=f"spend:user:{user_object.user_id}",
fallback_spend=user_object.spend or 0.0,
max_budget=user_budget,
)
if math.isfinite(user_budget) and user_spend >= user_budget:
raise litellm.BudgetExceededError(
current_cost=user_spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
)
if math.isfinite(user_budget) and user_spend >= user_budget:
raise litellm.BudgetExceededError(
current_cost=user_spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
)
# Each scope reads a distinct counter key with no cross-scope ordering
# dependency, so the per-scope Redis-first reads run concurrently instead

View file

@ -2442,6 +2442,7 @@ async def _reserve_budget_after_common_checks(
proxy_logging_obj=proxy_logging_obj,
end_user_id=end_user_id,
end_user_object=end_user_object,
skip_user_budget_on_team_key=general_settings.get("skip_user_budget_on_team_key") is True,
)

View file

@ -4034,30 +4034,41 @@ class MicrosoftSSOHandler:
base_url = MicrosoftSSOHandler.get_graph_api_base_url()
# Endpoint to get app role assignments for the given service principal
endpoint = f"/servicePrincipals/{service_principal_id}/appRoleAssignedTo"
url = base_url + endpoint
next_link: str | None = base_url + endpoint
headers = {
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
}
response = await async_client.get(url, headers=headers)
response_json = response.json()
verbose_proxy_logger.debug(f"Response from service principal app role assigned to: {response_json}")
group_ids: List[str] = []
service_principal_teams: List[MicrosoftServicePrincipalTeam] = []
page_count = 0
for _object in response_json.get("value", []):
if _object.get("principalType") == "Group":
# Append the group ID to the list
group_ids.append(_object.get("principalId"))
# Append the service principal team to the list
service_principal_teams.append(
MicrosoftServicePrincipalTeam(
principalDisplayName=_object.get("principalDisplayName"),
principalId=_object.get("principalId"),
while next_link is not None and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
response = await async_client.get(next_link, headers=headers)
response_json = response.json()
verbose_proxy_logger.debug(f"Response from service principal app role assigned to: {response_json}")
for _object in response_json.get("value", []):
if _object.get("principalType") == "Group":
# Append the group ID to the list
group_ids.append(_object.get("principalId"))
# Append the service principal team to the list
service_principal_teams.append(
MicrosoftServicePrincipalTeam(
principalDisplayName=_object.get("principalDisplayName"),
principalId=_object.get("principalId"),
)
)
)
next_link = response_json.get("@odata.nextLink")
page_count += 1
if next_link is not None and page_count >= MicrosoftSSOHandler.MAX_GRAPH_API_PAGES:
verbose_proxy_logger.warning(
f"Reached maximum page limit of {MicrosoftSSOHandler.MAX_GRAPH_API_PAGES}. Some service principal group assignments may not be included."
)
return group_ids, service_principal_teams

View file

@ -14805,6 +14805,7 @@ async def get_config_list(
"forward_client_headers_to_llm_api": {"type": "Boolean"},
"mcp_required_fields": {"type": "List"},
"cancel_on_disconnect": {"type": "Boolean"},
"skip_user_budget_on_team_key": {"type": "Boolean"},
}
return_val = []

View file

@ -123,6 +123,7 @@ async def reserve_budget_for_request(
proxy_logging_obj: ProxyLogging,
end_user_id: Optional[str] = None,
end_user_object: Optional[Any] = None,
skip_user_budget_on_team_key: bool = False,
) -> Optional[dict]:
if valid_token is None or not RouteChecks.is_llm_api_route(route=route):
return None
@ -141,6 +142,7 @@ async def reserve_budget_for_request(
proxy_logging_obj=proxy_logging_obj,
end_user_id=end_user_id,
end_user_object=end_user_object,
skip_user_budget_on_team_key=skip_user_budget_on_team_key,
)
if not counters:
return None
@ -296,6 +298,7 @@ async def _get_budget_counters(
proxy_logging_obj: ProxyLogging,
end_user_id: Optional[str] = None,
end_user_object: Optional[Any] = None,
skip_user_budget_on_team_key: bool = False,
) -> List[_BudgetCounter]:
counters: List[_BudgetCounter] = []
@ -344,8 +347,9 @@ async def _get_budget_counters(
)
)
is_team_key = team_object is not None and team_object.team_id is not None
if (
(team_object is None or team_object.team_id is None)
not (is_team_key and skip_user_budget_on_team_key)
and user_object is not None
and user_object.user_id is not None
and user_object.max_budget is not None

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.93.0"
version = "1.94.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
@ -62,8 +62,8 @@ proxy = [
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.26.0,<2.0",
"litellm-proxy-extras==0.4.76",
"litellm-enterprise==0.1.49",
"litellm-proxy-extras==0.4.77",
"litellm-enterprise==0.1.50",
"RestrictedPython>=8.1,<9.0",
"rich>=13.9.4,<14.0",
"polars>=1.38.1,<2.0",
@ -264,6 +264,8 @@ constraint-dependencies = [
"aiohttp>=3.14.1,<4.0",
"packaging>=24.0",
"soupsieve>=2.8.4",
"httplib2>=0.32.0",
"setuptools>=83.0.0",
]
override-dependencies = [
# a2a-sdk 1.x requires packaging>=24.0; lunary 1.4.x still caps at <24.0.
@ -284,7 +286,7 @@ members = ["enterprise", "litellm-proxy-extras"]
profile = "black"
[tool.commitizen]
version = "1.93.0"
version = "1.94.0"
version_files = [
"pyproject.toml:^version",
]

View file

@ -17,6 +17,7 @@ from __future__ import annotations
import json
import time
from datetime import datetime, timedelta, timezone
from typing import Callable
import pytest
@ -49,7 +50,7 @@ from e2e_http import (
unwrap,
)
from lifecycle import ResourceManager
from models import KeyGenerateBody, SpendLogRow, SpendLogsParams
from models import KeyGenerateBody, SpendLogRow
pytestmark = pytest.mark.e2e
@ -349,6 +350,10 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
the file-read path fires while the batch itself is not blocked.
``resources.key()`` cannot set limits, so the key is minted on the gateway
directly and its delete deferred.
Snapshots read /spend/logs/v2 over a bounded window around the test instead
of the unpaginated /spend/logs whole-table read, which grows with the
environment and OOMed the e2e runner on stage.
"""
user_id = f"e2e-batch-rl-{unique_marker()}"
key = client.gateway.generate_key(
@ -356,8 +361,13 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
)
resources.defer(lambda: client.gateway.delete_key(key))
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
window_end = window_start + timedelta(hours=2)
before = frozenset(
row.request_id for row in unattributed_rows(client.gateway.spend_logs(SpendLogsParams()))
row.request_id
for row in unattributed_rows(
client.gateway.spend_logs_window(start=window_start, end=window_end)
)
)
file = unwrap(
@ -379,7 +389,9 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
new_orphans = [
row
for row in unattributed_rows(client.gateway.spend_logs(SpendLogsParams()))
for row in unattributed_rows(
client.gateway.spend_logs_window(start=window_start, end=window_end)
)
if row.request_id not in before
]
assert not new_orphans, (

View file

@ -8,7 +8,7 @@
- {id: logging.datadog.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/datadog/datadog.py", rationale: "Powers dashboards/alerts; cardinality regressions common"}
- {id: logging.datadog.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions], source: "integrations/datadog/datadog.py", rationale: "Failure metrics for alerting/SLO"}
- {id: logging.prometheus.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/prometheus.py", rationale: "Standard OSS metrics; per-key cardinality (existing e2e)"}
- {id: logging.otel.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/otel/logger.py", rationale: "OTEL spans on every call path"}
- {id: logging.otel.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, responses, embeddings], source: "integrations/otel/logger.py", rationale: "OTEL spans on every call path"}
- {id: logging.otel.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions, messages], source: "integrations/otel/logger.py", rationale: "Error spans for observability continuity"}
- {id: logging.braintrust.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/braintrust_logging.py", rationale: "Evals platform spend"}
- {id: logging.langsmith.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/langsmith.py", rationale: "LangChain ecosystem"}

View file

@ -25,6 +25,7 @@ UI_USERNAME = os.environ.get("E2E_UI_USERNAME", "admin")
UI_PASSWORD = os.environ.get("E2E_UI_PASSWORD", MASTER_KEY)
CHEAP_ANTHROPIC_MODEL = os.environ.get("E2E_CHEAP_ANTHROPIC_MODEL", "claude-haiku-4-5")
CHEAP_OPENAI_MODEL = os.environ.get("E2E_CHEAP_OPENAI_MODEL", "gpt-5.5")
# Jaeger query API of the compose stack's OTEL trace destination (the `jaeger`
# service in docker-compose.yml maps it to host 16686). Trace-completeness tests

View file

@ -51,7 +51,9 @@ from models import (
OcrBody,
OcrResponse,
SpendLogRow,
SpendLogs,
SpendLogsPage,
SpendLogsPageParams,
SpendLogsParams,
)
from e2e_config import (
@ -289,6 +291,28 @@ class Gateway:
case _:
return None
def spend_logs_window(self, *, start: datetime, end: datetime) -> list[SpendLogRow]:
def fetch(page: int) -> SpendLogsPage:
return unwrap(
self.transport.get(
"/spend/logs/v2",
headers=self.transport.master,
params=SpendLogsPageParams(
start_date=start.strftime("%Y-%m-%d %H:%M:%S"),
end_date=end.strftime("%Y-%m-%d %H:%M:%S"),
page=page,
page_size=100,
),
response_type=SpendLogsPage,
)
)
first = fetch(1)
return [
*first.data,
*(row for page in range(2, first.total_pages + 1) for row in fetch(page).data),
]
def poll_logs_for_key(
self, key: str, *, min_rows: int = 1, predicate: RowsPredicate | None = None
) -> list[SpendLogRow]:

View file

@ -98,6 +98,14 @@ CLAUDE_CODE_TOOLS = [
]
class ResponsesRequestBody(BaseModel):
"""OpenAI Responses API /v1/responses request (non-streaming)."""
model: str
input: str
max_output_tokens: int
class TeamCallbackBody(BaseModel):
callback_name: Literal["langfuse_otel", "langfuse", "langsmith", "gcs"]
callback_type: Literal["success", "failure", "success_and_failure"]
@ -509,6 +517,18 @@ class LoggingClient:
)
return self.gateway.transport.send(
"/v1/messages", headers=self.gateway.transport.bearer(key), json=body
def responses_raw(
self, key: str, model: str, text: str, *, max_output_tokens: int = 64
) -> StreamingResponse:
"""Non-streaming POST /v1/responses (OpenAI Responses API): raw outcome
judged by status/body/headers, for tests that need x-litellm-call-id.
max_output_tokens caps reasoning-model output cost; a capped response is
still a 200 and still exports the trace."""
return self.gateway.transport.send(
"/v1/responses",
headers=self.gateway.transport.bearer(key),
json=ResponsesRequestBody(model=model, input=text, max_output_tokens=max_output_tokens),
)
def scrape_metrics(self) -> str:

View file

@ -23,7 +23,7 @@ from collections.abc import Callable
import pytest
from pydantic import BaseModel, ConfigDict
from e2e_config import CHEAP_ANTHROPIC_MODEL, unique_marker
from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, unique_marker
from e2e_http import NoBody, StreamingResponse, require_successful_call
from lifecycle import ResourceManager
from logging_client import LoggingClient
@ -151,7 +151,7 @@ def _settled_names(*, route: str, genai_span: str) -> set[str]:
class TestOtelTraceCompleteness:
@pytest.mark.covers("logging.otel.success.exports_metric")
@pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["chat_completions"])
def test_chat_completions_exports_complete_trace(
self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager
) -> None:
@ -189,3 +189,71 @@ class TestOtelTraceCompleteness:
settled_prefixes={DB_SPAN_PREFIX},
)
_assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}")
@pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["messages"])
def test_messages_exports_complete_trace(
self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager
) -> None:
"""This test verifies that one successful non-streaming /v1/messages request
produces exactly one complete OTEL trace.
The trace must have a single root span named "POST /v1/messages". The
authentication, database, cost-writing, and model-call spans must all belong to
the same trace and have valid parent relationships leading back to that root.
The model-call span is expected to be named "chat <model>". The test fails if
the request is split across multiple traces, if any span references a missing
parent, or if the model-call span cannot be connected back to the root."""
route = "/v1/messages"
_assert_otel_destination_configured(client)
key = client.key_with_alias(f"otel-trace-messages-{unique_marker()}", models=[MODEL])
resources.defer(lambda: client.delete_key(key))
marker = unique_marker()
outcome = _first_ok(
client, lambda: client.messages_raw(key, MODEL, f"reply with one word {marker}", max_tokens=16)
)
assert outcome.call_id is not None, "success response must carry x-litellm-call-id"
hits = otel_reader.poll_traces_for_call(
call_id=outcome.call_id,
settled_names=_settled_names(route=route, genai_span=f"chat {MODEL}"),
settled_prefixes={DB_SPAN_PREFIX},
)
_assert_complete_trace(hits, route=route, genai_span=f"chat {MODEL}")
@pytest.mark.covers("logging.otel.success.exports_metric", exercised_on=["responses"])
def test_responses_exports_complete_trace(
self, client: LoggingClient, otel_reader: OtelReader, resources: ResourceManager
) -> None:
"""This test verifies that one successful non-streaming /v1/responses request
produces exactly one complete OTEL trace.
The trace must have a single root span named "POST /v1/responses". The
authentication, database, cost-writing, and model-call spans must all belong to
the same trace and have valid parent relationships leading back to that root.
The model-call span is expected to be named "chat <model>". The test fails if
the request is split across multiple traces, if any span references a missing
parent, or if the model-call span cannot be connected back to the root."""
route = "/v1/responses"
_assert_otel_destination_configured(client)
key = client.key_with_alias(f"otel-trace-responses-{unique_marker()}", models=[CHEAP_OPENAI_MODEL])
resources.defer(lambda: client.delete_key(key))
marker = unique_marker()
outcome = _first_ok(
client,
lambda: client.responses_raw(key, CHEAP_OPENAI_MODEL, f"reply with one word {marker}"),
)
assert outcome.call_id is not None, "success response must carry x-litellm-call-id"
genai_span = f"chat {CHEAP_OPENAI_MODEL}"
hits = otel_reader.poll_traces_for_call(
call_id=outcome.call_id,
settled_names=_settled_names(route=route, genai_span=genai_span),
settled_prefixes={DB_SPAN_PREFIX},
)
_assert_complete_trace(hits, route=route, genai_span=genai_span)

View file

@ -8,7 +8,7 @@ from __future__ import annotations
from typing import Literal
from pydantic import BaseModel, ConfigDict, RootModel
from pydantic import BaseModel, ConfigDict, RootModel, model_validator
# ---------- keys ----------
@ -279,6 +279,16 @@ class SpendLogsParams(BaseModel):
page: int = 1
page_size: int = 100
@model_validator(mode="after")
def require_filter(self) -> SpendLogsParams:
if self.request_id is None and self.api_key is None:
raise ValueError(
"unfiltered /spend/logs returns the entire spend table and OOMs the "
"runner on long-lived environments; filter by request_id or api_key, "
"or use Gateway.spend_logs_window for a bounded /spend/logs/v2 read"
)
return self
class SpendLogsPageParams(BaseModel):
"""Query for /spend/logs/v2, which requires an explicit date window and

View file

@ -1,17 +1,23 @@
"""Unit coverage for the Gateway model-management surface (create_model /
delete_model).
delete_model) and the bounded spend read-back (spend_logs_window).
The batches conftest and several llm_translation tests register deployments at
runtime through gateway.create_model; when that method went missing, every batch
test errored at fixture setup (AttributeError) before a single request reached
the proxy. This pins the surface with a typed fake Transport so a rename or
signature drift fails here instead of in a live stage run.
spend_logs_window exists because the unpaginated /spend/logs whole-table read
grew past the e2e runner's memory limit on stage and OOMKilled every run; these
tests pin its /spend/logs/v2 pagination and that SpendLogsParams can no longer
express the unfiltered read.
"""
from dataclasses import dataclass, field
from datetime import datetime, timezone
import pytest
from pydantic import BaseModel
from pydantic import BaseModel, ValidationError
from batches.batch_client import BatchClient
from e2e_gateway import Gateway
@ -30,6 +36,9 @@ from models import (
ModelNewBody,
ModelNewResponse,
ModelsListResponse,
SpendLogsPage,
SpendLogsPageParams,
SpendLogsParams,
)
@ -46,6 +55,8 @@ class _RecordingTransport:
servable_after_gets: int = 0
models_error: UnknownApiError | None = None
model_gets: int = 0
spend_total: int = 0
spend_gets: list[SpendLogsPageParams] = field(default_factory=list)
_created: list[str] = field(default_factory=list)
def post[R: BaseModel](
@ -91,6 +102,22 @@ class _RecordingTransport:
return Success(
data=response_type.model_validate({"data": [{"id": name} for name in visible]})
)
if path == "/spend/logs/v2" and response_type is SpendLogsPage:
assert isinstance(params, SpendLogsPageParams)
self.spend_gets.append(params)
offset = (params.page - 1) * params.page_size
count = min(params.page_size, max(self.spend_total - offset, 0))
return Success(
data=response_type.model_validate(
{
"data": [{"request_id": f"req-{offset + i}"} for i in range(count)],
"total": self.spend_total,
"page": params.page,
"page_size": params.page_size,
"total_pages": (self.spend_total + params.page_size - 1) // params.page_size,
}
)
)
raise AssertionError(f"unexpected get: {path}")
def delete[R: BaseModel](
@ -202,3 +229,45 @@ def test_gateway_delete_model_posts_the_model_id() -> None:
assert path == "/model/delete"
assert isinstance(body, ModelDeleteBody)
assert body.id == "registered-id"
WINDOW_START = datetime(2026, 7, 14, 12, 0, 0, tzinfo=timezone.utc)
WINDOW_END = datetime(2026, 7, 14, 14, 0, 0, tzinfo=timezone.utc)
def test_gateway_spend_logs_window_pages_through_every_row_in_the_window() -> None:
transport = _RecordingTransport(spend_total=250)
gateway = Gateway(transport=transport)
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
assert len(rows) == 250
assert len({row.request_id for row in rows}) == 250
assert [params.page for params in transport.spend_gets] == [1, 2, 3]
assert all(params.start_date == "2026-07-14 12:00:00" for params in transport.spend_gets)
assert all(params.end_date == "2026-07-14 14:00:00" for params in transport.spend_gets)
def test_gateway_spend_logs_window_stops_at_an_exact_page_boundary() -> None:
transport = _RecordingTransport(spend_total=200)
gateway = Gateway(transport=transport)
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
assert len(rows) == 200
assert [params.page for params in transport.spend_gets] == [1, 2]
def test_gateway_spend_logs_window_returns_empty_for_an_empty_window() -> None:
transport = _RecordingTransport(spend_total=0)
gateway = Gateway(transport=transport)
rows = gateway.spend_logs_window(start=WINDOW_START, end=WINDOW_END)
assert rows == []
assert [params.page for params in transport.spend_gets] == [1]
def test_spend_logs_params_rejects_the_unfiltered_whole_table_read() -> None:
with pytest.raises(ValidationError, match="spend_logs_window"):
SpendLogsParams()

View file

@ -219,8 +219,8 @@ async def test_aaauser_personal_budgets(key_ownership):
"""
Set a personal budget on a user
- have it only apply when key belongs to user -> raises BudgetExceededError
- if key belongs to team, have key respect team budget -> allows call to go through
User budget is enforced regardless of key ownership (personal or team).
Both cases should raise BudgetExceededError when the user is over budget.
"""
import asyncio
import time
@ -229,7 +229,12 @@ async def test_aaauser_personal_budgets(key_ownership):
from starlette.datastructures import URL
import litellm
from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth
from litellm.proxy._types import (
LiteLLM_UserTable,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.proxy_server import hash_token, user_api_key_cache
@ -273,14 +278,9 @@ async def test_aaauser_personal_budgets(key_ownership):
== valid_token
)
try:
with pytest.raises(ProxyException) as exc_info:
await user_api_key_auth(request=request, api_key="Bearer " + user_key)
if key_ownership == "user_key":
pytest.fail("Expected this call to fail. User is over limit.")
except Exception:
if key_ownership == "team_key":
pytest.fail("Expected this call to work. Key is below team budget.")
assert exc_info.value.type == ProxyErrorTypes.budget_exceeded
@pytest.mark.asyncio

View file

@ -1614,3 +1614,105 @@ class TestGuardrailInterventionClassification:
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
assert slg["guardrail_status"] == "guardrail_intervened"
class _ApplyStyleGuardrail(CustomGuardrail):
"""Overrides only apply_guardrail, like openai_moderation; async_pre_call_hook stays the CustomLogger no-op."""
def __init__(self, block: bool):
from litellm.types.guardrails import GuardrailEventHooks
super().__init__(
guardrail_name="apply-style-guardrail",
event_hook=GuardrailEventHooks.pre_call,
default_on=False,
)
self.block = block
self.apply_called = False
self.seen_texts = None
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
from fastapi import HTTPException
self.apply_called = True
self.seen_texts = inputs.get("texts")
if self.block:
raise HTTPException(status_code=400, detail={"error": "Violated moderation policy"})
return inputs
class TestApplyGuardrailStyleDeploymentDispatch:
"""LIT-4217 regression: model-level guardrails that implement only the
unified apply_guardrail interface must execute in
async_pre_call_deployment_hook instead of silently hitting the
async_pre_call_hook no-op."""
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", [CallTypes.completion, CallTypes.acompletion])
async def test_blocks_when_requested_via_model_level_guardrails(self, call_type):
from fastapi import HTTPException
guardrail = _ApplyStyleGuardrail(block=True)
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "flagged content"}],
"guardrails": ["apply-style-guardrail"],
"metadata": {},
}
with pytest.raises(HTTPException):
await guardrail.async_pre_call_deployment_hook(kwargs, call_type)
assert guardrail.apply_called is True
assert guardrail.seen_texts == ["flagged content"]
@pytest.mark.asyncio
async def test_pass_path_runs_guardrail_and_strips_dispatch_key(self):
guardrail = _ApplyStyleGuardrail(block=False)
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "hello"}],
"guardrails": ["apply-style-guardrail"],
"metadata": {},
}
result = await guardrail.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert guardrail.apply_called is True
assert result is not None
assert "guardrail_to_apply" not in result
assert result["messages"] == [{"role": "user", "content": "hello"}]
@pytest.mark.asyncio
async def test_skips_when_not_requested(self):
guardrail = _ApplyStyleGuardrail(block=True)
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "hello"}],
"guardrails": ["some-other-guardrail"],
"metadata": {},
}
result = await guardrail.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert guardrail.apply_called is False
assert result is not None
@pytest.mark.asyncio
async def test_fails_closed_when_proxy_extras_missing(self):
import sys
from unittest.mock import patch
guardrail = _ApplyStyleGuardrail(block=True)
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "flagged content"}],
"guardrails": ["apply-style-guardrail"],
"metadata": {},
}
with patch.dict(sys.modules, {"litellm.proxy.utils": None}):
with pytest.raises(ImportError, match="litellm\\[proxy\\]"):
await guardrail.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert guardrail.apply_called is False

View file

@ -2,6 +2,7 @@ import pytest
from litellm.constants import (
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
)
@ -174,6 +175,81 @@ def test_unrecognized_effort_raises_clean_400():
assert exc_info.value.status_code == 400
def test_pinned_temperature_dropped_when_adaptive_downgraded_to_enabled():
"""Regression (#33203): Claude Code's safety classifier sends adaptive thinking +
temperature=0 to Haiku 4.5. The adaptive interface is downgraded to legacy enabled
thinking, but Anthropic rejects "temperature may only be set to 1 when thinking is
enabled". The pinned temperature must be dropped so the request succeeds while the
downgraded thinking is preserved."""
params = _claude_code_payload(effort="medium")
params["temperature"] = 0
result = _transform("claude-haiku-4-5", params)
assert result["thinking"] == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
}
assert "temperature" not in result
def test_temperature_one_preserved_with_enabled_thinking():
"""temperature=1 is compatible with extended thinking, so it must be kept."""
params = _claude_code_payload(effort="medium")
params["temperature"] = 1
result = _transform("claude-haiku-4-5", params)
assert result["thinking"]["type"] == "enabled"
assert result["temperature"] == 1
def test_pinned_temperature_preserved_when_thinking_dropped():
"""When thinking is dropped entirely (non-reasoning model), there is no thinking
conflict, so a pinned temperature must survive untouched."""
params = _claude_code_payload(effort="medium")
params["temperature"] = 0
result = _transform("claude-3-5-haiku-latest", params)
assert "thinking" not in result
assert result["temperature"] == 0
def test_pinned_temperature_preserved_for_adaptive_model():
"""Adaptive models (4.6+) own the thinking/temperature relationship natively, so
the passthrough must not strip a pinned temperature for them."""
params = _claude_code_payload(effort="high")
params["temperature"] = 0
result = _transform("claude-sonnet-4-6", params)
assert result["thinking"] == {"type": "adaptive"}
assert result["temperature"] == 0
def test_pinned_temperature_dropped_for_opus_4_5_effort():
"""Opus 4.5 keeps native output_config.effort (extended thinking), which is equally
incompatible with a pinned non-1 temperature, so the temperature must be dropped."""
params = _claude_code_payload(effort="medium")
params["temperature"] = 0
result = _transform("claude-opus-4-5", params)
assert result["output_config"] == {"effort": "medium"}
assert "temperature" not in result
def test_reasoning_effort_with_pinned_temperature_drops_temperature():
"""The reasoning_effort alias synthesizes legacy enabled thinking on a non-adaptive
model; a co-pinned non-1 temperature must be dropped to avoid the Anthropic 400."""
result = _transform(
"claude-haiku-4-5",
{"max_tokens": 8192, "reasoning_effort": "low", "temperature": 0},
)
assert result["thinking"] == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
}
assert "temperature" not in result
def test_non_adaptive_request_without_effort_is_untouched():
"""A non-adaptive model receiving a request with no adaptive interface (no
effort, no adaptive thinking) must pass through untouched."""

View file

@ -4374,10 +4374,8 @@ _BRIDGE_MASTER_KEY = "sk-bridge-producer-master-key-0123456789abcdef"
async def _exchange_for_bridge_server(server, upstream_body, key_hash, code="auth-code", fake_client_out=None):
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_ResolvedKey,
exchange_token_with_server,
)
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _ResolvedKey
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server
fake_http_response = MagicMock()
fake_http_response.json.return_value = upstream_body
@ -4397,7 +4395,7 @@ async def _exchange_for_bridge_server(server, upstream_body, key_hash, code="aut
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_active_litellm_key",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._resolve_active_litellm_key",
new=key_resolver,
),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
@ -4808,7 +4806,7 @@ async def _refresh_for_bridge_server(
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._revalidate_active_subject",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._revalidate_active_subject",
new=AsyncMock(return_value=revalidate_result),
),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
@ -5085,7 +5083,7 @@ async def test_bridge_refresh_grant_with_deactivated_user_is_invalid_grant_befor
async def test_revalidate_active_subject_dispatches_on_subject_type():
"""Subject re-validation routes a key_hash envelope to the key reload and a user_id envelope to the
user reload, so revocation gates renewal for either identity source through one dispatch point."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_ResolvedKey,
_revalidate_active_subject,
)
@ -5093,11 +5091,11 @@ async def test_revalidate_active_subject_dispatches_on_subject_type():
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_key_by_hash",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_key_by_hash",
new=AsyncMock(return_value=_ResolvedKey(key_hash="kh", key=MagicMock())),
) as key_reload,
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_user_by_id",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_user_by_id",
new=AsyncMock(return_value=None),
) as user_reload,
):
@ -5107,11 +5105,11 @@ async def test_revalidate_active_subject_dispatches_on_subject_type():
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_key_by_hash",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_key_by_hash",
new=AsyncMock(),
) as key_reload2,
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_user_by_id",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_user_by_id",
new=AsyncMock(return_value="no_active_key"),
) as user_reload2,
):
@ -5125,7 +5123,7 @@ def test_upstream_refresh_credential_expired_refresh_token_is_not_sealed():
not be sealed: _upstream_refresh_credential returns None so the exchange degrades to an access-only
response, mirroring how the access grant refuses an already-elapsed access token rather than capping a
dead token to the full refresh TTL. A live or unspecified lifetime still yields a credential."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _upstream_refresh_credential
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _upstream_refresh_credential
assert _upstream_refresh_credential({"access_token": "A", "refresh_token": "R", "refresh_expires_in": 0}) is None
assert _upstream_refresh_credential({"refresh_token": "R", "refresh_expires_in": -5}) is None
@ -5164,7 +5162,7 @@ async def test_bridge_refresh_upstream_invalid_grant_maps_to_invalid_grant():
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._revalidate_active_subject",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._revalidate_active_subject",
new=AsyncMock(return_value=None),
),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
@ -5217,7 +5215,7 @@ async def test_bridge_refresh_upstream_error_detection_parses_json_not_substring
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._revalidate_active_subject",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._revalidate_active_subject",
new=AsyncMock(return_value=None),
),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
@ -5244,7 +5242,7 @@ async def test_revalidate_key_subject_revoked_when_owner_scim_deactivated(proxy_
"""A key_hash refresh envelope whose key is still active but whose OWNING user was SCIM-deactivated must
fail closed to no_active_key, mirroring how admission's _reject_if_admitted_owner_scim_deactivated
revokes an offboarded owner's key. Without this, an offboarded user keeps renewing a live key."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _ResolvedKey, _revalidate_active_subject
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _ResolvedKey, _revalidate_active_subject
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import key_hash_identity
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -5254,7 +5252,7 @@ async def test_revalidate_key_subject_revoked_when_owner_scim_deactivated(proxy_
resolved = _ResolvedKey(key_hash="kh", key=MagicMock(user_id="offboarded-owner"))
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_key_by_hash",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_key_by_hash",
new=AsyncMock(return_value=resolved),
),
patch(
@ -5272,7 +5270,7 @@ async def test_revalidate_key_subject_active_owner_renews_and_missing_owner_fail
"""The key-owner SCIM gate blocks only an explicit scim_active False: an active owner renews (None), and
a missing owner (get_user_object's wrapped ValueError) fails OPEN, since a key may outlive its owner
record and a transient blip must not revoke a live key."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _ResolvedKey, _revalidate_active_subject
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _ResolvedKey, _revalidate_active_subject
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import key_hash_identity
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -5283,7 +5281,7 @@ async def test_revalidate_key_subject_active_owner_renews_and_missing_owner_fail
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_key_by_hash",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_key_by_hash",
new=AsyncMock(return_value=resolved),
),
patch(
@ -5295,7 +5293,7 @@ async def test_revalidate_key_subject_active_owner_renews_and_missing_owner_fail
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._reload_active_key_by_hash",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._reload_active_key_by_hash",
new=AsyncMock(return_value=resolved),
),
patch(
@ -5324,7 +5322,7 @@ async def test_bridge_mint_fails_closed_before_upstream_when_master_key_unset():
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_active_litellm_key",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._resolve_active_litellm_key",
new=AsyncMock(return_value="no_active_key"),
),
patch("litellm.proxy.proxy_server.master_key", None),
@ -5361,7 +5359,7 @@ async def _prepare_only_bridge_exchange(resolver_result):
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_active_litellm_key",
"litellm.proxy._experimental.mcp_server.bridge_token_flow._resolve_active_litellm_key",
new=AsyncMock(return_value=resolver_result),
),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
@ -5505,7 +5503,7 @@ def test_classify_upstream_lifetime():
oversized) is "unspecified" so the envelope caps it, while a parseable non-positive value is
"expired": the upstream reporting an already-dead token, which the mint must reject rather than
silently give the 1h cap."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _classify_upstream_lifetime
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _classify_upstream_lifetime
assert _classify_upstream_lifetime(300) == 300
assert _classify_upstream_lifetime(300.0) == 300
@ -5538,7 +5536,7 @@ def test_bridge_grant_honors_and_rejects_upstream_lifetime():
"""The grant validator honors a positive lifetime, leaves an unknown one None for the envelope to
cap, and rejects an explicitly-expired one with "expired_lifetime" so a dead upstream token is
never sealed into an hour-long envelope."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _bridge_grant_from_token_response
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _bridge_grant_from_token_response
def grant(v):
return _bridge_grant_from_token_response({"access_token": "x", "expires_in": v})
@ -5831,7 +5829,7 @@ async def test_extract_user_id_reads_x_litellm_api_key_header(proxy_globals):
"""The LiteLLM key arrives on x-litellm-api-key (what Claude Desktop/Code send), not
Authorization. Reading only Authorization dropped the identity, so the per-user token was
never stored and the egress 401'd forever. Resolution must honor x-litellm-api-key."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request,
)
from litellm.proxy._types import UserAPIKeyAuth, hash_token
@ -5856,7 +5854,7 @@ async def test_extract_user_id_rehydrates_cross_replica_dict_cache(proxy_globals
"""Cross-replica, async_get_cache hands back a serialized dict, not a UserAPIKeyAuth.
Resolution must rehydrate it; the old getattr(cached, "user_id") returned None on a dict,
which is exactly why a multi-replica gateway never found the stored token."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request,
)
from litellm.proxy._types import hash_token
@ -5877,7 +5875,7 @@ async def test_extract_user_id_falls_back_to_db_on_cache_miss(proxy_globals):
"""A cache miss must read the key from the DB rather than returning None; the old code did a
cache-only peek and skipped the DB, so any replica that hadn't just authenticated the key
failed to store the token."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -5899,7 +5897,7 @@ async def test_extract_user_id_falls_back_to_db_on_cache_miss(proxy_globals):
@pytest.mark.asyncio
async def test_extract_user_id_none_without_litellm_key(proxy_globals):
"""No LiteLLM key on the request resolves to None without consulting the resolver."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -5916,7 +5914,7 @@ async def test_extract_user_id_rejects_blocked_key(proxy_globals):
"""A blocked LiteLLM key must not resolve an identity. get_key_object returns the DB row without
checking blocked/expiry (the main auth pipeline does, and the public token endpoint bypasses it),
so a revoked key could otherwise overwrite the stored per-user OAuth token for its user."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -5938,7 +5936,7 @@ async def test_extract_user_id_rejects_expired_key(proxy_globals):
"""An expired LiteLLM key must not resolve an identity, for the same reason as a blocked key."""
from datetime import datetime, timedelta, timezone
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -5963,7 +5961,7 @@ async def test_resolve_active_litellm_key_returns_resolved_key_for_active_key(pr
record. For an active key the resolver returns exactly hash_token(key), the same value
get_key_object and the whole cache/DB layer key the record by, so the sealed reference resolves
back to this key at admission."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_resolve_active_litellm_key,
_ResolvedKey,
)
@ -5993,10 +5991,10 @@ async def test_resolve_active_litellm_key_resolves_key_without_user_id(proxy_glo
presence wrongly rejected these keys with invalid_request; the active-state gate now checks only
blocked and expiry, and the key hash (not the user) is what the mint seals. The per-user token
store still gets no user for such a key, since there is none to key a stored credential by."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_extract_user_id_from_request,
_resolve_active_litellm_key,
_ResolvedKey,
_resolve_active_litellm_key,
)
from litellm.proxy._types import UserAPIKeyAuth, hash_token
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -6022,7 +6020,7 @@ async def test_resolve_active_litellm_key_resolves_key_without_user_id(proxy_glo
async def test_resolve_active_litellm_key_rejects_blocked_key(proxy_globals):
"""A blocked key must not yield a hash, so no gateway-bound envelope is minted for a revoked key;
the mint fails closed with invalid_request instead."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_resolve_active_litellm_key,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -6045,7 +6043,7 @@ async def test_resolve_active_litellm_key_fails_closed_on_malformed_expiry(proxy
returns invalid_request), not surface an unhandled 500. The active-state check runs outside the
resolver's try, so it must be total over a bad expires rather than letting datetime.fromisoformat
raise. Before the fix this raised a ValueError instead of returning None."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_resolve_active_litellm_key,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -6065,7 +6063,7 @@ async def test_resolve_active_litellm_key_fails_closed_on_malformed_expiry(proxy
@pytest.mark.asyncio
async def test_resolve_active_litellm_key_no_active_key_without_litellm_key(proxy_globals):
"""No LiteLLM key on the request yields no hash without consulting the resolver."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_resolve_active_litellm_key,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -6083,7 +6081,7 @@ async def test_resolve_active_litellm_key_db_outage_is_unavailable(proxy_globals
the caller's fault, so the resolver reports "unavailable" (the mint statuses it 503) rather than
collapsing it to the same value as a missing credential. is_database_service_unavailable_error
classifies a connection error (an OSError) as an outage, matching admission's egress-side handling."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_resolve_active_litellm_key,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -6104,7 +6102,7 @@ async def test_resolve_active_litellm_key_no_database_is_unresolvable(proxy_glob
"""With no database connection configured the gateway cannot verify the presented key at all, so
the resolver reports "unresolvable" (the mint statuses it 500) instead of blaming the caller.
Mirrors admission, which 500s a missing prisma_client on the egress side."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_resolve_active_litellm_key,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -6138,7 +6136,7 @@ async def test_reload_active_user_by_id_missing_user_is_no_active_key(proxy_glob
refresh path maps it to invalid_grant), not unresolvable/500. get_user_object catches the missing row
and re-raises a bare ValueError, so a missing user must not be misclassified as a DB outage or an
opaque gateway fault."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _reload_active_user_by_id
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
proxy_globals.user_api_key_cache = UserApiKeyCache()
@ -6157,7 +6155,7 @@ async def test_reload_active_user_by_id_db_outage_is_unavailable(proxy_globals):
a missing user, so the refresh path surfaces "unavailable" (a 503) rather than blaming the caller.
get_user_object wraps the outage in a bare ValueError, so this exercises the chain-aware classifier; a
raw ConnectionError would falsely pass even a chain-blind check because it is an OSError."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import _reload_active_user_by_id
from litellm.proxy._experimental.mcp_server.bridge_token_flow import _reload_active_user_by_id
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
proxy_globals.user_api_key_cache = UserApiKeyCache()

View file

@ -4445,3 +4445,85 @@ async def test_common_checks_personal_user_budget_blocks_in_gather():
request=MagicMock(spec=Request),
)
assert "User=u1" in str(over.value)
@pytest.mark.asyncio
async def test_user_budget_enforced_on_team_key():
"""User budget must be enforced even when the key belongs to a team.
Previously _user_max_budget_check skipped enforcement for team keys,
letting a user with a $100 personal budget spend unlimited through a
team key. This regression test ensures that is no longer the case.
"""
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0)
team = LiteLLM_TeamTable(team_id="t1", max_budget=2100.0)
token = UserAPIKeyAuth(token="k1", user_id="u1", team_id="t1")
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
return 999.0 if counter_key == "spend:user:u1" else 0.0
async def _no_membership(*a, **kw):
return None
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership):
with pytest.raises(litellm.BudgetExceededError) as over:
await common_checks(
request_body={"messages": [{"role": "user", "content": "hi"}]},
team_object=team,
user_object=user,
end_user_object=None,
global_proxy_spend=None,
general_settings={},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
valid_token=token,
request=MagicMock(spec=Request),
)
assert "User=u1" in str(over.value)
@pytest.mark.asyncio
async def test_skip_user_budget_on_team_key_flag_restores_old_behavior():
"""Setting skip_user_budget_on_team_key=True skips user budget for team keys.
This is the opt-in escape hatch that restores the legacy behavior where
user budgets were not enforced when the key belonged to a team.
"""
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0)
team = LiteLLM_TeamTable(team_id="t1", max_budget=2100.0)
token = UserAPIKeyAuth(token="k1", user_id="u1", team_id="t1")
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
return 999.0 if counter_key == "spend:user:u1" else 0.0
async def _no_membership(*a, **kw):
return None
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership):
result = await common_checks(
request_body={"messages": [{"role": "user", "content": "hi"}]},
team_object=team,
user_object=user,
end_user_object=None,
global_proxy_spend=None,
general_settings={"skip_user_budget_on_team_key": True},
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
valid_token=token,
request=MagicMock(spec=Request),
)
assert result is True

View file

@ -461,6 +461,53 @@ async def test_get_group_ids_from_service_principal_uses_configured_graph_endpoi
]
@pytest.mark.asyncio
async def test_get_group_ids_from_service_principal_paginates_through_all_pages():
# Arrange
page_one = {
"@odata.nextLink": "https://graph.microsoft.com/v1.0/servicePrincipals/sp-123/appRoleAssignedTo?$skiptoken=page2",
"value": [
{
"principalType": "Group",
"principalId": "group-on-page-1",
"principalDisplayName": "Group On Page 1",
}
],
}
page_two = {
"value": [
{
"principalType": "Group",
"principalId": "group-on-page-2",
"principalDisplayName": "Group On Page 2",
}
],
}
responses = [page_one, page_two]
async def mock_get(url, *args, **kwargs):
mock = MagicMock()
mock.json.return_value = responses.pop(0)
return mock
async_client = MagicMock()
async_client.get = mock_get
# Act
group_ids, teams = await MicrosoftSSOHandler.get_group_ids_from_service_principal(
service_principal_id="sp-123",
async_client=async_client,
access_token="mock_token",
)
# Assert
assert group_ids == ["group-on-page-1", "group-on-page-2"]
assert [team["principalId"] for team in teams] == [
"group-on-page-1",
"group-on-page-2",
]
def test_get_group_ids_from_graph_api_response():
# Arrange
mock_response = MicrosoftGraphAPIUserGroupResponse(

View file

@ -576,6 +576,83 @@ async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_should_reserve_user_budget_counter_for_team_key(spend_counter_state):
"""A user's personal budget must be reserved even when the key belongs to a team.
Regression for GitHub issue #12905: previously the reservation path skipped the
user spend counter whenever the key had a team, so a team key could overshoot the
user's personal max_budget under concurrency.
"""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-user-on-team",
spend=0.0,
user_id="user-on-team",
team_id="team-no-budget",
)
team_object = LiteLLM_TeamTable(team_id="team-no-budget", spend=0.0, max_budget=None)
user_object = LiteLLM_UserTable(user_id="user-on-team", spend=0.0, max_budget=5.0)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=0.3,
):
reservation = await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=team_object,
user_object=user_object,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert counter_cache.in_memory_cache.get_cache(key="spend:user:user-on-team") == pytest.approx(0.3)
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_should_skip_user_budget_counter_for_team_key_when_flag_set(spend_counter_state):
"""skip_user_budget_on_team_key=True restores the legacy behavior where a user's
personal budget is not reserved for a team key."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-user-on-team-skip",
spend=0.0,
user_id="user-on-team-skip",
team_id="team-no-budget-skip",
)
team_object = LiteLLM_TeamTable(team_id="team-no-budget-skip", spend=0.0, max_budget=None)
user_object = LiteLLM_UserTable(user_id="user-on-team-skip", spend=0.0, max_budget=5.0)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=0.3,
):
reservation = await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=team_object,
user_object=user_object,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
skip_user_budget_on_team_key=True,
)
assert counter_cache.in_memory_cache.get_cache(key="spend:user:user-on-team-skip") is None
await release_budget_reservation(reservation)
@pytest.mark.asyncio
async def test_should_seed_org_counter_from_with_budget_cache(spend_counter_state):
counter_cache, key_cache = spend_counter_state

View file

@ -8649,6 +8649,38 @@ def test_get_config_list_includes_cancel_on_disconnect(monkeypatch):
app.dependency_overrides.clear()
def test_get_config_list_includes_skip_user_budget_on_team_key(monkeypatch):
"""Related to #12905: the opt-out flag must be discoverable via /config/list so
it renders as a Boolean toggle on the Admin UI General Settings table. This
requires both the ConfigGeneralSettings field and the allowed_args entry."""
import types
from unittest.mock import AsyncMock, MagicMock
from fastapi.testclient import TestClient
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.proxy_server import app
mock_prisma = MagicMock()
mock_config_table = MagicMock()
mock_config_table.find_first = AsyncMock(return_value=None)
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
client = TestClient(app)
resp = client.get("/config/list", params={"config_type": "general_settings"})
assert resp.status_code == 200, resp.text
fields = {item["field_name"]: item for item in resp.json()}
assert "skip_user_budget_on_team_key" in fields
assert fields["skip_user_budget_on_team_key"]["field_type"] == "Boolean"
finally:
app.dependency_overrides.clear()
def test_get_config_list_includes_budget_exceeded_throttle_percentage(monkeypatch):
"""The throttle fraction is a litellm_settings scalar surfaced on the General
Settings table as a Float field so it sits with the other global limits; it

View file

@ -436,6 +436,33 @@ describe("Teams - premium props", () => {
});
});
describe("Teams - Create Team CTA is grouped with the tabs on the left", () => {
beforeEach(() => {
vi.clearAllMocks();
mockUseOrganizations.mockReturnValue({ data: [] });
});
it("renders the Create Team button inside the tab bar, ahead of the tabs", () => {
const { container } = renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
const createButton = screen.getByTestId("create-team-button");
const tabNav = container.querySelector(".ant-tabs-nav");
// The CTA lives in the tab bar's left slot, not the standalone page header.
expect(tabNav).not.toBeNull();
expect(tabNav!.contains(createButton)).toBe(true);
// It reads as the left end of the cluster: it precedes the first tab in DOM order.
const firstTab = screen.getByRole("tab", { name: "Your Teams" });
expect(createButton.compareDocumentPosition(firstTab) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy();
});
it("omits the Create Team CTA for a role that cannot manage teams", () => {
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin Viewer" />);
expect(screen.queryByTestId("create-team-button")).not.toBeInTheDocument();
});
});
describe("Teams - Default Team Settings tab visibility", () => {
beforeEach(() => {
vi.clearAllMocks();

View file

@ -564,18 +564,23 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
icon={<Users className="size-5" />}
title="Teams"
subtitle="Manage teams, members, and their access to models and budgets"
actions={
canCreateOrManageTeams(userRole, userID, organizations) ? (
/>
</div>
<Tabs
items={tabItems}
tabBarExtraContent={{
left: canCreateOrManageTeams(userRole, userID, organizations) ? (
<div className="flex items-center gap-4 pr-4">
<UIButton onClick={() => setIsTeamModalVisible(true)} data-testid="create-team-button">
<Plus className="size-4" />
Create Team
</UIButton>
) : undefined
}
/>
</div>
<Tabs items={tabItems} />
<div className="h-6 w-px bg-gray-200" />
</div>
) : undefined,
}}
/>
</>
)}

View file

@ -174,10 +174,19 @@ it("should render VirtualKeysTable component", () => {
expect(screen.getByText("Test Key Alias")).toBeInTheDocument();
});
it("renders the page header with the create-key action slot", () => {
it("left-anchors the create-key CTA below the title, between the header and the table toolbar", () => {
renderWithProviders(<VirtualKeysTable headerActions={<button>Create New Key</button>} />);
expect(screen.getByRole("heading", { name: "Virtual Keys" })).toBeInTheDocument();
expect(screen.getByRole("button", { name: "Create New Key" })).toBeInTheDocument();
const heading = screen.getByRole("heading", { name: "Virtual Keys" });
const ctas = screen.getAllByRole("button", { name: "Create New Key" });
expect(ctas).toHaveLength(1);
const cta = ctas[0];
const search = screen.getByPlaceholderText(/Search by key alias/);
// The CTA follows the title row...
expect(heading.compareDocumentPosition(cta) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy();
// ...and precedes the table's search toolbar, so it sits in its own row above the table.
expect(cta.compareDocumentPosition(search) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy();
});
it("should display key information correctly", async () => {

View file

@ -162,8 +162,8 @@ export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) {
icon={<KeyRound className="size-5" />}
title="Virtual Keys"
subtitle="Every key that authenticates requests to the gateway."
actions={headerActions}
/>
{headerActions}
<DataTable
data={keyList}
columns={columns}

View file

@ -0,0 +1,134 @@
import { Card, InputNumber, Radio, Slider, Space, Switch, Typography } from "antd";
import React from "react";
import {
AdaptiveEligible,
ComplexityRouterConfigValue,
DEFAULT_ADAPTIVE_WEIGHTS,
DEFAULT_TIER_DISTANCE_PENALTY,
} from "./ComplexityRouterConfig";
const { Text } = Typography;
interface AdaptiveRoutingConfigProps {
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
}
const AdaptiveRoutingConfig: React.FC<AdaptiveRoutingConfigProps> = ({ value, onChange }) => {
const adaptiveWeights = value.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS;
const adaptiveEligible = value.adaptive_eligible ?? "all";
const tierDistancePenalty = value.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY;
const handleAdaptiveToggle = (adaptive: boolean) => {
const nextValue: ComplexityRouterConfigValue = {
...value,
adaptive,
adaptive_weights: adaptiveWeights,
adaptive_eligible: adaptiveEligible,
tier_distance_penalty: tierDistancePenalty,
};
onChange(nextValue);
};
const handleQualityWeightChange = (qualityPercent: number) => {
const quality = qualityPercent / 100;
onChange({ ...value, adaptive_weights: { quality, cost: Math.round((1 - quality) * 100) / 100 } });
};
const handleAdaptiveEligibleChange = (eligible: AdaptiveEligible) => {
onChange({ ...value, adaptive_eligible: eligible });
};
const handleTierDistancePenaltyChange = (penalty: number | null) => {
onChange({ ...value, tier_distance_penalty: penalty ?? DEFAULT_TIER_DISTANCE_PENALTY });
};
return (
<>
<div className="flex items-center gap-2 mb-2">
<Switch checked={value.adaptive ?? false} onChange={handleAdaptiveToggle} />
<Text strong>Enable adaptive bandit selection</Text>
</div>
<Text type="secondary" style={{ display: "block", fontSize: 12 }}>
When disabled, each request always uses the model assigned to its classified tier.
</Text>
<Card className="bg-gray-50 mt-4">
<Text strong style={{ display: "block", marginBottom: 8 }}>
How Adaptive Routing Works
</Text>
<Text type="secondary" style={{ fontSize: 13 }}>
It learns from how each conversation actually goes: does the user have to rephrase or correct the model, does
it get stuck repeating itself, does it run out of tool calls, does the user seem satisfied. Combined with
cost, this live feedback shifts future routing toward the models that are actually working well, and improves
as more conversations come in. Until there&apos;s enough feedback, it defaults to the classified tier&apos;s
model.
</Text>
</Card>
{value.adaptive && (
<div className="mt-4 space-y-4">
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Quality vs. Cost ({Math.round(adaptiveWeights.quality * 100)}% quality /{" "}
{Math.round(adaptiveWeights.cost * 100)}% cost)
</Text>
<Slider
min={0}
max={100}
value={Math.round(adaptiveWeights.quality * 100)}
onChange={handleQualityWeightChange}
tooltip={{ formatter: (v) => `${v}% quality / ${100 - (v ?? 0)}% cost` }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Higher quality weight favors more capable (pricier) models; higher cost weight favors cheaper models when
the bandit has feedback to act on. Recommended: 30% quality / 70% cost split.
</Text>
</div>
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Eligible Model Pool
</Text>
<Radio.Group
value={adaptiveEligible}
onChange={(e) => handleAdaptiveEligibleChange(e.target.value)}
className="w-full"
>
<Space direction="vertical" className="w-full">
<Radio value="all">
<Text strong>All tiers (soft floor)</Text>{" "}
<Text type="secondary"> router can pick across tiers, depending on the best fit for the prompt</Text>
</Radio>
<Radio value="classified_tier">
<Text strong>Classified tier only</Text>{" "}
<Text type="secondary"> router can only pick models within tier</Text>
</Radio>
</Space>
</Radio.Group>
</div>
{adaptiveEligible === "all" && (
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Tier Distance Penalty
</Text>
<InputNumber
value={tierDistancePenalty}
onChange={handleTierDistancePenaltyChange}
min={0}
step={0.1}
style={{ width: "100%" }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Score penalty applied per tier-step away from the classified tier.
</Text>
</div>
)}
</div>
)}
</>
);
};
export default AdaptiveRoutingConfig;

View file

@ -0,0 +1,171 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Select as AntdSelect, Card, InputNumber, Radio, Space, Tooltip, Typography } from "antd";
import React from "react";
import { ClassifierType, ComplexityRouterConfigValue, DEFAULT_CLASSIFIER_TIMEOUT_MS } from "./ComplexityRouterConfig";
const { Text } = Typography;
interface ClassificationMethodConfigProps {
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
modelOptions: { value: string; label: string }[];
customTechnicalKeywords?: string[];
onCustomTechnicalKeywordsChange?: (keywords: string[]) => void;
showValidationErrors?: boolean;
}
const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
value,
onChange,
modelOptions,
customTechnicalKeywords,
onCustomTechnicalKeywordsChange,
showValidationErrors = false,
}) => {
const classifierModelMissing =
showValidationErrors && value.classifier_type === "llm" && !value.classifier_llm_config?.model;
const handleClassifierTypeChange = (classifierType: ClassifierType) => {
onChange({
...value,
classifier_type: classifierType,
classifier_llm_config:
classifierType === "llm"
? value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }
: undefined,
});
};
const handleClassifierModelChange = (model: string) => {
onChange({
...value,
classifier_llm_config: {
model,
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
const handleClassifierTimeoutChange = (timeoutMs: number | null) => {
onChange({
...value,
classifier_llm_config: {
model: value.classifier_llm_config?.model ?? "",
timeout_ms: timeoutMs ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
return (
<>
<Radio.Group
value={value.classifier_type}
onChange={(e) => handleClassifierTypeChange(e.target.value)}
className="w-full"
>
<Space direction="vertical" className="w-full">
<Radio value="heuristic">
<Text strong>Heuristic</Text>{" "}
<Text type="secondary">(default) rule-based scoring, no API calls, &lt;1ms latency</Text>
</Radio>
<Radio value="llm">
<Text strong>LLM Classifier</Text>{" "}
<Text type="secondary"> use a model to decide the tier (e.g. a small/fast model)</Text>
</Radio>
</Space>
</Radio.Group>
{value.classifier_type === "llm" && (
<div className="mt-4 space-y-3">
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Classifier Model
</Text>
<AntdSelect
value={value.classifier_llm_config?.model || undefined}
onChange={handleClassifierModelChange}
placeholder="Select the model that will classify request complexity"
showSearch
style={{ width: "100%" }}
options={modelOptions}
status={classifierModelMissing ? "error" : undefined}
/>
{classifierModelMissing && (
<Text type="danger" style={{ fontSize: 12 }}>
A classifier model is required
</Text>
)}
</div>
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Timeout (ms)
</Text>
<InputNumber
value={value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS}
onChange={handleClassifierTimeoutChange}
min={1}
style={{ width: "100%" }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Falls back to the heuristic scorer if the classifier call errors, times out, or returns an unparseable
response.
</Text>
</div>
</div>
)}
{value.classifier_type === "heuristic" && (
<div className="mt-4">
<div className="flex items-center gap-2 mb-1">
<Text strong>Custom Technical Keywords</Text>
<Tooltip title="Domain-specific terms appended to the built-in technical keyword list. Prompts containing these terms score higher on the technical dimension and route to more capable models.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</div>
<Text type="secondary" style={{ display: "block", marginBottom: 8, fontSize: 12 }}>
Optional: Add terms to the built-in list to improve classification accuracy on the technical dimension.
(e.g., udp, kafka, terraform).
</Text>
<AntdSelect
mode="tags"
value={customTechnicalKeywords ?? []}
onChange={(keywords: string[]) => onCustomTechnicalKeywordsChange?.(keywords)}
placeholder="Type a keyword and press Enter, or paste a comma-separated list"
tokenSeparators={[","]}
open={false}
suffixIcon={null}
style={{ width: "100%" }}
allowClear
/>
</div>
)}
<Card className="bg-gray-50 mt-4">
<Text strong style={{ display: "block", marginBottom: 8 }}>
How Classification Works
</Text>
<Text type="secondary" style={{ fontSize: 13 }}>
The router scores each request across 7 dimensions: token count, code presence, reasoning markers, technical
terms, simple indicators, multi-step patterns, and question complexity. The weighted score determines the
tier:
</Text>
<ul style={{ marginTop: 8, marginBottom: 0, paddingLeft: 20, fontSize: 13, color: "rgba(0, 0, 0, 0.45)" }}>
<li>
<strong>SIMPLE</strong>: Score &lt; 0.15
</li>
<li>
<strong>MEDIUM</strong>: Score 0.15 - 0.35
</li>
<li>
<strong>COMPLEX</strong>: Score 0.35 - 0.60
</li>
<li>
<strong>REASONING</strong>: Score &gt; 0.60 (or 2+ reasoning markers)
</li>
</ul>
</Card>
</>
);
};
export default ClassificationMethodConfig;

View file

@ -12,10 +12,10 @@ const mockModelInfo = [
const defaultValue: ComplexityRouterConfigValue = {
tiers: {
SIMPLE: "gpt-3.5-turbo",
MEDIUM: "gpt-3.5-turbo",
COMPLEX: "gpt-4",
REASONING: "claude-3-opus",
SIMPLE: ["gpt-3.5-turbo"],
MEDIUM: ["gpt-3.5-turbo"],
COMPLEX: ["gpt-4"],
REASONING: ["claude-3-opus"],
},
classifier_type: "heuristic",
};
@ -58,11 +58,13 @@ describe("ComplexityRouterConfig", () => {
it("should display the how classification works section", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText("How Classification Works")).toBeInTheDocument();
});
it("should show score thresholds in the classification section", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText(/Score < 0.15/)).toBeInTheDocument();
expect(screen.getByText(/Score 0.15 - 0.35/)).toBeInTheDocument();
expect(screen.getByText(/Score 0.35 - 0.60/)).toBeInTheDocument();
@ -107,6 +109,7 @@ describe("ComplexityRouterConfig", () => {
it("should render the custom technical keywords field", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText("Custom Technical Keywords")).toBeInTheDocument();
});
@ -118,6 +121,7 @@ describe("ComplexityRouterConfig", () => {
onCustomTechnicalKeywordsChange={vi.fn()}
/>,
);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText("udp")).toBeInTheDocument();
expect(screen.getByText("kafka")).toBeInTheDocument();
});
@ -132,14 +136,16 @@ describe("ComplexityRouterConfig", () => {
onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange}
/>,
);
const keywordsCard = screen.getByText("Custom Technical Keywords").closest(".ant-card") as HTMLElement;
const input = within(keywordsCard).getByRole("combobox");
fireEvent.click(screen.getByText("Advanced: Classification Method"));
const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement;
const input = within(keywordsSection).getByRole("combobox");
await user.type(input, "udp,");
expect(onCustomTechnicalKeywordsChange).toHaveBeenCalledWith(["udp"]);
});
it("should render an empty state when no keyword tier rules exist", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("Keyword Tier Overrides")).toBeInTheDocument();
expect(screen.getByText("No keyword tier overrides configured")).toBeInTheDocument();
});
@ -158,6 +164,7 @@ describe("ComplexityRouterConfig", () => {
const user = userEvent.setup();
const onKeywordTierRulesChange = vi.fn();
renderWithProviders(<ComplexityRouterConfig {...baseProps} onKeywordTierRulesChange={onKeywordTierRulesChange} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
await user.click(screen.getByRole("button", { name: /add keyword rule/i }));
expect(onKeywordTierRulesChange).toHaveBeenCalledTimes(1);
const newRules = onKeywordTierRulesChange.mock.calls[0][0];
@ -175,6 +182,7 @@ describe("ComplexityRouterConfig", () => {
onKeywordTierRulesChange={onKeywordTierRulesChange}
/>,
);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("invoice")).toBeInTheDocument();
expect(screen.getByText("refund")).toBeInTheDocument();
@ -184,6 +192,7 @@ describe("ComplexityRouterConfig", () => {
it("should not show embedding model or match score fields when semantic matching is disabled", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} semanticMatchingEnabled={false} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("Semantic keyword matching")).toBeInTheDocument();
expect(screen.queryByText("Embedding model")).not.toBeInTheDocument();
expect(screen.queryByText("Minimum match score")).not.toBeInTheDocument();
@ -191,6 +200,7 @@ describe("ComplexityRouterConfig", () => {
it("should show embedding model and match score fields when semantic matching is enabled", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} semanticMatchingEnabled={true} />);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
expect(screen.getByText("Embedding model")).toBeInTheDocument();
expect(screen.getByText("Minimum match score")).toBeInTheDocument();
});
@ -205,6 +215,7 @@ describe("ComplexityRouterConfig", () => {
onSemanticMatchingEnabledChange={onSemanticMatchingEnabledChange}
/>,
);
fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching"));
await user.click(screen.getByRole("switch"));
expect(onSemanticMatchingEnabledChange).toHaveBeenCalledWith(true, expect.anything());
});
@ -252,7 +263,7 @@ describe("ComplexityRouterConfig", () => {
renderWithProviders(
<ComplexityRouterConfig
{...baseProps}
value={{ ...defaultValue, tiers: { ...defaultValue.tiers, REASONING: "" } }}
value={{ ...defaultValue, tiers: { ...defaultValue.tiers, REASONING: [] } }}
showValidationErrors={true}
/>,
);

View file

@ -1,19 +1,22 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Select as AntdSelect, Card, Collapse, Divider, InputNumber, Radio, Space, Tooltip, Typography } from "antd";
import { Select as AntdSelect, Card, Collapse, Divider, Space, Tooltip, Typography } from "antd";
import React from "react";
import { ModelGroup } from "@/components/llm_calls/fetch_models";
import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig";
import ClassificationMethodConfig from "./ClassificationMethodConfig";
import KeywordTierRules, { KeywordTierRule } from "./KeywordTierRules";
import SemanticKeywordMatching from "./SemanticKeywordMatching";
const { Text } = Typography;
export const DEFAULT_CLASSIFIER_TIMEOUT_MS = 3000;
export const DEFAULT_TIER_DISTANCE_PENALTY = 0.5;
export interface ComplexityTiers {
SIMPLE: string;
MEDIUM: string;
COMPLEX: string;
REASONING: string;
SIMPLE: string[];
MEDIUM: string[];
COMPLEX: string[];
REASONING: string[];
}
export interface ClassifierLLMConfig {
@ -23,10 +26,23 @@ export interface ClassifierLLMConfig {
export type ClassifierType = "heuristic" | "llm";
export interface AdaptiveRouterWeights {
quality: number;
cost: number;
}
export const DEFAULT_ADAPTIVE_WEIGHTS: AdaptiveRouterWeights = { quality: 0.3, cost: 0.7 };
export type AdaptiveEligible = "all" | "classified_tier";
export interface ComplexityRouterConfigValue {
tiers: ComplexityTiers;
classifier_type: ClassifierType;
classifier_llm_config?: ClassifierLLMConfig;
adaptive?: boolean;
adaptive_weights?: AdaptiveRouterWeights;
tier_distance_penalty?: number;
adaptive_eligible?: AdaptiveEligible;
}
interface ComplexityRouterConfigProps {
@ -95,44 +111,10 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
label: model.model_group,
}));
const classifierModelMissing =
showValidationErrors && value.classifier_type === "llm" && !value.classifier_llm_config?.model;
const handleTierChange = (tier: keyof ComplexityTiers, model: string) => {
const handleTierChange = (tier: keyof ComplexityTiers, models: string[]) => {
onChange({
...value,
tiers: { ...value.tiers, [tier]: model },
});
};
const handleClassifierTypeChange = (classifierType: ClassifierType) => {
onChange({
...value,
classifier_type: classifierType,
classifier_llm_config:
classifierType === "llm"
? (value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS })
: undefined,
});
};
const handleClassifierModelChange = (model: string) => {
onChange({
...value,
classifier_llm_config: {
model,
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
const handleClassifierTimeoutChange = (timeoutMs: number | null) => {
onChange({
...value,
classifier_llm_config: {
model: value.classifier_llm_config?.model ?? "",
timeout_ms: timeoutMs ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
tiers: { ...value.tiers, [tier]: models },
});
};
@ -142,20 +124,20 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
<Typography.Title level={4} style={{ margin: 0 }}>
Complexity Tier Configuration
</Typography.Title>
<Tooltip title="Map each complexity tier to a model. Simple queries use cheaper/faster models, complex queries use more capable models.">
<Tooltip title="Map each complexity tier to one or more models. Simple queries use cheaper/faster models, complex queries use more capable models.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</Space>
<Text type="secondary" style={{ display: "block", marginBottom: 24 }}>
The complexity router automatically classifies requests by complexity using rule-based scoring (no API calls,
&lt;1ms latency). Configure which model handles each tier.
&lt;1ms latency). Configure which model(s) handle each tier.
</Text>
<Card>
{(Object.keys(TIER_DESCRIPTIONS) as Array<keyof ComplexityTiers>).map((tier, index) => {
const tierInfo = TIER_DESCRIPTIONS[tier];
const tierMissing = showValidationErrors && !value.tiers[tier];
const tierMissing = showValidationErrors && value.tiers[tier].length === 0;
return (
<div key={tier}>
{index > 0 && <Divider style={{ margin: "16px 0" }} />}
@ -172,14 +154,21 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
Examples: {tierInfo.examples}
</Text>
<AntdSelect
mode="multiple"
value={value.tiers[tier]}
onChange={(model) => handleTierChange(tier, model)}
placeholder={`Select model for ${tierInfo.label.toLowerCase()} queries`}
onChange={(models) => handleTierChange(tier, models)}
placeholder={`Select model(s) for ${tierInfo.label.toLowerCase()} queries`}
showSearch
style={{ width: "100%" }}
options={modelOptions}
status={tierMissing ? "error" : undefined}
/>
{value.tiers[tier].length > 1 && (
<Text type="secondary" style={{ fontSize: 12 }}>
Multiple models selected the router randomly picks among them per request (or Thompson-samples
within the pool when adaptive routing is on).
</Text>
)}
{tierMissing && (
<Text type="danger" style={{ fontSize: 12 }}>
This tier is required
@ -205,148 +194,61 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
</Text>
),
children: (
<>
<Radio.Group
value={value.classifier_type}
onChange={(e) => handleClassifierTypeChange(e.target.value)}
className="w-full"
>
<Space direction="vertical" className="w-full">
<Radio value="heuristic">
<Text strong>Heuristic</Text>{" "}
<Text type="secondary">(default) rule-based scoring, no API calls, &lt;1ms latency</Text>
</Radio>
<Radio value="llm">
<Text strong>LLM Classifier</Text>{" "}
<Text type="secondary"> use a model to decide the tier (e.g. a small/fast model)</Text>
</Radio>
</Space>
</Radio.Group>
{value.classifier_type === "llm" && (
<div className="mt-4 space-y-3">
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Classifier Model
</Text>
<AntdSelect
value={value.classifier_llm_config?.model || undefined}
onChange={handleClassifierModelChange}
placeholder="Select the model that will classify request complexity"
showSearch
style={{ width: "100%" }}
options={modelOptions}
status={classifierModelMissing ? "error" : undefined}
/>
{classifierModelMissing && (
<Text type="danger" style={{ fontSize: 12 }}>
A classifier model is required
</Text>
)}
</div>
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Timeout (ms)
</Text>
<InputNumber
value={value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS}
onChange={handleClassifierTimeoutChange}
min={1}
style={{ width: "100%" }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Falls back to the heuristic scorer if the classifier call errors, times out, or returns an
unparseable response.
</Text>
</div>
</div>
)}
</>
<ClassificationMethodConfig
value={value}
onChange={onChange}
modelOptions={modelOptions}
customTechnicalKeywords={customTechnicalKeywords}
onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange}
showValidationErrors={showValidationErrors}
/>
),
},
{
key: "adaptive",
label: (
<Text strong style={{ color: "#374151" }}>
Advanced: Adaptive Routing
</Text>
),
children: <AdaptiveRoutingConfig value={value} onChange={onChange} />,
},
...(onKeywordTierRulesChange || onSemanticMatchingEnabledChange
? [
{
key: "keyword-semantic",
label: (
<Text strong style={{ color: "#374151" }}>
Advanced: Keyword/Semantic Matching
</Text>
),
children: (
<>
{onKeywordTierRulesChange && (
<KeywordTierRules rules={keywordTierRules} onChange={onKeywordTierRulesChange} />
)}
{onKeywordTierRulesChange && onSemanticMatchingEnabledChange && (
<Divider style={{ margin: "16px 0" }} />
)}
{onSemanticMatchingEnabledChange && (
<SemanticKeywordMatching
enabled={semanticMatchingEnabled}
onEnabledChange={onSemanticMatchingEnabledChange}
embeddingModel={embeddingModel}
onEmbeddingModelChange={onEmbeddingModelChange}
matchThreshold={matchThreshold}
onMatchThresholdChange={onMatchThresholdChange}
modelInfo={modelInfo}
showValidationErrors={showValidationErrors}
/>
)}
</>
),
},
]
: []),
]}
/>
<Divider />
<Card>
<div className="flex items-center gap-2 mb-2">
<Text strong style={{ fontSize: 16 }}>
Custom Technical Keywords
</Text>
<Tooltip title="Domain-specific terms appended to the built-in technical keyword list. Prompts containing these terms score higher on the technical dimension and route to more capable models.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</div>
<Text type="secondary" style={{ display: "block", marginBottom: 8, fontSize: 12 }}>
Optional: Add terms to the built-in list to improve classification accuracy on the technical dimension. (e.g.,
udp, kafka, terraform).
</Text>
<AntdSelect
mode="tags"
value={customTechnicalKeywords ?? []}
onChange={(keywords: string[]) => onCustomTechnicalKeywordsChange?.(keywords)}
placeholder="Type a keyword and press Enter, or paste a comma-separated list"
tokenSeparators={[","]}
open={false}
suffixIcon={null}
style={{ width: "100%" }}
allowClear
/>
</Card>
<Divider />
<Card className="bg-gray-50">
<Text strong style={{ display: "block", marginBottom: 8 }}>
How Classification Works
</Text>
<Text type="secondary" style={{ fontSize: 13 }}>
The router scores each request across 7 dimensions: token count, code presence, reasoning markers, technical
terms, simple indicators, multi-step patterns, and question complexity. The weighted score determines the
tier:
</Text>
<ul style={{ marginTop: 8, marginBottom: 0, paddingLeft: 20, fontSize: 13, color: "rgba(0, 0, 0, 0.45)" }}>
<li>
<strong>SIMPLE</strong>: Score &lt; 0.15
</li>
<li>
<strong>MEDIUM</strong>: Score 0.15 - 0.35
</li>
<li>
<strong>COMPLEX</strong>: Score 0.35 - 0.60
</li>
<li>
<strong>REASONING</strong>: Score &gt; 0.60 (or 2+ reasoning markers)
</li>
</ul>
</Card>
{/* Keyword-tier and semantic sections only render when their change handlers are
wired (the add-router flow). The edit-auto-router modal doesn't pass them yet, so
they stay hidden there rather than rendering interactive-but-dead controls. */}
{onKeywordTierRulesChange && (
<>
<Divider />
<KeywordTierRules rules={keywordTierRules} onChange={onKeywordTierRulesChange} />
</>
)}
{onSemanticMatchingEnabledChange && (
<>
<Divider />
<SemanticKeywordMatching
enabled={semanticMatchingEnabled}
onEnabledChange={onSemanticMatchingEnabledChange}
embeddingModel={embeddingModel}
onEmbeddingModelChange={onEmbeddingModelChange}
matchThreshold={matchThreshold}
onMatchThresholdChange={onMatchThresholdChange}
modelInfo={modelInfo}
showValidationErrors={showValidationErrors}
/>
</>
)}
</div>
);
};

View file

@ -1,5 +1,5 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Card, InputNumber, Select as AntdSelect, Switch, Tooltip, Typography } from "antd";
import { InputNumber, Select as AntdSelect, Switch, Tooltip, Typography } from "antd";
import React from "react";
import { ModelGroup } from "@/components/llm_calls/fetch_models";
@ -36,7 +36,7 @@ const SemanticKeywordMatching: React.FC<SemanticKeywordMatchingProps> = ({
const embeddingModelMissing = showValidationErrors && !embeddingModel;
return (
<Card className="mb-4">
<div>
<div className="flex items-center justify-between gap-4">
<div>
<div className="flex items-center gap-2">
@ -86,7 +86,7 @@ const SemanticKeywordMatching: React.FC<SemanticKeywordMatchingProps> = ({
</div>
</div>
)}
</Card>
</div>
);
};

View file

@ -8,7 +8,11 @@ import { all_admin_roles } from "@/utils/roles";
import { handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import RouterConfigBuilder from "./RouterConfigBuilder";
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
import ComplexityRouterConfig, {
ComplexityRouterConfigValue,
DEFAULT_ADAPTIVE_WEIGHTS,
DEFAULT_TIER_DISTANCE_PENALTY,
} from "./ComplexityRouterConfig";
import { KeywordTierRule } from "./KeywordTierRules";
import { DEFAULT_MATCH_THRESHOLD } from "./SemanticKeywordMatching";
import {
@ -38,7 +42,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
const [routerType, setRouterType] = useState<RouterType>("recommended");
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] },
classifier_type: "heuristic",
});
@ -89,6 +93,10 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
tiers,
classifier_type: classifierType,
classifier_llm_config: classifierLlmConfig,
adaptive = false,
adaptive_weights: adaptiveWeights = DEFAULT_ADAPTIVE_WEIGHTS,
tier_distance_penalty: tierDistancePenalty = DEFAULT_TIER_DISTANCE_PENALTY,
adaptive_eligible: adaptiveEligible = "all",
} = complexityRouterConfig;
const missingTiersError = getMissingTiersError(tiers);
@ -111,7 +119,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
return;
}
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
const defaultModel = tiers.MEDIUM[0] || tiers.SIMPLE[0] || tiers.COMPLEX[0] || tiers.REASONING[0];
form.setFieldsValue({
custom_llm_provider: "auto_router",
@ -132,6 +140,10 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
semanticMatchingEnabled,
embeddingModel,
matchThreshold,
adaptive,
adaptiveWeights,
tierDistancePenalty,
adaptiveEligible,
};
const submitValues = {

View file

@ -1,10 +1,10 @@
import { buildAutoRouterTestTargets } from "./build_auto_router_test_targets";
const tiers = {
SIMPLE: "gpt-4o-mini",
MEDIUM: "claude-sonnet-4",
COMPLEX: "claude-sonnet-4",
REASONING: "o3",
SIMPLE: ["gpt-4o-mini"],
MEDIUM: ["claude-sonnet-4"],
COMPLEX: ["claude-sonnet-4"],
REASONING: ["o3"],
};
describe("buildAutoRouterTestTargets", () => {
@ -17,9 +17,21 @@ describe("buildAutoRouterTestTargets", () => {
]);
});
it("emits a target per model when a tier has more than one, and dedups across tiers", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: ["gpt-4o-mini", "claude-sonnet-4"], MEDIUM: ["claude-sonnet-4"], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: undefined,
});
expect(targets).toEqual([
{ labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" },
{ labels: ["SIMPLE", "MEDIUM"], modelGroup: "claude-sonnet-4", mode: "chat" },
]);
});
it("drops empty/whitespace tiers", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: " ", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [" "], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: undefined,
});
@ -29,7 +41,7 @@ describe("buildAutoRouterTestTargets", () => {
it("returns [] when no tier is configured", () => {
expect(
buildAutoRouterTestTargets({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: undefined,
}),
@ -38,7 +50,7 @@ describe("buildAutoRouterTestTargets", () => {
it("appends an embedding target only when semantic matching is on and a model is set", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: true,
embeddingModel: "voyage-3-5",
});
@ -50,7 +62,7 @@ describe("buildAutoRouterTestTargets", () => {
it("omits the embedding target when semantic matching is on but no model is chosen", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: true,
embeddingModel: undefined,
});
@ -59,7 +71,7 @@ describe("buildAutoRouterTestTargets", () => {
it("omits the embedding target when a model is set but semantic matching is off", () => {
const targets = buildAutoRouterTestTargets({
tiers: { SIMPLE: "gpt-4o-mini", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [], REASONING: [] },
semanticMatchingEnabled: false,
embeddingModel: "voyage-3-5",
});

View file

@ -29,9 +29,11 @@ export const buildAutoRouterTestTargets = ({
embeddingModel,
}: BuildAutoRouterTestTargetsParams): AutoRouterTestTarget[] => {
const groupedByModel = TIER_ORDER.reduce<Record<string, string[]>>((acc, tier) => {
const modelGroup = tiers[tier]?.trim();
if (!modelGroup) return acc;
return { ...acc, [modelGroup]: [...(acc[modelGroup] ?? []), tier] };
return (tiers[tier] ?? []).reduce((tierAcc, rawModel) => {
const modelGroup = rawModel?.trim();
if (!modelGroup) return tierAcc;
return { ...tierAcc, [modelGroup]: [...(tierAcc[modelGroup] ?? []), tier] };
}, acc);
}, {});
const tierTargets: AutoRouterTestTarget[] = Object.entries(groupedByModel).map(([modelGroup, labels]) => ({

View file

@ -6,10 +6,10 @@ import {
} from "./build_complexity_router_config";
const tiers = {
SIMPLE: "gpt-4o-mini",
MEDIUM: "gpt-4o",
COMPLEX: "claude-sonnet-4",
REASONING: "o1-preview",
SIMPLE: ["gpt-4o-mini"],
MEDIUM: ["gpt-4o"],
COMPLEX: ["claude-sonnet-4"],
REASONING: ["o1-preview"],
};
const baseParams: BuildComplexityRouterConfigParams = {
@ -21,6 +21,10 @@ const baseParams: BuildComplexityRouterConfigParams = {
semanticMatchingEnabled: false,
embeddingModel: undefined,
matchThreshold: 0.5,
adaptive: false,
adaptiveWeights: { quality: 0.3, cost: 0.7 },
tierDistancePenalty: 0.5,
adaptiveEligible: "all",
};
describe("buildComplexityRouterConfig", () => {
@ -29,6 +33,14 @@ describe("buildComplexityRouterConfig", () => {
expect(config).toEqual({ tiers, classifier_type: "heuristic" });
});
it("passes through a tier configured with more than one model as a pool", () => {
const config = buildComplexityRouterConfig({
...baseParams,
tiers: { ...tiers, SIMPLE: ["gpt-4o-mini", "gpt-4o", "claude-haiku-4-5"] },
});
expect(config.tiers.SIMPLE).toEqual(["gpt-4o-mini", "gpt-4o", "claude-haiku-4-5"]);
});
it("includes classifier_llm_config only when classifier_type is llm", () => {
const config = buildComplexityRouterConfig({
...baseParams,
@ -123,6 +135,47 @@ describe("buildComplexityRouterConfig", () => {
const config = buildComplexityRouterConfig(params);
expect(config.keyword_tier_rules).toBeUndefined();
});
it("omits adaptive fields when adaptive is disabled even if weights linger in state", () => {
const config = buildComplexityRouterConfig({
...baseParams,
adaptive: false,
adaptiveWeights: { quality: 0.9, cost: 0.1 },
tierDistancePenalty: 2,
adaptiveEligible: "classified_tier",
});
expect(config.adaptive).toBeUndefined();
expect(config.adaptive_weights).toBeUndefined();
expect(config.tier_distance_penalty).toBeUndefined();
expect(config.adaptive_eligible).toBeUndefined();
});
it("includes tier_distance_penalty when adaptive is enabled with eligible='all'", () => {
const config = buildComplexityRouterConfig({
...baseParams,
adaptive: true,
adaptiveWeights: { quality: 0.6, cost: 0.4 },
tierDistancePenalty: 0.75,
adaptiveEligible: "all",
});
expect(config.adaptive).toBe(true);
expect(config.adaptive_weights).toEqual({ quality: 0.6, cost: 0.4 });
expect(config.tier_distance_penalty).toBe(0.75);
expect(config.adaptive_eligible).toBe("all");
});
it("omits tier_distance_penalty when eligible='classified_tier', since the penalty doesn't apply there", () => {
const config = buildComplexityRouterConfig({
...baseParams,
adaptive: true,
adaptiveWeights: { quality: 0.6, cost: 0.4 },
tierDistancePenalty: 0.75,
adaptiveEligible: "classified_tier",
});
expect(config.adaptive).toBe(true);
expect(config.adaptive_eligible).toBe("classified_tier");
expect(config.tier_distance_penalty).toBeUndefined();
});
});
describe("getMissingTiersError", () => {
@ -131,23 +184,27 @@ describe("getMissingTiersError", () => {
});
it("names the specific missing tier when only one is blank", () => {
expect(getMissingTiersError({ ...tiers, REASONING: "" })).toBe(
expect(getMissingTiersError({ ...tiers, REASONING: [] })).toBe(
"Select a model for the following tier(s): REASONING",
);
});
it("names multiple missing tiers in SIMPLE/MEDIUM/COMPLEX/REASONING order", () => {
expect(getMissingTiersError({ ...tiers, SIMPLE: "", REASONING: "" })).toBe(
expect(getMissingTiersError({ ...tiers, SIMPLE: [], REASONING: [] })).toBe(
"Select a model for the following tier(s): SIMPLE, REASONING",
);
});
it("names all four tiers when none are filled", () => {
const noTiers = { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" };
const noTiers = { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] };
expect(getMissingTiersError(noTiers)).toBe(
"Select a model for the following tier(s): SIMPLE, MEDIUM, COMPLEX, REASONING",
);
});
it("treats a tier with more than one model as filled", () => {
expect(getMissingTiersError({ ...tiers, SIMPLE: ["gpt-4o-mini", "gpt-4o"] })).toBeNull();
});
});
describe("getSemanticConfigError", () => {

View file

@ -1,12 +1,11 @@
import { KeywordTierRule } from "./KeywordTierRules";
import { ClassifierLLMConfig, ClassifierType } from "./ComplexityRouterConfig";
export interface ComplexityTiers {
SIMPLE: string;
MEDIUM: string;
COMPLEX: string;
REASONING: string;
}
import {
AdaptiveEligible,
AdaptiveRouterWeights,
ClassifierLLMConfig,
ClassifierType,
ComplexityTiers,
} from "./ComplexityRouterConfig";
export interface BuildComplexityRouterConfigParams {
tiers: ComplexityTiers;
@ -17,6 +16,10 @@ export interface BuildComplexityRouterConfigParams {
semanticMatchingEnabled: boolean;
embeddingModel: string | undefined;
matchThreshold: number;
adaptive: boolean;
adaptiveWeights: AdaptiveRouterWeights;
tierDistancePenalty: number;
adaptiveEligible: AdaptiveEligible;
}
export interface ComplexityRouterConfigPayload {
@ -28,12 +31,16 @@ export interface ComplexityRouterConfigPayload {
semantic_keyword_matching?: boolean;
embedding_model?: string;
match_threshold?: number;
adaptive?: boolean;
adaptive_weights?: AdaptiveRouterWeights;
tier_distance_penalty?: number;
adaptive_eligible?: AdaptiveEligible;
}
const TIER_KEYS: Array<keyof ComplexityTiers> = ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"];
export const getMissingTiersError = (tiers: ComplexityTiers): string | null => {
const missing = TIER_KEYS.filter((tier) => !tiers[tier]);
const missing = TIER_KEYS.filter((tier) => tiers[tier].length === 0);
if (missing.length === 0) return null;
return `Select a model for the following tier(s): ${missing.join(", ")}`;
};
@ -43,7 +50,8 @@ export const getSemanticConfigError = ({
embeddingModel,
keywordTierRules,
}: Pick<BuildComplexityRouterConfigParams, "semanticMatchingEnabled" | "embeddingModel" | "keywordTierRules">):
string | null => {
| string
| null => {
if (!semanticMatchingEnabled) return null;
if (!embeddingModel) return "Select an embedding model to use semantic keyword matching";
if (keywordTierRules.length === 0) return "Add at least one keyword tier rule to use semantic keyword matching";
@ -61,6 +69,10 @@ export const buildComplexityRouterConfig = ({
semanticMatchingEnabled,
embeddingModel,
matchThreshold,
adaptive,
adaptiveWeights,
tierDistancePenalty,
adaptiveEligible,
}: BuildComplexityRouterConfigParams): ComplexityRouterConfigPayload => {
// Trim keywords and drop empty ones; drop any rule left with no keywords. Clicking
// "Add keyword rule" seeds a rule with an empty keywords list, so without this an
@ -81,5 +93,11 @@ export const buildComplexityRouterConfig = ({
embedding_model: embeddingModel,
match_threshold: matchThreshold,
}),
...(adaptive && {
adaptive: true,
adaptive_weights: adaptiveWeights,
...(adaptiveEligible === "all" && { tier_distance_penalty: tierDistancePenalty }),
adaptive_eligible: adaptiveEligible,
}),
};
};

View file

@ -0,0 +1,104 @@
import { buildUpdatedComplexityRouterConfig } from "./edit_auto_router_modal";
const storedConfigValue = {
tiers: {
SIMPLE: "old-simple",
MEDIUM: "old-medium",
COMPLEX: "old-complex",
REASONING: "old-reasoning",
},
classifier_type: "llm",
classifier_llm_config: { model: "old-classifier", timeout_ms: 1200 },
custom_technical_keywords: ["kafka", "terraform"],
keyword_tier_rules: [{ keywords: ["invoice", "refund"], tier: "MEDIUM" }],
semantic_keyword_matching: true,
embedding_model: "voyage-4-large",
match_threshold: 0.65,
adaptive: true,
adaptive_weights: { quality: 0.3, cost: 0.7 },
tier_distance_penalty: 0.8,
adaptive_eligible: "all",
};
const storedConfig = JSON.stringify(storedConfigValue);
const tiers = {
SIMPLE: ["gpt-4o-mini"],
MEDIUM: ["gpt-4o-mini"],
COMPLEX: ["anthropic-sonnet-4-5"],
REASONING: ["anthropic-sonnet-4-5"],
};
const classifiedTierValue = {
tiers,
classifier_type: "heuristic" as const,
adaptive: true,
adaptive_weights: { quality: 0.4, cost: 0.6 },
tier_distance_penalty: 0.8,
adaptive_eligible: "classified_tier" as const,
};
const expectedClassifiedTierConfig = {
tiers,
classifier_type: "heuristic",
custom_technical_keywords: ["kafka", "terraform"],
keyword_tier_rules: [{ keywords: ["invoice", "refund"], tier: "MEDIUM" }],
semantic_keyword_matching: true,
embedding_model: "voyage-4-large",
match_threshold: 0.65,
adaptive: true,
adaptive_weights: { quality: 0.4, cost: 0.6 },
adaptive_eligible: "classified_tier",
};
const adaptiveDisabledValue = {
tiers,
classifier_type: "heuristic" as const,
adaptive: false,
};
const expectedAdaptiveDisabledConfig = {
tiers,
classifier_type: "heuristic",
custom_technical_keywords: ["kafka", "terraform"],
keyword_tier_rules: [{ keywords: ["invoice", "refund"], tier: "MEDIUM" }],
semantic_keyword_matching: true,
embedding_model: "voyage-4-large",
match_threshold: 0.65,
};
describe("buildUpdatedComplexityRouterConfig", () => {
it("preserves unrelated options and omits the penalty for classified-tier routing", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, classifiedTierValue);
expect(updatedConfig).toEqual(expectedClassifiedTierConfig);
});
it("removes managed adaptive and classifier fields when they are disabled", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, adaptiveDisabledValue);
expect(updatedConfig).toEqual(expectedAdaptiveDisabledConfig);
});
it("updates custom technical keywords when they are edited", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, classifiedTierValue, ["postgres"]);
expect(updatedConfig.custom_technical_keywords).toEqual(["postgres"]);
});
it("removes custom technical keywords when they are cleared", () => {
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, classifiedTierValue, []);
expect(updatedConfig.custom_technical_keywords).toBeUndefined();
});
it("preserves a tier configured with more than one model as a pool", () => {
const multiModelValue = {
...classifiedTierValue,
tiers: { ...tiers, SIMPLE: ["gpt-4o-mini", "claude-haiku-4-5"] },
};
const updatedConfig = buildUpdatedComplexityRouterConfig(storedConfig, multiModelValue);
expect(updatedConfig.tiers).toMatchObject({ SIMPLE: ["gpt-4o-mini", "claude-haiku-4-5"] });
});
});

View file

@ -4,13 +4,23 @@ import { Text, TextInput } from "@tremor/react";
import { modelAvailableCall, modelPatchUpdateCall } from "../networking";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import RouterConfigBuilder from "../add_model/RouterConfigBuilder";
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "../add_model/ComplexityRouterConfig";
import ComplexityRouterConfig, {
ComplexityRouterConfigValue,
DEFAULT_ADAPTIVE_WEIGHTS,
DEFAULT_TIER_DISTANCE_PENALTY,
} from "../add_model/ComplexityRouterConfig";
import NotificationsManager from "../molecules/notifications_manager";
const isComplexityRouterModel = (modelData: any): boolean =>
modelData?.litellm_params?.model?.startsWith("auto_router/complexity_router") ||
modelData?.litellm_params?.complexity_router_config != null;
const normalizeTierModels = (value: unknown): string[] => {
if (Array.isArray(value)) return value;
if (typeof value === "string" && value) return [value];
return [];
};
interface EditAutoRouterModalProps {
isVisible: boolean;
onCancel: () => void;
@ -20,6 +30,57 @@ interface EditAutoRouterModalProps {
userRole: string;
}
const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([
"tiers",
"classifier_type",
"classifier_llm_config",
"adaptive",
"adaptive_weights",
"tier_distance_penalty",
"adaptive_eligible",
]);
const toRecord = (value: unknown): Record<string, unknown> => {
const parsed: unknown = typeof value === "string" ? JSON.parse(value) : value;
return typeof parsed === "object" && parsed !== null && !Array.isArray(parsed)
? (parsed as Record<string, unknown>)
: {};
};
export const buildUpdatedComplexityRouterConfig = (
storedConfig: unknown,
value: ComplexityRouterConfigValue,
customTechnicalKeywords?: string[],
): Record<string, unknown> => {
const preservedConfig = Object.fromEntries(
Object.entries(toRecord(storedConfig)).filter(
([key]) =>
!MANAGED_COMPLEXITY_ROUTER_KEYS.has(key) &&
(customTechnicalKeywords === undefined || key !== "custom_technical_keywords"),
),
);
const adaptiveEligible = value.adaptive_eligible ?? "all";
return {
...preservedConfig,
tiers: value.tiers,
classifier_type: value.classifier_type,
...(value.classifier_type === "llm" ? { classifier_llm_config: value.classifier_llm_config } : {}),
...(customTechnicalKeywords &&
customTechnicalKeywords.length > 0 && {
custom_technical_keywords: customTechnicalKeywords,
}),
...(value.adaptive && {
adaptive: true,
adaptive_weights: value.adaptive_weights ?? DEFAULT_ADAPTIVE_WEIGHTS,
...(adaptiveEligible === "all" && {
tier_distance_penalty: value.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY,
}),
adaptive_eligible: adaptiveEligible,
}),
};
};
const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
isVisible,
onCancel,
@ -35,8 +96,9 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
const [showCustomDefaultModel, setShowCustomDefaultModel] = useState<boolean>(false);
const [showCustomEmbeddingModel, setShowCustomEmbeddingModel] = useState<boolean>(false);
const [routerConfig, setRouterConfig] = useState<any>(null);
const [customTechnicalKeywords, setCustomTechnicalKeywords] = useState<string[]>([]);
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] },
classifier_type: "heuristic",
});
const isComplexityRouter = isComplexityRouterModel(modelData);
@ -85,14 +147,21 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
setComplexityRouterConfig({
tiers: {
SIMPLE: parsedConfig.tiers?.SIMPLE || "",
MEDIUM: parsedConfig.tiers?.MEDIUM || "",
COMPLEX: parsedConfig.tiers?.COMPLEX || "",
REASONING: parsedConfig.tiers?.REASONING || "",
SIMPLE: normalizeTierModels(parsedConfig.tiers?.SIMPLE),
MEDIUM: normalizeTierModels(parsedConfig.tiers?.MEDIUM),
COMPLEX: normalizeTierModels(parsedConfig.tiers?.COMPLEX),
REASONING: normalizeTierModels(parsedConfig.tiers?.REASONING),
},
classifier_type: parsedConfig.classifier_type || "heuristic",
classifier_llm_config: parsedConfig.classifier_llm_config,
adaptive: parsedConfig.adaptive || false,
adaptive_weights: parsedConfig.adaptive_weights,
tier_distance_penalty: parsedConfig.tier_distance_penalty,
adaptive_eligible: parsedConfig.adaptive_eligible || "all",
});
setCustomTechnicalKeywords(
Array.isArray(parsedConfig.custom_technical_keywords) ? parsedConfig.custom_technical_keywords : [],
);
form.setFieldsValue({
auto_router_name: modelData.model_name,
@ -138,7 +207,7 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
if (isComplexityRouter) {
const { tiers, classifier_type, classifier_llm_config } = complexityRouterConfig;
if (Object.values(tiers).filter(Boolean).length === 0) {
if (Object.values(tiers).every((models) => models.length === 0)) {
NotificationsManager.fromBackend("Please select at least one model for a complexity tier");
return;
}
@ -147,14 +216,14 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
return;
}
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
const defaultModel = tiers.MEDIUM[0] || tiers.SIMPLE[0] || tiers.COMPLEX[0] || tiers.REASONING[0];
const updatedLitellmParams = {
...modelData.litellm_params,
complexity_router_config: {
tiers,
classifier_type,
...(classifier_type === "llm" ? { classifier_llm_config } : {}),
},
complexity_router_config: buildUpdatedComplexityRouterConfig(
modelData.litellm_params?.complexity_router_config,
complexityRouterConfig,
customTechnicalKeywords,
),
complexity_router_default_model: defaultModel,
};
const updatedModelInfo = {
@ -264,6 +333,8 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
onChange={(config) => {
setComplexityRouterConfig(config);
}}
customTechnicalKeywords={customTechnicalKeywords}
onCustomTechnicalKeywordsChange={setCustomTechnicalKeywords}
/>
</div>
) : (

View file

@ -22619,6 +22619,11 @@ export interface components {
* @description When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.
*/
reject_clientside_metadata_tags?: boolean | null;
/**
* Skip User Budget On Team Key
* @description If True, restores the legacy behavior where a user's personal max_budget is NOT enforced when their key belongs to a team; only the team (and team-member) budgets apply. Defaults to False, meaning the user's personal max_budget is always enforced regardless of whether the key belongs to a team (see GitHub issue #12905).
*/
skip_user_budget_on_team_key?: boolean | null;
/**
* Store Model In Db
* @description If True, models and config are stored in and loaded from the database. Default is False.

22
uv.lock generated
View file

@ -9,7 +9,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-07-10T16:47:58.286372Z"
exclude-newer = "2026-07-11T16:57:33.067258Z"
exclude-newer-span = "P3D"
[manifest]
@ -20,7 +20,9 @@ members = [
]
constraints = [
{ name = "aiohttp", specifier = ">=3.14.1,<4.0" },
{ name = "httplib2", specifier = ">=0.32.0" },
{ name = "packaging", specifier = ">=24.0" },
{ name = "setuptools", specifier = ">=83.0.0" },
{ name = "soupsieve", specifier = ">=2.8.4" },
{ name = "tornado", specifier = ">=6.5.6" },
]
@ -2504,14 +2506,14 @@ wheels = [
[[package]]
name = "httplib2"
version = "0.31.2"
version = "0.32.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "pyparsing" },
]
sdist = { url = "https://files.pythonhosted.org/packages/c1/1f/e86365613582c027dda5ddb64e1010e57a3d53e99ab8a72093fa13d565ec/httplib2-0.31.2.tar.gz", hash = "sha256:385e0869d7397484f4eab426197a4c020b606edd43372492337c0b4010ae5d24", size = 250800, upload-time = "2026-01-23T11:04:44.165Z" }
sdist = { url = "https://files.pythonhosted.org/packages/84/f5/ccf58de92d61e3ad921119668f54ed36ca1d0cf5dcc5c1657dfb164fd78b/httplib2-0.32.0.tar.gz", hash = "sha256:48a0ef30a42db65d8f3399045e1d09ab0ba66e3b9efc360d07f80ea55d286025", size = 254283, upload-time = "2026-06-26T10:13:56.265Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/2f/90/fd509079dfcab01102c0fdd87f3a9506894bc70afcf9e9785ef6b2b3aff6/httplib2-0.31.2-py3-none-any.whl", hash = "sha256:dbf0c2fa3862acf3c55c078ea9c0bc4481d7dc5117cae71be9514912cf9f8349", size = 91099, upload-time = "2026-01-23T11:04:42.78Z" },
{ url = "https://files.pythonhosted.org/packages/33/a0/550eec327e5f5c7b732531c489f5307efec41f047b0d703bd4ca1e5ad2db/httplib2-0.32.0-py3-none-any.whl", hash = "sha256:dc6705cacdf3fb0a2aba7629fa33c90fd93e30035db0c157325826be177e4816", size = 93148, upload-time = "2026-06-26T10:13:54.985Z" },
]
[[package]]
@ -3283,7 +3285,7 @@ wheels = [
[[package]]
name = "litellm"
version = "1.93.0"
version = "1.94.0"
source = { editable = "." }
dependencies = [
{ name = "aiohttp" },
@ -3656,12 +3658,12 @@ proxy-dev = [
[[package]]
name = "litellm-enterprise"
version = "0.1.49"
version = "0.1.50"
source = { editable = "enterprise" }
[[package]]
name = "litellm-proxy-extras"
version = "0.4.76"
version = "0.4.77"
source = { editable = "litellm-proxy-extras" }
[[package]]
@ -7003,11 +7005,11 @@ wheels = [
[[package]]
name = "setuptools"
version = "82.0.1"
version = "83.0.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/4f/db/cfac1baf10650ab4d1c111714410d2fbb77ac5a616db26775db562c8fab2/setuptools-82.0.1.tar.gz", hash = "sha256:7d872682c5d01cfde07da7bccc7b65469d3dca203318515ada1de5eda35efbf9", size = 1152316, upload-time = "2026-03-09T12:47:17.221Z" }
sdist = { url = "https://files.pythonhosted.org/packages/34/26/f5d29e25ffdb535afef2d35cdb55b325298f96debd670da4c325e08d70f4/setuptools-83.0.0.tar.gz", hash = "sha256:025bccbbf0fa05b6192bc64ae1e7b16e001fd6d6d4d5de03c97b1c1ade523bef", size = 1154254, upload-time = "2026-07-04T15:31:22.699Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/9d/76/f789f7a86709c6b087c5a2f52f911838cad707cc613162401badc665acfe/setuptools-82.0.1-py3-none-any.whl", hash = "sha256:a59e362652f08dcd477c78bb6e7bd9d80a7995bc73ce773050228a348ce2e5bb", size = 1006223, upload-time = "2026-03-09T12:47:15.026Z" },
{ url = "https://files.pythonhosted.org/packages/5d/40/e1e72872c6354b306daef1703549e8e83b4d43cfea356311bf722a043752/setuptools-83.0.0-py3-none-any.whl", hash = "sha256:29b23c360f22f414dc7336bb39178cc7bcbf6021ed2733cde173f09dba19abb3", size = 1008090, upload-time = "2026-07-04T15:31:20.885Z" },
]
[[package]]