Merge origin/litellm_internal_staging into litellm_lit_4725_cache_write_cost_tracking

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
milan 2026-07-24 02:32:32 +00:00
commit 961e52964b
159 changed files with 14351 additions and 8881 deletions

View file

@ -5,10 +5,24 @@ on:
branches:
- main
- litellm_internal_staging
paths:
- "litellm/**"
- "tests/benchmarks/**"
- "pyproject.toml"
- "uv.lock"
- ".github/workflows/codspeed.yml"
- ".github/actions/setup-uv-with-retries/**"
pull_request:
branches:
- main
- litellm_internal_staging
paths:
- "litellm/**"
- "tests/benchmarks/**"
- "pyproject.toml"
- "uv.lock"
- ".github/workflows/codspeed.yml"
- ".github/actions/setup-uv-with-retries/**"
# Allow CodSpeed to trigger backtest performance analysis
# in order to generate initial data
workflow_dispatch:

View file

@ -211,6 +211,9 @@ filter_invalid_headers: Optional[bool] = False
add_user_information_to_llm_headers: Optional[bool] = (
None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
)
overwrite_user_with_key_hash: bool = (
False # force the outgoing `user` param to the hashed api key, so providers see a stable, tamper-proof id
)
store_audit_logs = False # Enterprise feature, allow users to see audit logs
skip_system_message_in_guardrail: bool = False
skip_tool_message_in_guardrail: bool = False

View file

@ -2191,6 +2191,13 @@ def batch_cost_calculator(
return total_prompt_cost, total_completion_cost
def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> List[str]:
field_names = list(type(prompt_tokens_details).model_fields)
if getattr(prompt_tokens_details, "cache_write_tokens", None) is None:
return field_names
return [attr for attr in field_names if attr != "cache_creation_tokens"]
class BaseTokenUsageProcessor:
@staticmethod
def combine_usage_objects(usage_objects: List[Usage]) -> Usage:
@ -2225,7 +2232,7 @@ class BaseTokenUsageProcessor:
# Check what keys exist in the model's prompt_tokens_details
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
for attr in type(usage.prompt_tokens_details).model_fields:
for attr in _summable_prompt_token_fields(usage.prompt_tokens_details):
if (
hasattr(usage.prompt_tokens_details, attr)
and not attr.startswith("_")

View file

@ -470,8 +470,8 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
cache_creation_tokens = (
cast(
Optional[int],
getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0)
or getattr(usage.prompt_tokens_details, "cache_write_tokens", 0),
getattr(usage.prompt_tokens_details, "cache_write_tokens", 0)
or getattr(usage.prompt_tokens_details, "cache_creation_tokens", 0),
)
or 0
)
@ -920,10 +920,6 @@ def get_token_type_cost_breakdown(
cache_read_tokens = prompt_tokens_details["cache_hit_tokens"]
cache_creation_tokens = prompt_tokens_details["cache_creation_tokens"]
cache_creation_token_details = prompt_tokens_details["cache_creation_token_details"]
# Some OpenAI-compatible providers (e.g. kimi-k2) report cache-write tokens
# under `cache_write_tokens`; mirror the total-cost normalization path.
if not cache_creation_tokens:
cache_creation_tokens = _coerce_token_count(getattr(usage.prompt_tokens_details, "cache_write_tokens", 0))
# Fall back to the private top-level counters the Usage constructor mirrors cache
# tokens onto, so providers/callers that bypass prompt_tokens_details are covered.
if not cache_read_tokens:

View file

@ -1,192 +0,0 @@
"""
OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers.
Exchanges a user's incoming JWT (subject_token) for a scoped access token
at an IDP's token exchange endpoint. The exchanged token is then used to
authenticate requests to the upstream MCP server.
See: https://datatracker.ietf.org/doc/html/rfc8693
"""
import asyncio
import hashlib
import weakref
from typing import TYPE_CHECKING, Dict, Tuple
import httpx
from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
build_token_endpoint_client_auth,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
# RFC 8693 grant type constant
TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange"
class TokenExchangeHandler:
"""Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers.
Caches exchanged tokens keyed by ``hash(subject_token + server_id)`` so
repeated calls with the same user token skip the IDP round-trip.
"""
def __init__(self) -> None:
self._cache = InMemoryCache(
max_size_in_memory=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
)
# WeakValueDictionary so locks are GC'd once no coroutine holds a reference,
# preventing unbounded growth with many rotating user tokens.
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = weakref.WeakValueDictionary()
def _get_lock(self, cache_key: str) -> asyncio.Lock:
lock = self._locks.get(cache_key)
if lock is None:
lock = asyncio.Lock()
self._locks[cache_key] = lock
return lock
@staticmethod
def _cache_key(subject_token: str, server_id: str) -> str:
raw = f"{subject_token}:{server_id}"
return hashlib.sha256(raw.encode()).hexdigest()
async def exchange_token(
self,
subject_token: str,
server: "MCPServer",
) -> str:
"""Exchange *subject_token* for a scoped access token.
Returns the exchanged ``access_token`` string (suitable for a
``Bearer`` header).
Raises ``ValueError`` on configuration or IDP errors.
"""
cache_key = self._cache_key(subject_token, server.server_id)
# Fast path
cached = self._cache.get_cache(cache_key)
if cached is not None:
return cached
# Slow path — one exchange at a time per (user, server) pair
async with self._get_lock(cache_key):
cached = self._cache.get_cache(cache_key)
if cached is not None:
return cached
token, ttl = await self._do_exchange(subject_token, server)
self._cache.set_cache(cache_key, token, ttl=ttl)
return token
async def _do_exchange(
self,
subject_token: str,
server: "MCPServer",
) -> Tuple[str, int]:
"""POST to the token exchange endpoint with RFC 8693 parameters.
Returns ``(access_token, ttl_seconds)``.
"""
endpoint = server.token_exchange_endpoint or server.token_url
if not endpoint:
raise ValueError(
f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange "
f"but no token_exchange_endpoint or token_url configured"
)
if not server.client_id or not server.client_secret:
raise ValueError(
f"MCP server '{server.server_id}' has auth_type=oauth2_token_exchange "
f"but missing client_id or client_secret"
)
client_auth = build_token_endpoint_client_auth(
auth_method=server.token_endpoint_auth_method,
client_id=server.client_id,
client_secret=server.client_secret,
)
data: Dict[str, str] = {
"grant_type": TOKEN_EXCHANGE_GRANT_TYPE,
"subject_token": subject_token,
"subject_token_type": server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
**client_auth.body,
}
if server.audience:
data["audience"] = server.audience
if server.scopes:
data["scope"] = " ".join(server.scopes)
verbose_logger.debug(
"Exchanging token for MCP server %s at %s (audience=%s)",
server.server_id,
endpoint,
server.audience,
)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})}
try:
response = await client.post(endpoint, **post_kwargs)
response.raise_for_status()
except httpx.HTTPStatusError as exc:
verbose_logger.debug(
"Token exchange IDP error for MCP server %s (status %d)",
server.server_id,
exc.response.status_code,
)
raise ValueError(
f"Token exchange for MCP server '{server.server_id}' failed with status {exc.response.status_code}"
) from exc
body = response.json()
if not isinstance(body, dict):
raise ValueError(
f"Token exchange response for MCP server '{server.server_id}' "
f"returned non-object JSON (got {type(body).__name__})"
)
access_token = body.get("access_token")
if not access_token:
raise ValueError(f"Token exchange response for MCP server '{server.server_id}' missing 'access_token'")
raw_expires_in = body.get("expires_in")
try:
expires_in = int(raw_expires_in) if raw_expires_in is not None else MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
except (TypeError, ValueError):
expires_in = MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL
ttl = max(
expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS,
MCP_OAUTH2_TOKEN_CACHE_MIN_TTL,
)
verbose_logger.info(
"Token exchange succeeded for MCP server %s (expires in %ds)",
server.server_id,
expires_in,
)
return access_token, ttl
def invalidate(self, subject_token: str, server_id: str) -> None:
"""Remove a cached exchanged token (e.g. after a 401)."""
cache_key = self._cache_key(subject_token, server_id)
self._cache.delete_cache(cache_key)
# Module-level singleton
mcp_token_exchange_handler = TokenExchangeHandler()

View file

@ -33,6 +33,7 @@ from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
_finish_bridge_mint,
_prepare_bridge_mint,
_prepare_bridge_refresh,
_reload_active_user_by_id,
)
from litellm.proxy._experimental.mcp_server.faults import (
CallerRejected,
@ -43,6 +44,14 @@ from litellm.proxy._experimental.mcp_server.faults import (
dcr_fault_detail,
render_token_fault,
)
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
aggregate_authorize,
aggregate_token,
complete_connect_flow,
is_gateway_dcr_client_id,
register_aggregate_client,
relative_request_url,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
get_request_base_url,
@ -324,14 +333,25 @@ def redeem_passthrough_authorization_code(
return sealed
def _session_cookie_user_id(request: Request) -> str | None:
"""The signed-in litellm user for a browser request, or ``None``. Thin wrapper so the
aggregate DCR flow's verbs receive the identity as a plain value instead of parsing
cookies themselves."""
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import ( # noqa: PLC0415 # circular import at module load
_user_id_from_session_cookie,
)
return _user_id_from_session_cookie(request)
def _redirect_to_litellm_login(request: Request) -> RedirectResponse:
"""Send an unauthenticated browser through litellm login before the interactive bridge authorize
can capture its identity. The bridge oauth_delegate flow seals the SSO user into the gateway code,
so a session is required; without one there is nothing to bind. After login the user re-initiates
the connection, which then finds the session cookie (the seamless return-to round-trip, which is
origin-validated against the control-plane URL, is a follow-up)."""
so a session is required; without one there is nothing to bind. A same-origin relative
``return_to`` (honored by the SSO callback) brings the browser straight back to this authorize
request after login instead of stranding it on the dashboard."""
base_url = get_request_base_url(request)
return RedirectResponse(f"{base_url}/sso/key/generate")
return RedirectResponse(f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}")
# LIT-4197: some upstream authorization servers reject an over-long ``state``
@ -606,6 +626,35 @@ def _raise_if_not_oauth2(mcp_server: MCPServer) -> None:
)
def _endpoint_not_configured_detail(
mcp_server: MCPServer,
endpoint_label: str,
manual_remedy: str,
issuer_remedy: str,
) -> str:
"""The 400 detail for an unresolved OAuth endpoint, naming the likely cause for this server's
shape (LIT-4658): an anchored issuer whose metadata fell short, a configured (possibly
misconfigured) server url whose discovery failed, or no discovery source at all. Kept free of
URLs and issuer values because these endpoints are reachable pre-auth."""
if mcp_server.issuer_is_anchored:
return (
f"MCP server {endpoint_label} is not configured. Endpoint discovery anchored on the configured "
f"Issuer (RFC 8414) failed or its metadata did not include this endpoint; check the proxy logs "
f"for 'MCP OAuth' warnings from server load, verify the Issuer, or {manual_remedy}."
)
if mcp_server.url:
return (
f"MCP server {endpoint_label} is not configured. OAuth endpoint discovery against the configured "
f"server url did not resolve it; the url may be misconfigured. Check the proxy logs for "
f"'MCP OAuth' warnings from server load, verify the server url, or {manual_remedy}, or "
f"{issuer_remedy}."
)
return (
f"MCP server {endpoint_label} is not configured. Servers with no url (OpenAPI spec or stdio) run no "
f"resource discovery, so {manual_remedy}, or {issuer_remedy}."
)
def _raise_unless_oauth2_discovery_server(
mcp_server: Optional[MCPServer],
mcp_server_name: Optional[str],
@ -707,10 +756,11 @@ async def authorize_with_server(
if mcp_server.authorization_url is None:
raise HTTPException(
status_code=400,
detail=(
"MCP server authorization url is not configured. Servers with no url (OpenAPI "
"spec or stdio) run no resource discovery, so set Authorization URL and Token URL "
"manually, or set Issuer to discover them from the identity provider (RFC 8414)."
detail=_endpoint_not_configured_detail(
mcp_server,
"authorization url",
"set Authorization URL and Token URL manually",
"set Issuer to discover them from the identity provider (RFC 8414)",
),
)
@ -828,10 +878,11 @@ async def exchange_token_with_server(
if mcp_server.token_url is None:
raise HTTPException(
status_code=400,
detail=(
"MCP server token url is not configured. Servers with no url (OpenAPI spec or "
"stdio) run no resource discovery, so set Token URL manually, or set Issuer to "
"discover it from the identity provider (RFC 8414)."
detail=_endpoint_not_configured_detail(
mcp_server,
"token url",
"set Token URL manually",
"set Issuer to discover it from the identity provider (RFC 8414)",
),
)
@ -1540,10 +1591,11 @@ async def register_client_with_server(
if mcp_server.authorization_url is None:
raise HTTPException(
status_code=400,
detail=(
"MCP server authorization url is not configured. Servers with no url (OpenAPI "
"spec or stdio) run no resource discovery, so set Authorization URL and Token URL "
"manually, or set Issuer to discover them from the identity provider (RFC 8414)."
detail=_endpoint_not_configured_detail(
mcp_server,
"authorization url",
"set Authorization URL and Token URL manually",
"set Issuer to discover them from the identity provider (RFC 8414)",
),
)
@ -1601,6 +1653,18 @@ async def authorize(
global_mcp_server_manager,
)
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id):
return aggregate_authorize(
request=request,
client_id=client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
code_challenge_method=code_challenge_method,
response_type=response_type,
session_user_id=_session_cookie_user_id(request),
)
lookup_name: Optional[str] = mcp_server_name or client_id
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = (
@ -1664,6 +1728,25 @@ async def token_endpoint(
global_mcp_server_manager,
)
if mcp_server_name is None and is_gateway_dcr_client_id(client_id):
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
master_key,
user_api_key_cache,
)
return await aggregate_token(
request=request,
grant_type=grant_type,
code=code,
redirect_uri=redirect_uri,
client_id=client_id,
code_verifier=code_verifier,
refresh_token=refresh_token,
master_key=master_key,
reload_user=_reload_active_user_by_id,
cache=user_api_key_cache,
)
lookup_name = mcp_server_name or client_id
client_ip = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip)
@ -1685,6 +1768,21 @@ async def token_endpoint(
)
@router.post("/authorize/complete")
async def authorize_complete(request: Request, flow: str = Form(...)):
"""Finish an aggregate connect flow: mint the gateway authorization code for the
signed-in user and redirect back to the DCR client. POST plus the per-flow HttpOnly
cookie set at /authorize; an anonymous or bad-flow request just 400s."""
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load
return await complete_connect_flow(
request=request,
flow_handle=flow,
session_user_id=_session_cookie_user_id(request),
cache=user_api_key_cache,
)
# Per RFC 6749 §4.1.2.1, an IdP that rejects an OAuth authorization request
# redirects back to the configured redirect URI with ``error`` /
# ``error_description`` / ``error_uri`` query params and no ``code``. The MCP
@ -2422,6 +2520,13 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non
}
client_ip = IPAddressUtils.get_mcp_client_ip(request)
if not mcp_server_name:
# A real DCR request carries redirect_uris (RFC 7591): route it to the aggregate DCR
# endpoint the aggregate authorization-server metadata advertises. A single-server
# deployment registers at /{server}/register instead (its bare-origin discovery
# advertises that), so this does not affect it. A request without redirect_uris is not
# a DCR request, so the legacy single-server-or-dummy fallback is kept for it.
if data.get("redirect_uris"):
return await register_aggregate_client(request=request, request_body=data)
resolved = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if resolved:
return await register_client_with_server(

View file

@ -0,0 +1,637 @@
"""The gateway-level DCR flow for the aggregate ``/mcp`` endpoint (``mcp_gateway_dcr``).
An OAuth-only DCR client (Claude Desktop, Claude Code, MCP Inspector) pointed at the
aggregate ``/mcp`` endpoint discovers the gateway as its authorization server (PR 1 of
this track) and then walks the flow implemented here:
1. ``POST /register``: stateless dynamic client registration. The ``client_id`` IS the
registration: the client's redirect URIs are sealed into it with the repo's
authenticated symmetric helper, so nothing is persisted and a forged or tampered
client_id simply fails to open. Clients are always public (``token_endpoint_auth_method
"none"``); PKCE S256 is what protects the code.
2. ``GET /authorize``: validates the client and redirect URI, requires S256 PKCE, and
interposes LiteLLM sign-in. Without a session cookie the browser is sent through
``/sso/key/generate`` with a same-origin ``return_to`` so it lands back here after
login. With a session, the flow parameters and the SSO user are sealed into a per-flow
HttpOnly cookie (the same pattern as the upstream OAuth state relay) and the browser is
sent to the connect page, where the user authorizes individual servers (vaulting those
tokens server-side) before finishing.
3. ``POST /authorize/complete``: the deliberate finish step. A POST (not GET) bound to the
SameSite=Lax flow cookie, so a cross-site link cannot silently mint a code with the
victim's session, and the signed-in user must match the user sealed into the flow.
Mints a short-lived, single-use, gateway-sealed authorization code and redirects to the
client's registered redirect URI.
4. ``POST /token``: exchanges the code (PKCE-verified, client- and redirect-bound,
single-use) for the identity-only session tokens of
:mod:`.outbound_credentials.session_token`, re-validating that the litellm user is
still active first; the ``refresh_token`` grant rotates the pair the same way.
Nothing here stores state server-side except the single-use code guard (a TTL cache
entry). Every sealed value is authenticated encryption over the proxy salt/master key
family, opened totally (bad input maps to an OAuth error, never a raise), and every
identity is a stable reference re-validated live at mint, refresh, and (in the admission
PR) tool-call time. Upstream server credentials never appear anywhere in this flow; they
are vaulted per user by the existing ``/v1/mcp`` authorize endpoints and resolved at
egress by user id.
"""
from __future__ import annotations
import hashlib
import hmac
import secrets
from base64 import urlsafe_b64encode
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import Awaitable, Callable, Literal, TypeVar
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, RedirectResponse, Response
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
get_request_base_url,
is_loopback_redirect_host,
validate_redirect_uri_shape,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
SessionRefreshOpened,
open_session_refresh_bearer,
session_keys_from_master_key,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
SESSION_REFRESH_TTL_SECONDS,
MintedSessionToken,
SessionKeys,
SessionPrincipal,
mint_session_refresh_token,
mint_session_token,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
GATEWAY_DCR_CLIENT_ID_PREFIX = "llm_dcrc_"
"""Marker prefix on every gateway-issued DCR client_id so the root authorize/token
endpoints can route an aggregate-flow request without decrypting, and existing per-server
flows (whose client_ids are upstream-issued) are never captured by the aggregate arm."""
GATEWAY_AUTH_CODE_PREFIX = "llm_gcode_"
"""Marker prefix on the gateway-sealed authorization code, distinct from the bridge
``llm_bcode_`` so neither flow can consume the other's codes."""
CONNECT_FLOW_COOKIE_PREFIX = "mcp_connect_flow_"
"""Per-flow HttpOnly cookie holding the sealed connect flow, keyed by a short random
handle carried in the connect-page URL (the same handle-plus-cookie pattern as the
``mcp_oauth_state_`` upstream relay, for the same reasons: replica-safe with no
server-side session store, and the sealed value never appears in a URL)."""
CONNECT_FLOW_TTL_SECONDS = 600
GATEWAY_AUTH_CODE_TTL_SECONDS = 120
_CLAIM_TTL_BUFFER_SECONDS = 60
_USED_CODE_CACHE_PREFIX = "mcp_gateway_dcr_code_used:"
_USED_FLOW_CACHE_PREFIX = "mcp_gateway_dcr_flow_used:"
_USED_REFRESH_CACHE_PREFIX = "mcp_gateway_dcr_refresh_used:"
MAX_REDIRECT_URIS = 3
MAX_REDIRECT_URI_LENGTH = 256
MAX_CLIENT_ID_LENGTH = 2048
"""Registration bounds. They exist to bound the sealed client_id, which rides inside
every session-token claim set: 3 URIs of 256 bytes seal to roughly 1.2KB, comfortably
under this cap and under the session token's own 4KB ceiling. Claude Desktop and MCP
Inspector register one or two redirect URIs."""
MAX_STATE_LENGTH = 1024
"""Bound on the client ``state`` sealed into the flow cookie and echoed on the auth-code
redirect. An unbounded ``state`` can push the sealed cookie past the browser's ~4KB cap
(silently dropped, breaking the flow); spec clients send a short opaque value."""
MIN_CODE_VERIFIER_LENGTH = 43
MAX_CODE_VERIFIER_LENGTH = 128
"""RFC 7636 section 4.1 bounds for the PKCE ``code_verifier``. Enforced so an out-of-range
verifier gets a clean ``invalid_request`` instead of an opaque PKCE-mismatch."""
_UNPREFIXED = ""
"""Prefix for a sealed value that carries no wire marker because it is never routed by
prefix (the connect flow lives only in its own per-handle cookie, opened by that one
handle). Named so the empty-string argument to ``_seal`` / ``_open_sealed`` reads as
deliberate rather than a typo."""
_CLIENT_RECORD_DEBUG_KEY = "gateway_dcr_client"
_CONNECT_FLOW_DEBUG_KEY = "gateway_connect_flow"
_AUTH_CODE_DEBUG_KEY = "gateway_authorization_code"
ReloadUserFailure = Literal["unresolvable", "unavailable", "no_active_key"]
ReloadUser = Callable[[str], Awaitable[ReloadUserFailure | None]]
"""Injected live-user revalidation (the token endpoint's mirror of admission):
``None`` means the user is active; ``unavailable`` is a retryable DB outage; anything
else fails the grant closed."""
class GatewayDcrClient(BaseModel):
"""The registration record sealed into a gateway DCR ``client_id``.
``extra="forbid"`` so a sealed value of another type (an auth code, a connect flow)
that happened to decrypt under the shared key can never validate as a client record:
cross-type confusion is rejected at the model boundary, not left to differing required
fields."""
model_config = ConfigDict(frozen=True, extra="forbid")
redirect_uris: tuple[str, ...] = Field(min_length=1, max_length=MAX_REDIRECT_URIS)
iat: int
class _ConnectFlow(BaseModel):
"""One in-flight authorize: the SSO user it belongs to and the client parameters
needed to mint the code at the finish step. Sealed into the per-flow cookie. ``jti``
makes the flow single-use at complete; ``extra="forbid"`` rejects cross-type
confusion."""
model_config = ConfigDict(frozen=True, extra="forbid")
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
redirect_uri: str = Field(min_length=1)
state: str
code_challenge: str = Field(min_length=1)
jti: str = Field(min_length=1)
exp: int
class _GatewayAuthCode(BaseModel):
"""The gateway-sealed authorization code: the user consent it represents and the
bindings the token endpoint must verify (client, redirect URI, PKCE challenge),
plus a ``jti`` for the single-use guard. ``extra="forbid"`` rejects cross-type
confusion."""
model_config = ConfigDict(frozen=True, extra="forbid")
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
redirect_uri: str = Field(min_length=1)
code_challenge: str = Field(min_length=1)
jti: str = Field(min_length=1)
iat: int
exp: int
def is_gateway_dcr_client_id(client_id: str | None) -> bool:
"""Cheap prefix routing test so the root endpoints only enter the aggregate arm for
clients this flow registered; every other client_id keeps today's behavior."""
return client_id is not None and client_id.startswith(GATEWAY_DCR_CLIENT_ID_PREFIX)
def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse:
"""RFC 6749 section 5.2 / RFC 7591 section 3.2.2 error body. Descriptions carry no
token, code, or URL material so they are safe to relay to any client."""
return JSONResponse(
status_code=status_code,
content={"error": error, "error_description": description},
headers=TOKEN_NO_CACHE_HEADERS,
)
def _seal(prefix: str, payload: BaseModel) -> str:
return prefix + encrypt_value_helper(payload.model_dump_json())
_SealedModelT = TypeVar("_SealedModelT", bound=BaseModel)
def _open_sealed(value: str, prefix: str, model: type[_SealedModelT], debug_key: str) -> _SealedModelT | None:
"""Open a sealed value totally: anything that is not prefix-shaped, does not decrypt,
or does not validate returns ``None`` for the caller to map onto an OAuth error."""
if not value.startswith(prefix):
return None
decrypted = decrypt_value_helper(value[len(prefix) :], debug_key, return_original_value=False)
if not isinstance(decrypted, str):
return None
try:
return model.model_validate_json(decrypted)
except ValidationError:
return None
def open_gateway_dcr_client(client_id: str) -> GatewayDcrClient | None:
return _open_sealed(client_id, GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient, _CLIENT_RECORD_DEBUG_KEY)
async def register_aggregate_client(request: Request, request_body: Mapping[str, object]) -> Response:
"""RFC 7591 dynamic registration against the gateway itself, statelessly.
Only ``redirect_uris`` is authoritative; every client is registered as a public
``token_endpoint_auth_method "none"`` client regardless of what it asked for (RFC
7591 lets the server override metadata), because the gateway never issues client
secrets: possession of a secret would add nothing over the mandatory S256 PKCE, and a
stateless registration has nowhere to keep one. Nothing is persisted, so open
registration cannot be used to fill storage.
Redirect-URI *hygiene* is not decided here: :func:`validate_redirect_uri_shape` is
the single owner of that rule across the MCP OAuth surface, so allowlisted native
callbacks (``cursor://``) are accepted and fragments, missing hosts, userinfo
(``https://claude.ai@attacker.example/cb``) and backslash hosts are rejected exactly
as they are on /authorize and /callback.
What this endpoint does decide is its own trust policy, which is deliberately wider
than :func:`validate_trusted_redirect_uri`'s: registration is *public*, so any https
client may register (that is what lets a hosted MCP client register at all), and the
controls are mandatory S256 PKCE plus the consent screen showing the client origin.
http is confined to loopback per RFC 8252 section 7.3.
"""
raw_uris = request_body.get("redirect_uris")
if not isinstance(raw_uris, list) or not raw_uris or len(raw_uris) > MAX_REDIRECT_URIS:
return _oauth_error(
400,
"invalid_redirect_uri",
f"redirect_uris must be a list of 1 to {MAX_REDIRECT_URIS} URIs",
)
if not all(isinstance(uri, str) and len(uri) <= MAX_REDIRECT_URI_LENGTH for uri in raw_uris):
return _oauth_error(
400,
"invalid_redirect_uri",
f"each redirect URI must be a string of at most {MAX_REDIRECT_URI_LENGTH} characters",
)
for uri in raw_uris:
parsed = urlparse(uri)
try:
if validate_redirect_uri_shape(parsed):
continue # allowlisted native callback, e.g. cursor://
except HTTPException as exc:
# The shared validator speaks HTTP; RFC 7591 registration answers with an OAuth
# error object, so translate the shape without re-deciding the rule.
return _oauth_error(400, "invalid_redirect_uri", str(exc.detail))
if parsed.scheme == "https" or (parsed.scheme == "http" and is_loopback_redirect_host(parsed)):
continue
return _oauth_error(
400,
"invalid_redirect_uri",
"each redirect URI must be https, http on a loopback host, or a registered native callback",
)
now = datetime.now(timezone.utc)
client_id = _seal(
GATEWAY_DCR_CLIENT_ID_PREFIX, GatewayDcrClient(redirect_uris=tuple(raw_uris), iat=int(now.timestamp()))
)
if len(client_id) > MAX_CLIENT_ID_LENGTH:
return _oauth_error(400, "invalid_client_metadata", "registered metadata is too large")
return JSONResponse(
status_code=201,
content={
"client_id": client_id,
"client_id_issued_at": int(now.timestamp()),
"redirect_uris": list(raw_uris),
"token_endpoint_auth_method": "none",
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
},
)
def _flow_cookie_name(handle: str) -> str:
return f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}"
def _cookie_path_and_secure(request: Request) -> tuple[str, bool]:
parsed = urlparse(get_request_base_url(request))
return parsed.path or "/", parsed.scheme == "https"
def _append_query_params(url: str, params: dict[str, str]) -> str:
parsed = urlparse(url)
query = parse_qsl(parsed.query, keep_blank_values=True) + list(params.items())
return urlunparse(parsed._replace(query=urlencode(query)))
def relative_request_url(request: Request) -> str:
"""The request's own path and query as a same-origin ``return_to`` target for the
login round-trip; relative by construction, so it can never leave the gateway."""
path = request.url.path
return f"{path}?{request.url.query}" if request.url.query else path
def aggregate_authorize(
request: Request,
client_id: str,
redirect_uri: str,
state: str,
code_challenge: str | None,
code_challenge_method: str | None,
response_type: str | None,
session_user_id: str | None,
) -> Response:
"""The aggregate authorize verb: validate the client, require S256 PKCE, interpose
LiteLLM sign-in, and hand the browser to the connect page with the flow sealed into a
per-flow cookie.
Validation failures respond directly with 400 and never redirect: per RFC 6749
section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and
once the client is at fault there is no trusted place to send the browser.
"""
client = open_gateway_dcr_client(client_id)
if client is None:
return _oauth_error(400, "invalid_client", "unknown or malformed client_id")
if redirect_uri not in client.redirect_uris:
return _oauth_error(400, "invalid_request", "redirect_uri is not registered for this client")
if response_type != "code":
return _oauth_error(400, "unsupported_response_type", "response_type must be 'code'")
if not code_challenge or code_challenge_method != "S256":
return _oauth_error(
400,
"invalid_request",
"PKCE is required: send code_challenge with code_challenge_method=S256",
)
if len(state) > MAX_STATE_LENGTH:
return _oauth_error(400, "invalid_request", f"state must be at most {MAX_STATE_LENGTH} characters")
base_url = get_request_base_url(request)
if session_user_id is None:
login_url = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}"
return RedirectResponse(login_url, status_code=303)
now = datetime.now(timezone.utc)
handle = secrets.token_urlsafe(24)
flow = _ConnectFlow(
user_id=session_user_id,
client_id=client_id,
redirect_uri=redirect_uri,
state=state,
code_challenge=code_challenge,
jti=secrets.token_urlsafe(24),
exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS,
)
connect_url = _append_query_params(
f"{base_url}/ui/chat/integrations",
{"connect_flow": handle, "connect_client": _origin_only(redirect_uri)},
)
response = RedirectResponse(connect_url, status_code=303)
path, secure = _cookie_path_and_secure(request)
response.set_cookie(
key=_flow_cookie_name(handle),
value=_seal(_UNPREFIXED, flow),
max_age=CONNECT_FLOW_TTL_SECONDS,
path=path,
secure=secure,
httponly=True,
samesite="lax",
)
return response
def _origin_only(url: str) -> str:
"""Scheme+host for display on the connect page; never the full redirect URI, whose
path or query could carry values that do not belong in a page URL or logs."""
parsed = urlparse(url)
return f"{parsed.scheme}://{parsed.netloc}" if parsed.netloc else ""
async def complete_connect_flow(
request: Request,
flow_handle: str,
session_user_id: str | None,
cache: DualCache,
) -> Response:
"""The deliberate finish step of the connect flow: mint the gateway authorization
code and send the browser back to the client.
Reached by POST so a cross-site GET cannot trigger it, and bound to the HttpOnly
per-flow cookie plus an exact match between the signed-in user and the user sealed
into the flow: a link crafted by another party dies here with ``access_denied``
instead of minting a code for the victim's identity. The flow is single-use (an atomic
claim on its ``jti``), so a double-submit cannot mint two codes from one sign-in.
"""
sealed_flow = request.cookies.get(_flow_cookie_name(flow_handle))
if sealed_flow is None:
return _oauth_error(400, "invalid_request", "unknown or expired connect flow")
flow = _open_sealed(sealed_flow, _UNPREFIXED, _ConnectFlow, _CONNECT_FLOW_DEBUG_KEY)
if flow is None:
return _oauth_error(400, "invalid_request", "unknown or expired connect flow")
now = datetime.now(timezone.utc)
if now.timestamp() >= flow.exp:
return _oauth_error(400, "invalid_request", "the connect flow has expired; restart the connection")
if session_user_id is None:
return _oauth_error(401, "login_required", "sign in to LiteLLM to finish connecting")
if session_user_id != flow.user_id:
return _oauth_error(403, "access_denied", "the signed-in user does not match this connect flow")
if not await _SingleUseGuard(cache).claim(
f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection")
code = _seal(
GATEWAY_AUTH_CODE_PREFIX,
_GatewayAuthCode(
user_id=flow.user_id,
client_id=flow.client_id,
redirect_uri=flow.redirect_uri,
code_challenge=flow.code_challenge,
jti=secrets.token_urlsafe(24),
iat=int(now.timestamp()),
exp=int(now.timestamp()) + GATEWAY_AUTH_CODE_TTL_SECONDS,
),
)
params = {"code": code, **({"state": flow.state} if flow.state else {})}
response = RedirectResponse(_append_query_params(flow.redirect_uri, params), status_code=303)
path, secure = _cookie_path_and_secure(request)
response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax")
return response
def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool:
"""RFC 7636 S256 verification, total over hostile input. The comparison is over bytes
so a non-ASCII ``code_challenge`` (which reaches here unvalidated from the client's
authorize request) simply fails to match instead of raising ``TypeError`` the way
``hmac.compare_digest`` does on two ``str`` with non-ASCII content. The verifier is
ASCII per spec; a compliant client's challenge is base64url and matches."""
digest = hashlib.sha256(code_verifier.encode("ascii", "replace")).digest()
computed = urlsafe_b64encode(digest).rstrip(b"=")
return hmac.compare_digest(computed, code_challenge.encode("utf-8"))
class _SingleUseGuard:
"""Atomic single-use claim for a one-time id (an auth-code, connect-flow ``jti``, or refresh-token
``jti``) over the injected proxy cache.
Uses an atomic increment rather than a get-then-set: two concurrent redemptions of the same id
cannot both observe "unused", because exactly one increment returns 1. The claim IS the gate, so it
fails closed. Crucially, the increment must be recorded in a backend SHARED across replicas, or the
single-use property is per-worker only (each replica's in-memory counter returns 1, so a captured
id replays through a different worker):
- When a Redis backend is configured it is the SOLE authority: the claim goes straight to Redis
(``INCR`` is atomic across replicas), and any Redis fault fails the claim CLOSED — it never falls
back to the per-worker in-memory count (``DualCache.async_increment_cache`` does fall back, which
is exactly the replay window this avoids).
- With no Redis configured (single-replica) the in-memory increment is authoritative within the one
process. A multi-worker deployment must run Redis for the guarantee to hold across workers.
The id's own TTL is the outer bound. For the auth code, PKCE binding is the primary defense against
interception; this makes the RFC 6749 4.1.2 single-use property reliable on top of it."""
def __init__(self, cache: DualCache) -> None:
self._cache = cache
async def claim(self, key: str, ttl_seconds: int) -> bool:
"""Atomically claim ``key``. ``True`` iff this caller is the first (increment to 1); ``False``
on a replay (>1) or when the claim could not be recorded in the shared backend (fail closed)."""
from litellm.proxy.proxy_server import redis_usage_cache # noqa: PLC0415 # circular import at module load
# Resolve the shared authority HERE rather than trusting the injected cache: callers pass
# user_api_key_cache, which only carries a redis_cache when enable_redis_auth_cache is set
# (off by default), so a guard that read its injected cache silently degraded every claim to
# a per-worker count on a stock multi-worker deployment. redis_usage_cache is the store the
# proxy already treats as cross-worker, so no call site can wire the guarantee away.
redis_cache = redis_usage_cache or getattr(self._cache, "redis_cache", None)
if redis_cache is not None:
# Shared, atomic authority for multi-replica deployments. Claim ONLY against Redis and fail
# CLOSED on any Redis fault (async_increment re-raises) rather than fall back to the
# per-worker in-memory count, which would let each replica observe count==1 and replay the id.
try:
count = await redis_cache.async_increment(key, 1, ttl=ttl_seconds)
except Exception as e: # noqa: BLE001 # ANY Redis fault fails the single-use claim closed
verbose_logger.warning(
"mcp gateway single-use claim: shared cache backend unavailable, failing closed: %s", e
)
return False
return count == 1
# No shared backend configured (single-replica): the in-memory increment is authoritative.
count = await self._cache.async_increment_cache(key, 1, ttl=ttl_seconds, local_only=True)
return count == 1
def _session_token_pair(principal: SessionPrincipal, keys: SessionKeys, now: datetime) -> Response:
access = mint_session_token(principal, keys, now)
refresh = mint_session_refresh_token(principal, keys, now)
if not isinstance(access, MintedSessionToken) or not isinstance(refresh, MintedSessionToken):
return _oauth_error(500, "server_error", "failed to mint the session credential")
return JSONResponse(
status_code=200,
content={
"access_token": access.token.get_secret_value(),
"token_type": "Bearer",
"expires_in": int((access.expires_at - now).total_seconds()),
"refresh_token": refresh.token.get_secret_value(),
},
headers=TOKEN_NO_CACHE_HEADERS,
)
def _reload_failure_response(failure: ReloadUserFailure) -> Response:
"""Map the live-user revalidation failure onto its OAuth error, exhaustively, so a new
``ReloadUserFailure`` member is a type error here rather than silently 400ing."""
match failure:
case "unavailable":
return _oauth_error(503, "temporarily_unavailable", "the gateway database is unavailable; retry")
case "unresolvable":
return _oauth_error(500, "server_error", "the gateway is not configured to resolve users")
case "no_active_key":
return _oauth_error(400, "invalid_grant", "the user for this grant is no longer active")
case _:
assert_never(failure)
async def aggregate_token(
request: Request,
grant_type: str,
code: str | None,
redirect_uri: str | None,
client_id: str,
code_verifier: str | None,
refresh_token: str | None,
master_key: str | None,
reload_user: ReloadUser,
cache: DualCache,
) -> Response:
"""The aggregate token verb: authorization_code and refresh_token grants for the
identity-only session pair. Every path re-validates the litellm user live before
minting, so a deactivated user cannot obtain or renew a session."""
if master_key is None:
verbose_logger.error("mcp_gateway_dcr token grant rejected: no master_key configured")
return _oauth_error(500, "server_error", "the gateway has no master key configured")
keys = session_keys_from_master_key(master_key)
now = datetime.now(timezone.utc)
if grant_type == "authorization_code":
return await _authorization_code_grant(
code=code,
redirect_uri=redirect_uri,
client_id=client_id,
code_verifier=code_verifier,
keys=keys,
now=now,
reload_user=reload_user,
guard=_SingleUseGuard(cache),
)
if grant_type == "refresh_token":
return await _refresh_token_grant(
refresh_token=refresh_token,
client_id=client_id,
keys=keys,
now=now,
reload_user=reload_user,
guard=_SingleUseGuard(cache),
)
return _oauth_error(400, "unsupported_grant_type", "grant_type must be authorization_code or refresh_token")
async def _authorization_code_grant(
code: str | None,
redirect_uri: str | None,
client_id: str,
code_verifier: str | None,
keys: SessionKeys,
now: datetime,
reload_user: ReloadUser,
guard: _SingleUseGuard,
) -> Response:
if not code or not redirect_uri or not code_verifier:
return _oauth_error(400, "invalid_request", "code, redirect_uri, and code_verifier are required")
if not MIN_CODE_VERIFIER_LENGTH <= len(code_verifier) <= MAX_CODE_VERIFIER_LENGTH:
return _oauth_error(400, "invalid_request", "code_verifier must be 43 to 128 characters (RFC 7636)")
parsed = _open_sealed(code, GATEWAY_AUTH_CODE_PREFIX, _GatewayAuthCode, _AUTH_CODE_DEBUG_KEY)
if parsed is None:
return _oauth_error(400, "invalid_grant", "the authorization code is invalid")
if now.timestamp() >= parsed.exp:
return _oauth_error(400, "invalid_grant", "the authorization code has expired")
if client_id != parsed.client_id or redirect_uri != parsed.redirect_uri:
return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client")
if not _pkce_verifier_matches(code_verifier, parsed.code_challenge):
return _oauth_error(400, "invalid_grant", "PKCE verification failed")
# Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable
# 503) does not consume a still-valid code and force the client to restart sign-in.
failure = await reload_user(parsed.user_id)
if failure is not None:
return _reload_failure_response(failure)
# Atomic single-use claim is the gate: on a concurrent double-redeem exactly one caller
# wins, and a claim that cannot be recorded fails closed.
if not await guard.claim(
f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", GATEWAY_AUTH_CODE_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_grant", "the authorization code was already used")
return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now)
async def _refresh_token_grant(
refresh_token: str | None,
client_id: str,
keys: SessionKeys,
now: datetime,
reload_user: ReloadUser,
guard: _SingleUseGuard,
) -> Response:
if not refresh_token:
return _oauth_error(400, "invalid_request", "refresh_token is required")
opened = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id)
if not isinstance(opened, SessionRefreshOpened):
return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client")
failure = await reload_user(opened.principal.user_id)
if failure is not None:
return _reload_failure_response(failure)
# Refresh-token rotation (OAuth 2.0 Security BCP section 4.13): the presented refresh token is
# single-use. Claim its jti before issuing the replacement pair, so a captured or replayed
# refresh token cannot mint a second pair after the legitimate holder rotated. Claimed AFTER
# user revalidation so a transient DB 503 does not burn a still-valid token; a claim that
# cannot be recorded fails closed, exactly like the authorization-code path.
if not await guard.claim(
f"{_USED_REFRESH_CACHE_PREFIX}{opened.jti}", SESSION_REFRESH_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_grant", "the refresh token was already used")
return _session_token_pair(opened.principal, keys, now)

View file

@ -13,6 +13,7 @@ import json
import os
import re
import time
from collections.abc import Sequence
from contextlib import asynccontextmanager
from typing import Any, AsyncIterator, Callable, Literal, Optional, Union, cast
from urllib.parse import urlparse
@ -49,6 +50,10 @@ from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
)
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
@ -59,17 +64,14 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
raise_classified_list_failure,
upstream_auth_challenge,
)
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
)
from litellm.proxy._experimental.mcp_server.sampling_handler import (
MCP_SAMPLING_AVAILABLE,
)
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
MCPPerUserTokenCache,
mcp_per_user_token_cache,
resolve_mcp_auth,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials import (
Error,
Ok,
@ -100,6 +102,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ServerSpec,
TokenExchangeConfig,
)
from litellm.proxy._experimental.mcp_server.sampling_handler import (
MCP_SAMPLING_AVAILABLE,
)
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
MCPMissingUserEnvVarsError,
@ -143,11 +148,9 @@ from litellm.types.mcp_server.mcp_server_manager import (
from litellm.types.utils import CallTypes
try:
from mcp.shared.tool_name_validation import (
validate_tool_name, # pyright: ignore[reportAssignmentType]
)
from mcp.shared.tool_name_validation import (
SEP_986_URL,
validate_tool_name, # pyright: ignore[reportAssignmentType]
)
except ImportError:
from pydantic import BaseModel
@ -408,6 +411,88 @@ def _restrict_discovery_to_corroborated_authorization_server(
return metadata.model_copy(update={"token_url": None, "registration_url": None})
def _redacted_origin_list(urls: Sequence[str]) -> str:
return ", ".join(_redact_mcp_resource_url(url) or "<unparseable url>" for url in urls)
def _sanitized_error_text(exc: Exception) -> str:
return re.sub(r"https?://\S+", "<url>", str(exc))[:200]
def _discovery_failure_leaves_needs_unresolved(
*,
needs_authorization_url: bool,
needs_token_url: bool,
manual_authorization_url: str | None,
manual_token_url: str | None,
) -> bool:
return (needs_authorization_url and not manual_authorization_url) or (needs_token_url and not manual_token_url)
def _warn_oauth_endpoints_unresolved(
*,
server_ref: str,
server_url: str | None,
discovery_attempted: bool,
issuer_anchored: bool,
metadata: MCPOAuthMetadata | None,
needs_authorization_url: bool,
needs_token_url: bool,
manual_authorization_url: str | None,
manual_token_url: str | None,
) -> None:
"""Log one actionable warning when a server that depends on OAuth endpoint discovery finishes a
build without the endpoints that its flows need (LIT-4658).
This is the operator-facing signal for a misconfigured server url: discovery failures themselves
are logged where they happen (``_descovery_metadata``), and this names WHICH server is affected,
which endpoints stayed unresolved after manual configuration was considered, and the remedies.
Scopes never trigger the warning on their own: scope-less metadata is normal for many servers and
warning on it every rebuild would be noise. Callers own the per-flow policy of which endpoints
are needed (client_credentials never needs authorization_url; OBO needs only token_url); the
issuer-anchored arm is excluded here because it has its own RFC 8414 §3.3 warning.
"""
if issuer_anchored:
return
unresolved = tuple(
field
for field, needed, value in (
(
"authorization_url",
needs_authorization_url,
manual_authorization_url or (metadata.authorization_url if metadata else None),
),
(
"token_url",
needs_token_url,
manual_token_url or (metadata.token_url if metadata else None),
),
)
if needed and not value
)
if not unresolved:
return
if discovery_attempted:
verbose_logger.warning(
"MCP server %s: OAuth endpoint discovery left %s unresolved (server url origin: %s). OAuth flows "
"that need them will fail with 'not configured' errors until they resolve. Check the preceding "
"'MCP OAuth' log lines for why discovery failed, verify the configured server url, or set the "
"unresolved endpoint urls manually, or set issuer to discover them from the identity provider "
"(RFC 8414)",
server_ref,
", ".join(unresolved),
_redact_mcp_resource_url(server_url) or "<no url>",
)
return
verbose_logger.warning(
"MCP server %s uses OAuth but has no discovery source (no server url or pinned issuer), and %s not "
"set manually. Set the missing endpoint urls on the server, or set issuer to discover them from the "
"identity provider (RFC 8414)",
server_ref,
" and ".join(unresolved) + (" is" if len(unresolved) == 1 else " are"),
)
def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None:
"""Drop a cached entry after the user stores or clears their env var values
so the next request reads the fresh value instead of a stale one."""
@ -884,10 +969,10 @@ def _create_sampling_callback(user_api_key_auth: Optional[Any] = None):
return None
async def _sampling_callback(context, params):
import litellm
from litellm.proxy._experimental.mcp_server.sampling_handler import (
handle_sampling_create_message,
)
import litellm
from litellm.proxy._experimental.mcp_server.server import (
get_active_auth_context,
)
@ -1284,6 +1369,15 @@ class MCPServerManager:
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
is_discovery_auth_type or obo_needs_discovery
)
config_oauth2_flow = server_config.get("oauth2_flow", None)
needs_authorization_url = is_discovery_auth_type and config_oauth2_flow != "client_credentials"
needs_token_url = is_discovery_auth_type or obo_needs_discovery
warn_on_empty_discovery = _discovery_failure_leaves_needs_unresolved(
needs_authorization_url=needs_authorization_url,
needs_token_url=needs_token_url,
manual_authorization_url=manual_authorization_url,
manual_token_url=manual_token_url,
)
if not should_discover:
mcp_oauth_metadata = None
elif use_issuer_anchor and manual_issuer is not None:
@ -1292,6 +1386,7 @@ class MCPServerManager:
mcp_oauth_metadata = await self._descovery_metadata(
server_url=server_url,
allow_origin_fallback=is_discovery_auth_type,
warn_when_no_metadata=warn_on_empty_discovery,
)
if use_issuer_anchor:
@ -1326,7 +1421,6 @@ class MCPServerManager:
)
effective_issuer = manual_issuer or discovered_issuer
config_oauth2_flow = server_config.get("oauth2_flow", None)
if auth_type == MCPAuth.oauth2 and config_oauth2_flow not in (
"client_credentials",
"authorization_code",
@ -1358,6 +1452,18 @@ class MCPServerManager:
"authorization-code flow."
)
_warn_oauth_endpoints_unresolved(
server_ref=server_name or server_id,
server_url=server_url,
discovery_attempted=should_discover,
issuer_anchored=use_issuer_anchor,
metadata=gated_oauth_metadata,
needs_authorization_url=needs_authorization_url,
needs_token_url=needs_token_url,
manual_authorization_url=manual_authorization_url,
manual_token_url=manual_token_url,
)
new_server = MCPServer(
server_id=server_id,
name=name_for_prefix,
@ -1485,14 +1591,12 @@ class MCPServerManager:
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
build_input_schema,
create_tool_function,
load_openapi_spec_async,
resolve_operation_params,
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
get_base_url as get_openapi_base_url,
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
load_openapi_spec_async,
resolve_operation_params,
)
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
)
@ -1681,10 +1785,20 @@ class MCPServerManager:
scopes: Optional[list[str]],
token_exchange_endpoint: Optional[str],
) -> Optional[MCPOAuthMetadata]:
obo_needs_discovery = self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url)
needs_authorization_url = (
is_discovery_auth_type and getattr(mcp_server, "oauth2_flow", None) != "client_credentials"
)
needs_token_url = is_discovery_auth_type or obo_needs_discovery
warn_on_empty_discovery = _discovery_failure_leaves_needs_unresolved(
needs_authorization_url=needs_authorization_url,
needs_token_url=needs_token_url,
manual_authorization_url=manual_authorization_url,
manual_token_url=manual_token_url,
)
has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes)
needs_discovery = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
(is_discovery_auth_type and not has_all_upstream_oauth_fields)
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url)
(is_discovery_auth_type and not has_all_upstream_oauth_fields) or obo_needs_discovery
)
if not needs_discovery:
mcp_oauth_metadata: Optional[MCPOAuthMetadata] = None
@ -1694,24 +1808,32 @@ class MCPServerManager:
mcp_oauth_metadata = await self._descovery_metadata(
server_url=server_url, # type: ignore[arg-type]
allow_origin_fallback=is_discovery_auth_type,
)
if needs_discovery and not use_issuer_anchor and mcp_oauth_metadata is None:
verbose_logger.warning(
"MCP OAuth discovery yielded no metadata for server %s (%s); "
"OAuth endpoints/scopes stay unresolved until a rebuild succeeds",
mcp_server.server_id,
server_url,
warn_when_no_metadata=warn_on_empty_discovery,
)
if use_issuer_anchor:
return mcp_oauth_metadata
if is_discovery_auth_type:
return _restrict_discovery_to_corroborated_authorization_server(
gated_metadata = (
_restrict_discovery_to_corroborated_authorization_server(
mcp_oauth_metadata,
manual_authorization_url,
mcp_server.server_id,
bool(getattr(mcp_server, "dcr_bridge", None)),
)
return mcp_oauth_metadata
if is_discovery_auth_type
else mcp_oauth_metadata
)
_warn_oauth_endpoints_unresolved(
server_ref=mcp_server.alias or mcp_server.server_name or mcp_server.server_id,
server_url=server_url,
discovery_attempted=needs_discovery,
issuer_anchored=False,
metadata=gated_metadata,
needs_authorization_url=needs_authorization_url,
needs_token_url=needs_token_url,
manual_authorization_url=manual_authorization_url,
manual_token_url=manual_token_url,
)
return gated_metadata
async def build_mcp_server_from_table(
self,
@ -2197,6 +2319,56 @@ class MCPServerManager:
return [server_id for server_id in submitted_server_ids if self.get_mcp_server_by_id(server_id) is not None]
async def operator_open_server_ids(
self,
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
allow_all_server_ids: list[str] | None = None,
submitted_server_ids: list[str] | None = None,
) -> set:
"""Servers reachable through OPEN channels rather than a grant: operator-opened
``allow_all_keys`` servers, plus the caller's own active BYOM submissions when the caller
carries no explicit ``mcp_servers`` scope.
The single owner of that question for BOTH axes. The server union in
``get_allowed_mcp_servers`` adds these ids, and the admitted subject's tool resolution asks
the same question to treat an open-channel server as default-open for tools — exactly how a
virtual key experiences it. Encoding the channel membership twice is how a server ends up
listable but uninvokable.
Empty inside a toolset scope: toolset_mcp_route / dynamic_mcp_route set
``_mcp_active_toolset_id`` before calling the handler, pinning the request to the toolset's
own servers (checking op.mcp_toolsets==[] instead would false-positive on DB-default rows
where Postgres initialises the column to ARRAY[]::TEXT[]).
``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union,
which precomputes both for its fallback path, does not compute them twice."""
from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415
_mcp_active_toolset_id,
)
if _mcp_active_toolset_id.get() is not None:
return set()
if allow_all_server_ids is None:
allow_all_server_ids = self.get_allow_all_keys_server_ids()
open_ids = set(allow_all_server_ids)
key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None
# "Explicitly scoped, so do not widen with BYOM" is a rule about a CREDENTIAL that carries
# its own mcp_servers list. It does not describe a keyless admitted subject: its
# object_permission is the user's own row, whose mcp_servers column is [] by DB default, so
# applying this rule would hide almost every admitted user's OWN submitted servers. Their
# submissions are theirs by authorship, and their scope comes from the per-source union.
has_explicit_object_permission = (
not _is_mcp_admitted_user_subject(user_api_key_auth)
and key_object_permission is not None
and (key_object_permission.mcp_servers is not None)
)
if not has_explicit_object_permission:
if submitted_server_ids is None:
submitted_server_ids = await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth)
open_ids.update(submitted_server_ids)
return open_ids
async def get_allowed_mcp_servers(self, user_api_key_auth: Optional[UserAPIKeyAuth] = None) -> list[str]:
"""
Get the allowed MCP Servers for the user.
@ -2210,11 +2382,22 @@ class MCPServerManager:
allow_all_server_ids = self.get_allow_all_keys_server_ids()
# A keyless admitted subject is resolved per grant source, and channel decisions that are
# absolute for a scoped KEY credential are not absolute for it: its own opt-out silences its
# own source (handled per source in the resolver), never its teams' grants, and its admin
# role does not swallow the grant model — a session bearer is a third-party client
# credential, not the dashboard, so an admin signing in through the connect flow gets their
# grants like anyone else rather than handing the client the full registry ahead of every
# per-team org ceiling.
is_admitted_subject = _is_mcp_admitted_user_subject(user_api_key_auth)
# The key explicitly opted out of every MCP server. Return zero before
# layering on allow_all_keys or submitted servers so the opt-out is absolute.
key_object_permission = user_api_key_auth.object_permission if user_api_key_auth else None
if key_object_permission is not None and (
SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or [])
if (
not is_admitted_subject
and key_object_permission is not None
and (SpecialMCPServerNames.no_mcp_servers.value in (key_object_permission.mcp_servers or []))
):
return []
@ -2234,8 +2417,14 @@ class MCPServerManager:
)
try:
# If admin but NO explicit object permission, get all servers
if user_api_key_auth and _user_has_admin_view(user_api_key_auth) and not has_explicit_object_permission:
# If admin but NO explicit object permission, get all servers (never for an admitted
# subject — see is_admitted_subject above)
if (
user_api_key_auth
and not is_admitted_subject
and _user_has_admin_view(user_api_key_auth)
and not has_explicit_object_permission
):
verbose_logger.debug("Admin user without explicit object_permission - returning all servers")
return list(self.get_registry().keys())
@ -2243,20 +2432,14 @@ class MCPServerManager:
allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
verbose_logger.debug(f"Allowed MCP Servers for user api key auth: {allowed_mcp_servers}")
combined_servers = set(allowed_mcp_servers)
# Only skip allow_all_keys servers when the request is inside a toolset
# scope. toolset_mcp_route / dynamic_mcp_route set _mcp_active_toolset_id
# before calling the handler — that ContextVar is the reliable signal.
# Using op.mcp_toolsets==[] would false-positive on DB-default rows where
# Postgres initialises the column to ARRAY[]::TEXT[].
from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415
_mcp_active_toolset_id,
combined_servers.update(
await self.operator_open_server_ids(
user_api_key_auth,
allow_all_server_ids=allow_all_server_ids,
submitted_server_ids=submitted_server_ids,
)
)
in_toolset_scope = _mcp_active_toolset_id.get() is not None
if not in_toolset_scope:
combined_servers.update(allow_all_server_ids)
combined_servers.update(submitted_server_ids)
# For anonymous callers (no user_id, no role), also surface any
# servers the operator has opted into upstream-delegated auth.
# These servers handle their own auth at the upstream level, so
@ -2903,9 +3086,7 @@ class MCPServerManager:
)
):
spec = None
auth_value = (
await resolve_mcp_auth(server, mcp_auth_header, subject_token=subject_token) if spec is None else None
)
auth_value = await resolve_mcp_auth(server, mcp_auth_header) if spec is None else None
# Create sampling and elicitation callbacks for this client
sampling_cb = _create_sampling_callback(user_api_key_auth=user_api_key_auth) if server.allow_sampling else None
@ -3430,6 +3611,7 @@ class MCPServerManager:
server_url: str,
*,
allow_origin_fallback: bool = True,
warn_when_no_metadata: bool = False,
) -> Optional[MCPOAuthMetadata]:
"""Discover OAuth metadata by following RFC 9728 (protected resource metadata discovery).
@ -3438,8 +3620,32 @@ class MCPServerManager:
it (a human sees the redirect), but token_exchange (OBO) sets it False so the gateway never
exchanges a subject token against an endpoint it inferred rather than one explicitly configured
or authoritatively advertised via RFC 9728 / RFC 8414.
"""
``warn_when_no_metadata`` makes an all-empty result log one WARNING with the per-step attempt
outcomes (LIT-4658), so a misconfigured server url is diagnosable from default-level logs. The
server loaders set it; the issuer-anchored resource-scopes lookup keeps it off because empty
scopes are not a fault there.
"""
metadata, attempts = await self._discover_metadata_recording_attempts(
server_url, allow_origin_fallback=allow_origin_fallback
)
if metadata is None and warn_when_no_metadata:
verbose_logger.warning(
"MCP OAuth endpoint discovery against %s found no authorization server metadata. Attempts: %s. "
"The MCP server url may be misconfigured, or the upstream may not support OAuth discovery "
"(RFC 9728 / RFC 8414)",
_redact_mcp_resource_url(server_url) or "<unparseable url>",
"; ".join(attempts) if attempts else "none recorded",
)
return metadata
async def _discover_metadata_recording_attempts(
self,
server_url: str,
*,
allow_origin_fallback: bool,
) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]:
origin = _redact_mcp_resource_url(server_url) or "<unparseable url>"
try:
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
response = await client.get(server_url)
@ -3452,67 +3658,112 @@ class MCPServerManager:
if metadata is None and not resource_scopes and authorization_servers and response.status_code == 200:
verbose_logger.warning(
"MCP OAuth discovery for %s received 200 OK without RFC 9728 challenge and no discoverable authorization metadata.",
server_url,
origin,
)
attempts = (
f"GET {origin}: HTTP {response.status_code} (no RFC 9728 challenge)",
*(
("well-known protected-resource lookup found no authorization servers",)
if not authorization_servers
else ()
),
*(
(f"authorization server metadata fetch failed for: {_redacted_origin_list(authorization_servers)}",)
if authorization_servers and metadata is None
else ()
),
)
if metadata is None and resource_scopes:
return MCPOAuthMetadata(scopes=resource_scopes)
return MCPOAuthMetadata(scopes=resource_scopes), attempts
if metadata is not None and resource_scopes:
metadata.scopes = resource_scopes
return metadata
return metadata, attempts
except HTTPStatusError as exc:
verbose_logger.debug(
"MCP OAuth discovery for %s received status error: %s",
server_url,
exc,
)
header_value: Optional[str] = None
if exc.response is not None:
header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get(
"www-authenticate"
)
resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value)
authorization_servers = []
resource_scopes = None
if resource_metadata_url:
(
authorization_servers,
resource_scopes,
) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url, server_url)
else:
(
authorization_servers,
resource_scopes,
) = await self._attempt_well_known_discovery(server_url)
metadata = None
used_origin_fallback = False
if allow_origin_fallback and not authorization_servers:
try:
parsed_url = urlparse(server_url)
if parsed_url.scheme and parsed_url.netloc:
authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"]
used_origin_fallback = True
except Exception:
authorization_servers = []
if authorization_servers:
metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url)
if metadata is not None and used_origin_fallback:
metadata.from_origin_fallback = True
preferred_scopes = scopes or resource_scopes
if metadata is None and preferred_scopes:
metadata = MCPOAuthMetadata(scopes=preferred_scopes)
elif metadata is not None and preferred_scopes:
metadata.scopes = preferred_scopes
return metadata
return await self._discover_after_status_error(server_url, exc, allow_origin_fallback=allow_origin_fallback)
except Exception as exc: # pragma: no cover - network/transient issues
verbose_logger.debug("MCP OAuth discovery failed for %s: %s", server_url, exc)
return None
return None, (f"GET {origin}: {type(exc).__name__}: {_sanitized_error_text(exc)}",)
async def _discover_after_status_error(
self,
server_url: str,
exc: HTTPStatusError,
*,
allow_origin_fallback: bool,
) -> tuple[MCPOAuthMetadata | None, tuple[str, ...]]:
origin = _redact_mcp_resource_url(server_url) or "<unparseable url>"
verbose_logger.debug(
"MCP OAuth discovery for %s received status error: %s",
server_url,
exc,
)
header_value: Optional[str] = None
if exc.response is not None:
header_value = exc.response.headers.get("WWW-Authenticate") or exc.response.headers.get("www-authenticate")
status_attempt = (
f"GET {origin}: HTTP {exc.response.status_code}"
if exc.response is not None
else f"GET {origin}: status error"
)
resource_metadata_url, scopes = self._parse_www_authenticate_header(header_value)
authorization_servers = []
resource_scopes = None
if resource_metadata_url:
(
authorization_servers,
resource_scopes,
) = await self._fetch_oauth_metadata_from_resource(resource_metadata_url, server_url)
lookup_attempt = (
None
if authorization_servers
else "challenge-advertised resource metadata yielded no authorization servers"
)
else:
(
authorization_servers,
resource_scopes,
) = await self._attempt_well_known_discovery(server_url)
lookup_attempt = (
None
if authorization_servers
else "no challenge-advertised resource metadata; well-known protected-resource lookup found no authorization servers"
)
metadata = None
used_origin_fallback = False
if allow_origin_fallback and not authorization_servers:
try:
parsed_url = urlparse(server_url)
if parsed_url.scheme and parsed_url.netloc:
authorization_servers = [f"{parsed_url.scheme}://{parsed_url.netloc}"]
used_origin_fallback = True
except Exception:
authorization_servers = []
fallback_attempt = None
if authorization_servers:
metadata = await self._fetch_authorization_server_metadata(authorization_servers, server_url)
if metadata is not None and used_origin_fallback:
metadata.from_origin_fallback = True
if metadata is None:
fallback_attempt = (
f"origin fallback: no authorization server metadata at {origin}"
if used_origin_fallback
else f"authorization server metadata fetch failed for: {_redacted_origin_list(authorization_servers)}"
)
attempts = tuple(entry for entry in (status_attempt, lookup_attempt, fallback_attempt) if entry)
preferred_scopes = scopes or resource_scopes
if metadata is None and preferred_scopes:
return MCPOAuthMetadata(scopes=preferred_scopes), attempts
if metadata is not None and preferred_scopes:
metadata.scopes = preferred_scopes
return metadata, attempts
def _parse_www_authenticate_header(self, header_value: Optional[str]) -> tuple[Optional[str], Optional[list[str]]]:
if not header_value:

View file

@ -26,7 +26,6 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy._experimental.mcp_server.auth import token_exchange
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
build_token_endpoint_client_auth,
)
@ -58,17 +57,12 @@ class MCPOAuth2TokenCache(InMemoryCache):
def _has_client_credentials_config(server: "MCPServer") -> bool:
return bool(server.client_id and server.client_secret and server.token_url)
async def async_get_token(
self,
server: "MCPServer",
*,
require_client_credentials_flow: bool = True,
) -> Optional[str]:
async def async_get_token(self, server: "MCPServer") -> Optional[str]:
"""Return a valid access token, fetching or refreshing as needed.
Returns ``None`` when the server lacks client credentials config.
"""
if require_client_credentials_flow and not server.has_client_credentials:
if not server.has_client_credentials:
return None
if not self._has_client_credentials_config(server):
return None
@ -278,36 +272,16 @@ mcp_per_user_token_cache = MCPPerUserTokenCache()
async def resolve_mcp_auth(
server: "MCPServer",
mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None,
subject_token: Optional[str] = None,
) -> Optional[Union[str, Dict[str, str]]]:
"""Resolve the auth value for an MCP server.
Priority:
1. ``mcp_auth_header`` — per-request/per-user override
2. OAuth2 Token Exchange (OBO / RFC 8693) — exchange user token for scoped token
3. OAuth2 client_credentials token — auto-fetched and cached
4. ``server.authentication_token`` — static token from config/DB
2. OAuth2 client_credentials token — auto-fetched and cached
3. ``server.authentication_token`` — static token from config/DB
"""
if mcp_auth_header:
return mcp_auth_header
if server.has_token_exchange_config:
if subject_token:
return await token_exchange.mcp_token_exchange_handler.exchange_token(subject_token, server)
# No subject_token — fall back to client_credentials using the same client
# credentials and token_url so M2M scenarios still work.
if server.client_id and server.client_secret and server.token_url:
return await mcp_oauth2_token_cache.async_get_token(
server,
require_client_credentials_flow=False,
)
# OBO configured but no subject_token and missing client credentials — warn
# rather than silently proceeding unauthenticated.
verbose_logger.warning(
"MCP server '%s' is configured for token exchange (OBO) but no subject_token "
"was provided and client credentials (client_id/client_secret/token_url) are "
"incomplete. The request will proceed without authentication.",
server.server_id,
)
if server.has_client_credentials:
return await mcp_oauth2_token_cache.async_get_token(server)
return server.authentication_token

View file

@ -4,7 +4,7 @@
import os
from ipaddress import ip_address
from typing import Any, Dict, List, NoReturn, Optional
from urllib.parse import ParseResult, urlparse, urlunparse
from urllib.parse import ParseResult, urlparse, urlsplit, urlunparse, urlunsplit
from fastapi import HTTPException, Request
@ -70,6 +70,29 @@ def _origin_label(scheme: str, netloc: str) -> str:
return f"{scheme}://{netloc}" if netloc else f"{scheme}://"
def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]:
"""Reduce an MCP server URL to its origin (scheme + host + port) for logging.
Everything else is dropped: userinfo (``user:pass@``), the query string, the
fragment, and the path, because hosted MCP servers routinely embed the
credential in the path (e.g. ``/mcp/s/<token>``) and this value is persisted
in spend-log metadata that a caller who can invoke the tool can read back.
Returns None when the URL has no host to identify (nothing safe to log).
"""
if not isinstance(url, str) or not url:
return None
try:
parts = urlsplit(url)
hostname = parts.hostname
port = parts.port
except ValueError:
return None
if not hostname:
return None
netloc = f"{hostname}:{port}" if port else hostname
return urlunsplit((parts.scheme, netloc, "", "", "")) or None
def _resolve_proxy_base_url_env() -> Optional[str]:
global _warned_invalid_proxy_base_url
configured = os.environ.get("PROXY_BASE_URL", "").strip()
@ -343,8 +366,36 @@ def _parse_redirect_uri_for_validation(redirect_uri: str) -> ParseResult:
)
def _validate_trusted_http_redirect_shape(parsed: ParseResult) -> bool:
"""Return True when ``parsed`` is an allowlisted native callback (caller may return)."""
def is_loopback_redirect_host(parsed: ParseResult) -> bool:
"""True when the redirect host is loopback (RFC 8252 section 7.3).
Shared by every redirect-URI policy in the MCP OAuth surface so that none of them
hand-rolls its own host list: a literal ``("localhost", "127.0.0.1", "::1")`` tuple
silently misses the rest of 127.0.0.0/8 and IPv6-mapped forms.
"""
host = (parsed.hostname or "").lower()
if host == "localhost":
return True
try:
return ip_address(host).is_loopback
except ValueError:
return False
def validate_redirect_uri_shape(parsed: ParseResult) -> bool:
"""Validate redirect-URI *hygiene* and resolve allowlisted native callbacks.
Returns True when ``parsed`` is an allowlisted native callback (the caller may accept
it outright); returns False for http/https, leaving the trust decision to the caller;
raises for a URI that no policy should ever accept (bad scheme, fragment, missing
host, userinfo, backslash in the host).
This is deliberately separate from :func:`validate_trusted_redirect_uri`, which adds
the *first-party* trust policy (same-origin, loopback, ops allowlist) appropriate to
the proxy's own OAuth endpoints. Public dynamic-client registration accepts any https
client and relies on PKCE plus the consent screen instead, so it shares this hygiene
rule but not that trust policy.
"""
if parsed.scheme not in ("http", "https"):
if _matches_trusted_native_redirect_uri(parsed):
return True
@ -396,14 +447,8 @@ def _trusted_redirect_uri_is_allowed(
):
return True
host = (parsed.hostname or "").lower()
if host == "localhost":
if is_loopback_redirect_host(parsed):
return True
try:
if ip_address(host).is_loopback:
return True
except ValueError:
pass
if parsed.scheme == "https":
for entry in _parse_trusted_redirect_origins():
@ -522,7 +567,7 @@ def validate_trusted_redirect_uri(request: Request, redirect_uri: str) -> None:
:func:`validate_loopback_redirect_uri`.
"""
parsed = _parse_redirect_uri_for_validation(redirect_uri)
if _validate_trusted_http_redirect_shape(parsed):
if validate_redirect_uri_shape(parsed):
return
redirect_netloc = _strip_default_port(parsed.scheme, parsed.netloc)
proxy_base = _resolve_proxy_base_for_redirect(request)

View file

@ -149,6 +149,7 @@ class SessionRefreshOpened(BaseModel):
model_config = ConfigDict(frozen=True)
tag: Literal["opened"] = "opened"
principal: SessionPrincipal
jti: str
class SessionRefreshInvalid(BaseModel):
@ -187,4 +188,4 @@ def open_session_refresh_bearer(
return SessionRefreshInvalid()
if opened.principal.client_id != expected_client_id:
return SessionRefreshInvalid()
return SessionRefreshOpened(principal=opened.principal)
return SessionRefreshOpened(principal=opened.principal, jti=opened.jti)

View file

@ -113,10 +113,12 @@ class MintedSessionToken(BaseModel):
class OpenedSessionToken(BaseModel):
"""A validated session token of either kind: the principal it was minted for."""
"""A validated session token of either kind: the principal it was minted for, plus the
``jti`` so the token endpoint can enforce single-use rotation on a refresh token."""
model_config = ConfigDict(frozen=True)
principal: SessionPrincipal
jti: str
class SessionTokenTooLarge(BaseModel):
@ -320,7 +322,9 @@ def _open(
return SessionMalformed()
if now.timestamp() >= claims.exp:
return SessionExpired()
return OpenedSessionToken(principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id))
return OpenedSessionToken(
principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id), jti=claims.jti
)
def _decode_claims(

View file

@ -230,16 +230,33 @@ if MCP_AVAILABLE:
return server_auth
return mcp_auth_header
def _get_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]:
"""Return the subset of *allowed_server_ids* whose servers use OAuth2 auth.
def _is_v1_resolved_oauth2_server(server: Optional[MCPServer]) -> bool:
"""Whether this server's per-user OAuth2 token is still resolved by v1.
Used as a cheap pre-flight check to skip bulk credential fetching when no
OAuth2 servers are involved in the current request.
A server the v2 resolver owns reads its stored token from the resolver at connect
time and drops any Authorization built for it here, so the v1 lookup would be a DB
round-trip whose result is discarded. Mirrors the same guard on the protocol listing
path and in ``_resolve_oauth2_headers_for_tool_call``.
"""
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
to_server_spec,
)
if getattr(server, "auth_type", None) != MCPAuth.oauth2:
return False
return to_server_spec(server) is None
def _v1_resolved_oauth2_server_ids(allowed_server_ids: List[str]) -> Set[str]:
"""Return the subset of *allowed_server_ids* whose per-user OAuth2 token is still
resolved by v1.
Used as a cheap pre-flight check to skip bulk credential fetching when no such
server is involved in the current request.
"""
return {
sid
for sid in allowed_server_ids
if getattr(global_mcp_server_manager.get_mcp_server_by_id(sid), "auth_type", None) == MCPAuth.oauth2
if _is_v1_resolved_oauth2_server(global_mcp_server_manager.get_mcp_server_by_id(sid))
}
async def _get_user_oauth_extra_headers(
@ -253,11 +270,13 @@ if MCP_AVAILABLE:
the MCP server the same way the admin "Add MCP / Authorize and Fetch" flow does.
Returns None for non-OAuth2 servers or when no credential is stored.
A server the v2 resolver owns is skipped; see ``_is_v1_resolved_oauth2_server``.
Args:
prefetched_creds: Optional dict keyed by server_id with credential payloads.
When provided, avoids a per-server DB round-trip.
"""
if getattr(server, "auth_type", None) != MCPAuth.oauth2:
if not _is_v1_resolved_oauth2_server(server):
return None
user_id = getattr(user_api_key_dict, "user_id", None)
server_id = getattr(server, "server_id", None)
@ -320,38 +339,6 @@ if MCP_AVAILABLE:
verbose_logger.warning(f"_prefetch_user_oauth_creds: failed to prefetch for user={user_id}: {e}")
return {}
async def _get_bulk_user_oauth_headers(
user_api_key_dict: UserAPIKeyAuth,
) -> Dict[str, Dict[str, str]]:
"""
Fetch ALL OAuth2 credentials for the current user in a single DB query and
return a mapping of server_id → {"Authorization": "Bearer <token>"}.
This is the batch alternative to calling _get_user_oauth_extra_headers
per-server inside a loop (N+1 DB queries).
"""
user_id = getattr(user_api_key_dict, "user_id", None)
if not user_id:
return {}
try:
from litellm.proxy._experimental.mcp_server.db import (
list_user_oauth_credentials,
)
from litellm.proxy.utils import get_prisma_client_or_throw
prisma_client = get_prisma_client_or_throw(
"Database not connected. Connect a database to use OAuth2 MCP tools."
)
creds = await list_user_oauth_credentials(prisma_client, user_id)
return {
c["server_id"]: {"Authorization": f"Bearer {c['access_token']}"}
for c in creds
if c.get("access_token") and c.get("server_id")
}
except Exception:
verbose_logger.debug("Failed to bulk-fetch OAuth credentials", exc_info=True)
return {}
def _create_tool_response_objects(tools, server: MCPServer):
"""Helper function to create tool response objects.
@ -825,7 +812,7 @@ if MCP_AVAILABLE:
# to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers.
prefetched_oauth_creds = (
await _prefetch_user_oauth_creds(user_api_key_dict)
if _get_oauth2_server_ids(allowed_server_ids)
if _v1_resolved_oauth2_server_ids(allowed_server_ids)
else {}
)

View file

@ -27,7 +27,6 @@ from typing import (
Union,
cast,
)
from urllib.parse import urlsplit, urlunsplit
import httpx
from fastapi import FastAPI, HTTPException
@ -59,6 +58,9 @@ from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_gateway_server_name,
)
from litellm.proxy._experimental.mcp_server.mcp_debug import MCPDebug
from litellm.proxy._experimental.mcp_server.oauth_utils import (
_redact_mcp_resource_url,
)
from litellm.proxy._experimental.mcp_server.utils import (
LITELLM_MCP_SERVER_DESCRIPTION,
LITELLM_MCP_SERVER_NAME,
@ -106,27 +108,6 @@ _MAX_STATEFUL_SESSIONS_PER_OWNER = 100
_MCP_ROUTING_PEEK_MAX_BYTES = 4096
def _redact_mcp_resource_url(url: Optional[str]) -> Optional[str]:
"""Reduce an MCP server URL to its origin (scheme + host + port) for logging.
Everything else is dropped: userinfo (``user:pass@``), the query string, the
fragment, and the path, because hosted MCP servers routinely embed the
credential in the path (e.g. ``/mcp/s/<token>``) and this value is persisted
in spend-log metadata that a caller who can invoke the tool can read back.
Returns None when the URL has no host to identify (nothing safe to log).
"""
if not isinstance(url, str) or not url:
return None
try:
parts = urlsplit(url)
except ValueError:
return None
if not parts.hostname:
return None
netloc = f"{parts.hostname}:{parts.port}" if parts.port else parts.hostname
return urlunsplit((parts.scheme, netloc, "", "", "")) or None
def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
"""Remove a (user_id, server_id) entry from the BYOK credential cache.
@ -977,7 +958,17 @@ if MCP_AVAILABLE:
data = await add_litellm_data_to_request(
data=body_data,
request=request,
user_api_key_dict=user_api_key_auth,
# Bill a team-derived call to the team that granted it. A keyless admitted
# subject carries no team_id, so spend skipped team updates entirely and
# charged the user's PRIMARY org — the granting team's budget never
# accumulated (so it could never begin to block) and, cross-org, the wrong
# organization was charged. This is the ACCOUNTING half; the enforcement
# half (an already-over-budget team stops granting) lives in the source gate.
# Authorization is unaffected: it ran before this, and the union is resolved
# from the untouched auth object passed to call_mcp_tool below.
user_api_key_dict=await MCPRequestHandler.billing_auth_for_tool_call(
user_api_key_auth, tool_name=name
),
proxy_config=proxy_config,
)
else:

View file

@ -26762,6 +26762,113 @@
"title": "ToolPolicyUpdateResponse",
"type": "object"
},
"ToolSpendDailyEntry": {
"description": "Spend attributed to one tool on one UTC day.",
"properties": {
"call_count": {
"default": 0,
"title": "Call Count",
"type": "integer"
},
"date": {
"title": "Date",
"type": "string"
},
"spend": {
"default": 0.0,
"title": "Spend",
"type": "number"
},
"tool_name": {
"title": "Tool Name",
"type": "string"
}
},
"required": [
"date",
"tool_name"
],
"title": "ToolSpendDailyEntry",
"type": "object"
},
"ToolSpendEntry": {
"description": "Total spend attributed to one tool over the requested window.",
"properties": {
"call_count": {
"default": 0,
"title": "Call Count",
"type": "integer"
},
"spend": {
"default": 0.0,
"description": "Attributed spend: a request that used several tools counts its full spend toward each of them",
"title": "Spend",
"type": "number"
},
"tool_name": {
"title": "Tool Name",
"type": "string"
},
"total_tokens": {
"default": 0,
"title": "Total Tokens",
"type": "integer"
}
},
"required": [
"tool_name"
],
"title": "ToolSpendEntry",
"type": "object"
},
"ToolSpendResponse": {
"properties": {
"by_tool": {
"items": {
"$ref": "#/components/schemas/ToolSpendEntry"
},
"title": "By Tool",
"type": "array"
},
"daily": {
"items": {
"$ref": "#/components/schemas/ToolSpendDailyEntry"
},
"title": "Daily",
"type": "array"
},
"end_date": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "End Date"
},
"start_date": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"title": "Start Date"
},
"total_spend": {
"default": 0.0,
"description": "Deduplicated spend of every request that called at least one tool in the window; less than the sum of per-tool attributed spend whenever multi-tool requests exist",
"title": "Total Spend",
"type": "number"
}
},
"title": "ToolSpendResponse",
"type": "object"
},
"ToolUsageLogEntry": {
"description": "One spend log row for a tool call (for UI \"recent logs\" table).",
"properties": {
@ -26858,6 +26965,13 @@
},
"ValidationError": {
"properties": {
"ctx": {
"title": "Context",
"type": "object"
},
"input": {
"title": "Input"
},
"loc": {
"items": {
"anyOf": [
@ -27301,6 +27415,81 @@
]
}
},
"/v1/tool/spend": {
"get": {
"description": "Spend attributed to each tool over a date range, for the Cost Optimization dashboard.\n\nJoins ``LiteLLM_SpendLogToolIndex`` (which tool names ran on which request) to\n``LiteLLM_SpendLogs`` (what the request cost). A request that used multiple tools\ncounts its full spend toward each of those tools, so per-tool numbers are\nattributions. ``total_spend`` is the deduplicated spend of every request that\ncalled at least one tool in the window, so it never double counts.",
"operationId": "get_tool_spend_v1_tool_spend_get",
"parameters": [
{
"description": "YYYY-MM-DD (defaults to 30 days ago)",
"in": "query",
"name": "start_date",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "YYYY-MM-DD (defaults to 30 days ago)",
"title": "Start Date"
}
},
{
"description": "YYYY-MM-DD (defaults to today)",
"in": "query",
"name": "end_date",
"required": false,
"schema": {
"anyOf": [
{
"type": "string"
},
{
"type": "null"
}
],
"description": "YYYY-MM-DD (defaults to today)",
"title": "End Date"
}
}
],
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ToolSpendResponse"
}
}
},
"description": "Successful Response"
},
"422": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/HTTPValidationError"
}
}
},
"description": "Validation Error"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Get Tool Spend",
"tags": [
"tools"
]
}
},
"/v1/tool/{tool_name}": {
"get": {
"description": "Get details for a single tool.",

View file

@ -2605,6 +2605,28 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
user_max_budget: Optional[float] = None
request_route: Optional[str] = None
is_session_token: bool = False
# Server-only marker set exclusively by the MCP gateway admission path
# (_reload_admitted_user) for a keyless user-subject admitted via a gateway DCR session
# bearer or bridge envelope. Not a DB column and never populated from caller-controlled key
# metadata or JWT claims, so it cannot be forged to gain the team-inherited MCP grant union
# or to escape the caller-Authorization egress scrub. exclude=True keeps it out of serialization.
mcp_admitted_user_subject: bool = Field(default=False, exclude=True)
# team_id -> that team's mcp_rpm_limit map, for a keyless admitted subject that reaches MCP
# servers through several teams at once and therefore has no single team_id for the limiter to
# key off. Server-only and stripped from validated input for the same reason as the marker
# above: a forged entry would let a caller pick which team's rpm bucket it is charged against.
mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True)
via_virtual_key: bool = Field(
default=False,
exclude=True,
description=(
"Server-only marker set exclusively by the DB virtual-key and master-key auth paths via "
"post-construction assignment. Stripped from validated input so custom auth handlers, JWT "
"claims, or key metadata cannot forge it. Gates overwrite_user_with_key_hash stamping: only "
"a credential the proxy itself validated as a key may be forwarded as the provider-facing "
"user id."
),
)
budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True)
budget_throttle_pct: Optional[float] = Field(default=None, exclude=True)
user: Optional[Any] = None # Expanded user object when expand=user is used
@ -2625,6 +2647,12 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# If values is already an instance (not a dict), return it as-is
if not isinstance(values, dict):
return values
# mcp_admitted_user_subject is a server-only marker, set ONLY by the MCP gateway admission
# path via post-construction assignment. Strip it from any validated input (constructor
# kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data.
values.pop("mcp_admitted_user_subject", None)
values.pop("mcp_source_team_rpm_limits", None)
values.pop("via_virtual_key", None)
if values.get("api_key") is not None:
values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))})
if isinstance(values.get("api_key"), str):

View file

@ -2661,6 +2661,15 @@ async def get_managed_vector_store_rows_by_uuids(
return result
class OrganizationNotFoundError(Exception):
"""The organization row is CONFIRMED absent, as opposed to a lookup that failed.
Subclasses Exception so every existing except Exception caller keeps its current
behavior; it exists so a caller that wants to treat "no such org" as "no restriction" can do
that WITHOUT also swallowing an outage and silently dropping a real org ceiling.
"""
@log_db_metrics
async def get_org_object(
org_id: str,
@ -2707,25 +2716,30 @@ async def get_org_object(
query_kwargs["include"] = {"litellm_budget_table": True}
response = await OrganizationRepository(prisma_client).table.find_unique(**query_kwargs)
if response is None:
raise Exception
_org_obj = LiteLLM_OrganizationTable(**response.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
key=cache_key,
value=_org_obj,
model_type=LiteLLM_OrganizationTable,
ttl=DEFAULT_IN_MEMORY_TTL,
)
return _org_obj
except Exception:
raise Exception(
# An operational failure (DB down, timeout, cache fault) is NOT the same fact as a confirmed
# missing row, and relabelling it as "doesn't exist" made every caller unable to tell them
# apart — a caller that treats absence as "this org places no restriction" then drops a real
# org ceiling during an outage. Propagate the real error; callers that already catch
# Exception are unaffected.
raise
if response is None:
raise OrganizationNotFoundError(
f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call."
)
_org_obj = LiteLLM_OrganizationTable(**response.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
key=cache_key,
value=_org_obj,
model_type=LiteLLM_OrganizationTable,
ttl=DEFAULT_IN_MEMORY_TTL,
)
return _org_obj
async def _get_resources_from_access_groups(
access_group_ids: List[str],

View file

@ -7,12 +7,15 @@ login endpoints (e.g., /login and /v2/login).
import os
import secrets
from datetime import datetime, timedelta, timezone
from typing import Literal, Optional, cast
import jwt
from fastapi import HTTPException
import litellm
from litellm.constants import LITELLM_PROXY_ADMIN_NAME, LITELLM_UI_SESSION_DURATION
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy._types import (
LiteLLM_UserTable,
LitellmUserRoles,
@ -313,6 +316,29 @@ async def authenticate_user(
)
def _ui_session_exp_timestamp() -> int:
"""The ``exp`` claim (unix seconds) for a UI session cookie, ``LITELLM_UI_SESSION_DURATION``
from now. The virtual key sealed inside the cookie already expires after this same
duration; stamping the JWT itself gives the cookie the bounded lifetime the dashboard's
client-side expiry check and the server-side session-cookie readers both assume, instead
of a token that stays signature-valid until the master key rotates."""
ttl_seconds = duration_in_seconds(LITELLM_UI_SESSION_DURATION)
return int((datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds)).timestamp())
def encode_ui_session_jwt(returned_ui_token_object: ReturnedUITokenObject, master_key: str) -> str:
"""Encode a UI session cookie JWT with a bounded ``exp``.
The single choke point every UI login path (SSO and username/password /login, /v2,
/v3) uses to mint the ``token`` cookie, so the cookie's lifetime is set in exactly one
place and cannot drift between paths. Without the ``exp`` the cookie is valid until the
master key rotates, and the session-cookie readers that require a bounded lifetime
(the MCP interactive sign-in) reject it.
"""
claims = {**cast(dict, returned_ui_token_object), "exp": _ui_session_exp_timestamp()}
return jwt.encode(claims, master_key, algorithm="HS256")
def create_ui_token_object(
login_result: LoginResult,
general_settings: dict,

View file

@ -1497,6 +1497,13 @@ async def _user_api_key_auth_builder(
check_cache_only=True,
).resolve(hashed_token=hash_token(api_key))
)
# Key-cache entries are written only after the proxy validated a
# virtual key or the master key, but via_virtual_key is exclude=True
# so serialization drops it; restore it at this trusted boundary.
# The UI-login JWT fallback below constructs its token from a
# decrypted blob, not this cache, and stays unmarked.
if isinstance(valid_token, UserAPIKeyAuth):
valid_token.via_virtual_key = True
except Exception:
verbose_logger.debug("api key not found in cache.")
valid_token = None
@ -1614,6 +1621,7 @@ async def _user_api_key_auth_builder(
_user_api_key_obj = update_valid_token_with_end_user_params(
valid_token=_user_api_key_obj, end_user_params=end_user_params
)
_user_api_key_obj.via_virtual_key = True
return _user_api_key_obj
@ -2021,7 +2029,7 @@ async def _user_api_key_auth_builder(
# No token was found when looking up in the DB
raise Exception("Invalid proxy server token passed")
if valid_token_dict is not None:
return await _return_user_api_key_auth_obj(
virtual_key_auth_obj = await _return_user_api_key_auth_obj(
user_obj=user_obj,
api_key=api_key,
parent_otel_span=parent_otel_span,
@ -2029,6 +2037,8 @@ async def _user_api_key_auth_builder(
route=route,
start_time=start_time,
)
virtual_key_auth_obj.via_virtual_key = True
return virtual_key_auth_obj
except Exception as e:
return await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
e=e,
@ -2442,6 +2452,7 @@ async def _reserve_budget_after_common_checks(
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,
fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
)

View file

@ -1781,28 +1781,38 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit
if not mcp_server_name or not user_api_key_dict.team_id:
if not mcp_server_name:
return
mcp_rpm_limit = get_team_mcp_rpm_limit(user_api_key_dict)
if not mcp_rpm_limit:
return
# Which teams' buckets does this call charge? A key is pinned to exactly one team. A keyless
# MCP-admitted subject reaches servers through SEVERAL teams at once and has no team_id, so
# without the second source below its calls charged no team bucket at all and it outran every
# team's mcp_rpm_limit. Every applicable team is charged rather than one being picked: the
# limiter enforces all descriptors, so each team's own ceiling binds on a call made through
# its grant, and there is no arbitrary attribution when several teams grant the same server.
team_limits: list[tuple[str | None, dict[str, int] | None]] = []
if user_api_key_dict.team_id:
team_limits.append((user_api_key_dict.team_id, get_team_mcp_rpm_limit(user_api_key_dict)))
for source_team_id, source_limit in (user_api_key_dict.mcp_source_team_rpm_limits or {}).items():
team_limits.append((source_team_id, source_limit))
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
return
descriptors.append(
RateLimitDescriptor(
key="mcp_per_team",
value=f"{user_api_key_dict.team_id}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
for team_id, mcp_rpm_limit in team_limits:
if not team_id or not mcp_rpm_limit:
continue
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
continue
descriptors.append(
RateLimitDescriptor(
key="mcp_per_team",
value=f"{team_id}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
)
def _should_enforce_rate_limit(
self,

View file

@ -13,7 +13,7 @@ from starlette.datastructures import Headers
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS, PRE_CALL_EXECUTED_GUARDRAILS_KEY
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
iter_client_callback_metadata_dicts,
@ -48,6 +48,24 @@ _EXPLICIT_SESSION_HEADERS = frozenset({"x-litellm-trace-id", "x-litellm-session-
# Session-id values must be non-empty strings of alphanumerics, hyphens, or underscores
# (covers UUIDs and most common session-id formats).
_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
_SHA256_HEX_RE = re.compile(r"^[0-9a-f]{64}$")
def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None:
"""Only proxy-validated keys are stamped, proven by the unforgeable
via_virtual_key marker AND a known non-secret shape: the sha256 hex digest
UserAPIKeyAuth stores virtual keys in, or the master key's stable alias.
Custom-auth credentials arrive raw (never forward auth material) and hashed
JWTs rotate on re-issue (useless as a stable ban id), so both are skipped."""
api_key = user_api_key_dict.api_key
if not user_api_key_dict.via_virtual_key or api_key is None:
return None
if api_key == LITELLM_PROXY_MASTER_KEY_ALIAS or _SHA256_HEX_RE.fullmatch(api_key):
return api_key
return None
_ANTHROPIC_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]+$")
@ -1447,6 +1465,11 @@ async def add_litellm_data_to_request(
if "user" not in data:
data["user"] = user
if litellm.overwrite_user_with_key_hash is True:
stampable_hash = _stampable_key_hash(user_api_key_dict)
if stampable_hash is not None:
data["user"] = stampable_hash
data["secret_fields"] = SecretFields(raw_headers=_raw_headers)
## Dynamic api version (Azure OpenAI endpoints) ##

View file

@ -10,16 +10,18 @@ POST /v1/tool/policy - Update the input_policy / output_policy for a
"""
import uuid
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, List, Optional
from datetime import datetime, timedelta, timezone
from itertools import groupby
from typing import TYPE_CHECKING, Annotated, Any, List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel, TypeAdapter
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.table_repositories import (
@ -39,6 +41,9 @@ from litellm.types.tool_management import (
ToolPolicyOptionsResponse,
ToolPolicyUpdateRequest,
ToolPolicyUpdateResponse,
ToolSpendDailyEntry,
ToolSpendEntry,
ToolSpendResponse,
ToolUsageLogEntry,
ToolUsageLogsResponse,
)
@ -124,6 +129,147 @@ async def list_tools(
raise HTTPException(status_code=500, detail=str(e))
def _parse_day_start(value: str | None) -> datetime | None:
if not value:
return None
try:
return datetime.strptime(value.strip(), "%Y-%m-%d").replace(tzinfo=timezone.utc)
except ValueError:
raise HTTPException(
status_code=400,
detail=f"Invalid date format: {value}. Expected: 'YYYY-MM-DD'",
)
class _ToolSpendRow(BaseModel):
date: str
tool_name: str
call_count: int
spend: float
total_tokens: int
class _RequestTotalRow(BaseModel):
total_spend: float
_TOOL_SPEND_ROWS = TypeAdapter(list[_ToolSpendRow])
_REQUEST_TOTAL_ROWS = TypeAdapter(list[_RequestTotalRow])
def _summarize_tool(name: str, grp: tuple[_ToolSpendRow, ...]) -> ToolSpendEntry:
return ToolSpendEntry(
tool_name=name,
spend=sum(r.spend for r in grp),
call_count=sum(r.call_count for r in grp),
total_tokens=sum(r.total_tokens for r in grp),
)
def _build_tool_spend_response(
rows: list[_ToolSpendRow],
total_spend: float,
start_date: str,
end_date: str,
) -> ToolSpendResponse:
daily = [
ToolSpendDailyEntry(date=r.date, tool_name=r.tool_name, spend=r.spend, call_count=r.call_count) for r in rows
]
grouped = groupby(sorted(rows, key=lambda r: r.tool_name), key=lambda r: r.tool_name)
by_tool = sorted(
(_summarize_tool(name, tuple(grp)) for name, grp in grouped),
key=lambda e: e.spend,
reverse=True,
)
return ToolSpendResponse(
by_tool=by_tool,
daily=daily,
total_spend=total_spend,
start_date=start_date,
end_date=end_date,
)
@router.get(
"/v1/tool/spend",
tags=["tool management"],
dependencies=[Depends(user_api_key_auth)],
response_model=ToolSpendResponse,
)
async def get_tool_spend(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
start_date: Annotated[str | None, Query(description="YYYY-MM-DD (defaults to 30 days ago)")] = None,
end_date: Annotated[str | None, Query(description="YYYY-MM-DD (defaults to today)")] = None,
):
"""
Spend attributed to each tool over a date range, for the Cost Optimization dashboard.
Joins ``LiteLLM_SpendLogToolIndex`` (which tool names ran on which request) to
``LiteLLM_SpendLogs`` (what the request cost). A request that used multiple tools
counts its full spend toward each of those tools, so per-tool numbers are
attributions. ``total_spend`` is the deduplicated spend of every request that
called at least one tool in the window, so it never double counts.
"""
from litellm.proxy.proxy_server import prisma_client
if user_api_key_dict.user_role not in (
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
):
raise HTTPException(
status_code=403,
detail="Only proxy admin roles can view tool spend across the deployment",
)
if prisma_client is None:
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
now = datetime.now(timezone.utc)
end_day = _parse_day_start(end_date)
start_dt = _parse_day_start(start_date) or ((end_day or now) - timedelta(days=30))
end_exclusive = (end_day + timedelta(days=1)) if end_day else now
rows = await prisma_client.db.query_raw(
"""
SELECT to_char(ti.start_time, 'YYYY-MM-DD') AS date,
ti.tool_name AS tool_name,
COUNT(*)::int AS call_count,
COALESCE(SUM(sl.spend), 0)::double precision AS spend,
COALESCE(SUM(sl.total_tokens), 0)::bigint AS total_tokens
FROM "LiteLLM_SpendLogToolIndex" ti
JOIN "LiteLLM_SpendLogs" sl ON sl.request_id = ti.request_id
WHERE ti.start_time >= ($1::timestamptz AT TIME ZONE 'UTC')
AND ti.start_time < ($2::timestamptz AT TIME ZONE 'UTC')
GROUP BY date, ti.tool_name
ORDER BY date ASC, spend DESC
""",
start_dt.isoformat(),
end_exclusive.isoformat(),
)
totals = await prisma_client.db.query_raw(
"""
SELECT COALESCE(SUM(sl.spend), 0)::double precision AS total_spend
FROM "LiteLLM_SpendLogs" sl
WHERE EXISTS (
SELECT 1
FROM "LiteLLM_SpendLogToolIndex" ti
WHERE ti.request_id = sl.request_id
AND ti.start_time >= ($1::timestamptz AT TIME ZONE 'UTC')
AND ti.start_time < ($2::timestamptz AT TIME ZONE 'UTC')
)
""",
start_dt.isoformat(),
end_exclusive.isoformat(),
)
total_rows = _REQUEST_TOTAL_ROWS.validate_python(totals or [])
return _build_tool_spend_response(
rows=_TOOL_SPEND_ROWS.validate_python(rows or []),
total_spend=total_rows[0].total_spend if total_rows else 0.0,
start_date=start_dt.strftime("%Y-%m-%d"),
end_date=(end_day or now).strftime("%Y-%m-%d"),
)
@router.get(
"/v1/tool/{tool_name:path}/detail",
tags=["tool management"],

View file

@ -36,7 +36,7 @@ if TYPE_CHECKING:
import httpx
import jwt
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
from fastapi.responses import RedirectResponse
import litellm
@ -965,15 +965,8 @@ async def google_login(
state=cli_state,
request=request,
)
if return_to is not None and sso_redirect is not None:
if SSOAuthenticationHandler._validate_return_to(return_to):
sso_redirect.set_cookie(
key="litellm_cp_return_to",
value=return_to,
max_age=600,
httponly=True,
samesite="lax",
)
if sso_redirect is not None:
_persist_return_to_cookie(sso_redirect, return_to)
return sso_redirect
from fastapi.responses import HTMLResponse
@ -982,13 +975,19 @@ async def google_login(
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
or general_settings.get("hide_default_credentials_hint", False) is True
)
return HTMLResponse(
form_response = HTMLResponse(
content=build_ui_login_form(
show_deprecation_banner=True,
hide_default_credentials_hint=hide_default_credentials_hint,
),
status_code=200,
)
# Preserve return_to across the username/password sign-in too, via the SAME shared, never-raising
# helper the SSO branch uses, so /login can resume the connect flow instead of dead-ending at the
# dashboard. One implementation → the two sign-in branches cannot diverge (and the login form always
# renders, since the helper never raises on a bad return_to).
_persist_return_to_cookie(form_response, return_to)
return form_response
def generic_response_convertor(
@ -2418,6 +2417,92 @@ async def sso_readiness():
)
def _is_same_origin_return_path(return_to: str) -> bool:
"""True for a strictly relative return path that stays on the gateway's own origin by
construction, and is therefore safe to honor without a configured ``control_plane_url``.
Used by the MCP gateway DCR authorize round-trip so a browser sent through login lands
back on the authorize request.
Requires a single leading ``/`` (not protocol-relative ``//``), no backslash (browsers
fold ``\\`` to ``/``, so ``/\\evil.com`` would escape the origin), and no control or
whitespace characters. Rejecting control chars keeps a ``\\r\\n``/tab-bearing value out
of the redirect ``Location`` and the ``litellm_cp_return_to`` cookie entirely, rather
than relying on downstream header encoding to neutralize it."""
if not return_to.startswith("/") or return_to.startswith("//") or "\\" in return_to:
return False
return not any(ord(ch) < 0x20 or ch in (" ", "\x7f") for ch in return_to)
async def _sso_return_to_redirect(
return_to: str | None,
jwt_token: str,
redis_usage_cache,
user_api_key_cache,
) -> RedirectResponse | None:
"""Resolve the post-SSO redirect for a ``return_to``, or None to fall through to the dashboard.
Two arms, both clearing the one-shot ``litellm_cp_return_to`` cookie:
- **Same-origin relative path** (the MCP gateway DCR authorize round-trip): set the session cookie
exactly like the dashboard path, then send the browser back where it came from.
- **Control-plane cross-origin** (``control_plane_url``): stash the JWT behind a single-use opaque
code (60s TTL) so the token never lands in browser history/logs; the control plane redeems it via
``POST /v3/login/exchange``.
Extracted from ``get_redirect_response_from_openid`` to keep that method inside the complexity
budget; behavior is identical to the inline arms it replaces (including letting
``_validate_return_to`` raise for a mismatched absolute return_to, as before)."""
if return_to is None:
return None
if _is_same_origin_return_path(return_to):
redirect_response = RedirectResponse(url=return_to, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
if SSOAuthenticationHandler._validate_return_to(return_to):
code = secrets.token_urlsafe(32)
cache_key = f"login_code:{code}"
cache_value = {"token": jwt_token, "redirect_url": return_to}
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
else:
await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
separator = "&" if "?" in return_to else "?"
redirect_url = return_to + separator + urlencode({"login": "success", "code": code})
verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code")
redirect_response = RedirectResponse(url=redirect_url, status_code=303)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
return None
def _persist_return_to_cookie(response: Response, return_to: str | None) -> None:
"""Best-effort: persist a SAFE ``return_to`` on ``response`` as the one-shot ``litellm_cp_return_to``
cookie so ANY sign-in path — SSO / Okta / generic OR the username/password form — can resume there
afterwards. THIS is the single source of truth, called by every sign-in branch so they cannot
diverge (a per-branch reimplementation is exactly how the two drifted before). Honors a strictly
relative same-origin path, and (when ``control_plane_url`` is configured) a return_to matching that
origin. It NEVER raises: a mismatched or invalid ``return_to`` is simply not stored, so it can never
block sign-in — the login entrypoint must always render."""
if return_to is None:
return
try:
safe = _is_same_origin_return_path(return_to) or SSOAuthenticationHandler._validate_return_to(return_to)
except HTTPException:
return # a non-matching absolute return_to is ignored, never blocks sign-in
if safe:
response.set_cookie(
key="litellm_cp_return_to",
value=return_to,
max_age=600,
httponly=True,
samesite="lax",
)
class SSOAuthenticationHandler:
"""
Handler for SSO Authentication across all SSO providers
@ -3055,7 +3140,6 @@ class SSOAuthenticationHandler:
return_to: Optional[str] = None,
sso_assertion: SSOIdentityAssertion | None = None,
) -> RedirectResponse:
import jwt
from litellm.proxy.proxy_server import (
general_settings,
@ -3219,30 +3303,21 @@ class SSOAuthenticationHandler:
server_root_path=get_server_root_path(),
)
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
master_key or "",
algorithm="HS256",
from litellm.proxy.auth.login_utils import encode_ui_session_jwt
jwt_token = encode_ui_session_jwt(returned_ui_token_object, master_key or "")
# Post-SSO return_to handling (the same-origin DCR round-trip and the control-plane
# cross-origin code exchange) lives in one shared helper so this method stays inside the
# complexity budget. None falls through to the dashboard redirect below.
return_to_redirect = await _sso_return_to_redirect(
return_to=return_to,
jwt_token=jwt_token,
redis_usage_cache=redis_usage_cache,
user_api_key_cache=user_api_key_cache,
)
# Control-plane cross-origin: store JWT behind a single-use opaque
# code (60s TTL) so the token never appears in browser history / logs.
# The control plane redeems it via POST /v3/login/exchange.
if return_to is not None and SSOAuthenticationHandler._validate_return_to(return_to):
code = secrets.token_urlsafe(32)
cache_key = f"login_code:{code}"
cache_value = {"token": jwt_token, "redirect_url": return_to}
if redis_usage_cache is not None:
await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
else:
await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60)
separator = "&" if "?" in return_to else "?"
redirect_url = return_to + separator + urlencode({"login": "success", "code": code})
verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code")
redirect_response = RedirectResponse(url=redirect_url, status_code=303)
redirect_response.delete_cookie("litellm_cp_return_to")
return redirect_response
if return_to_redirect is not None:
return return_to_redirect
if user_id is not None and isinstance(user_id, str):
litellm_dashboard_ui += "?login=success"

View file

@ -13472,7 +13472,7 @@ async def fallback_login(request: Request):
@router.post("/login", include_in_schema=False) # hidden since this is a helper for UI sso login
async def login(request: Request):
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt
from litellm.proxy.utils import get_custom_url
form = await request.form()
@ -13495,13 +13495,7 @@ async def login(request: Request):
)
# Generate JWT token
import jwt
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
cast(str, master_key),
algorithm="HS256",
)
jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key))
# Build redirect URL
litellm_dashboard_ui = get_custom_url(str(request.base_url))
@ -13511,16 +13505,51 @@ async def login(request: Request):
litellm_dashboard_ui += "/ui/"
litellm_dashboard_ui += "?login=success"
# Honor a same-origin return_to preserved by the sign-in page (e.g. the aggregate DCR connect flow's
# authorize round-trip), mirroring the SSO callback; otherwise land on the dashboard. Gated by
# _is_same_origin_return_path (strictly relative path) so it can never be an open redirect, and the
# one-shot cookie is cleared after use.
from litellm.proxy.management_endpoints.ui_sso import _sso_return_to_redirect
# Resume through the SAME resumer the SSO callback uses, rather than a second, narrower arm.
# _persist_return_to_cookie stores both shapes it accepts (a relative same-origin path AND a
# control_plane_url-matching absolute URL); honoring only the relative one here silently dropped
# the control-plane case, landing the user on the dashboard. One function decides how a stored
# return_to is honored for EVERY sign-in branch, so the write and read sets cannot diverge: it
# sets the token cookie on the same-origin arm and hands off via a one-time login code on the
# cross-origin arm, and clears the one-shot cookie in both.
cp_return_to = request.cookies.get("litellm_cp_return_to")
if cp_return_to:
try:
resumed = await _sso_return_to_redirect(
return_to=cp_return_to,
jwt_token=jwt_token,
redis_usage_cache=redis_usage_cache,
user_api_key_cache=user_api_key_cache,
)
except Exception: # noqa: BLE001 # resuming must NEVER block a completed sign-in
# The symmetric half of _persist_return_to_cookie's "never raises" contract. The resumer
# rejects a return_to that no longer matches control_plane_url (a config change between
# the cookie's write and this read), and the user has ALREADY authenticated here —
# failing their login over a stale one-shot cookie is the worst possible outcome. Land
# on the dashboard instead; the cookie is cleared below either way.
verbose_proxy_logger.info("Ignoring stale litellm_cp_return_to cookie; landing on dashboard")
resumed = None
if resumed is not None:
return resumed
# Create redirect response with cookie
redirect_response = RedirectResponse(url=litellm_dashboard_ui, status_code=303)
redirect_response.set_cookie(key="token", value=jwt_token)
if cp_return_to:
redirect_response.delete_cookie(key="litellm_cp_return_to")
return redirect_response
@router.post("/v2/login", include_in_schema=False) # hidden helper for UI logins via API
async def login_v2(request: Request):
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt
from litellm.proxy.utils import get_custom_url
try:
@ -13541,13 +13570,7 @@ async def login_v2(request: Request):
premium_user=premium_user,
)
import jwt
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
cast(str, master_key),
algorithm="HS256",
)
jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key))
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):
@ -13591,7 +13614,7 @@ async def login_v2(request: Request):
) # control-plane login — always returns token in body for cross-origin use
async def login_v3(request: Request):
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object, encode_ui_session_jwt
from litellm.proxy.utils import get_custom_url
try:
@ -13620,13 +13643,7 @@ async def login_v3(request: Request):
premium_user=premium_user,
)
import jwt
jwt_token = jwt.encode(
cast(dict, returned_ui_token_object),
cast(str, master_key),
algorithm="HS256",
)
jwt_token = encode_ui_session_jwt(returned_ui_token_object, cast(str, master_key))
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):

View file

@ -4,7 +4,9 @@ import asyncio
import json
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Mapping, Optional, Sequence, cast
from typing import Any, Dict, List, Mapping, NoReturn, Optional, Sequence, cast
from fastapi import HTTPException, status
import litellm
from litellm._logging import verbose_proxy_logger
@ -59,6 +61,22 @@ class _CounterReservationUnavailable(Exception):
super().__init__("Counter reservation unavailable")
def _raise_reservation_unavailable(counter_key: str) -> NoReturn:
verbose_proxy_logger.warning(
"fail_closed_budget_enforcement: rejecting request — budget reservation for %s could not be written",
counter_key,
)
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=(
"Budget enforcement unavailable: the budget reservation could not "
"be written to the spend counter backend, and "
"fail_closed_budget_enforcement is enabled, so the request was "
"rejected to avoid exceeding the configured budget. Retry shortly."
),
)
def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set:
if not budget_reservation:
return set()
@ -138,6 +156,7 @@ async def reserve_budget_for_request(
end_user_id: Optional[str] = None,
end_user_object: Optional[Any] = None,
skip_user_budget_on_team_key: bool = False,
fail_closed_budget_enforcement: bool = False,
) -> Optional[dict]:
if valid_token is None or not RouteChecks.is_llm_api_route(route=route):
return None
@ -193,6 +212,8 @@ async def reserve_budget_for_request(
default_reserved_cost=reservation_cost,
)
applied_entries.remove(entry)
if fail_closed_budget_enforcement:
_raise_reservation_unavailable(counter_key=counter.counter_key)
continue
if reserved_value is not None:

View file

@ -1695,13 +1695,10 @@ async def ui_view_spend_logs(
code=status.HTTP_401_UNAUTHORIZED,
)
if start_date is None or end_date is None:
raise ProxyException(
message="Start date and end date are required",
type="bad_request",
param="None",
code=status.HTTP_400_BAD_REQUEST,
)
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
is_v2 = "/spend/logs/v2" in get_request_route(request)
# Validate sort_by and sort_order
valid_sort_fields = {
@ -1729,36 +1726,50 @@ async def ui_view_spend_logs(
)
try:
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
is_admin_view = _is_admin_view_safe(user_api_key_dict=user_api_key_dict)
is_request_id_lookup = request_id is not None and not is_v2
is_v2 = "/spend/logs/v2" in get_request_route(request)
formats = ["%Y-%m-%d %H:%M:%S", "%Y-%m-%d"] if is_v2 else ["%Y-%m-%d %H:%M:%S"]
if is_request_id_lookup:
# request_id is the @id primary key: it identifies a single row, so a
# time window is meaningless. The dashboard always sends a default 24h
# window, which hid ids copied from an older page (LIT-3981). Drop the
# window for the id lookup so it resolves across all time; every other
# query, including the public v2 route, still requires one (below).
start_date_obj: datetime | None = None
end_date_obj: datetime | None = None
else:
if start_date is None or end_date is None:
raise ProxyException(
message="Start date and end date are required",
type="bad_request",
param="None",
code=status.HTTP_400_BAD_REQUEST,
)
formats = ["%Y-%m-%d %H:%M:%S", "%Y-%m-%d"] if is_v2 else ["%Y-%m-%d %H:%M:%S"]
def parse_date(date_str: str) -> datetime:
date_str = date_str.strip()
for fmt in formats:
try:
return datetime.strptime(date_str, fmt).replace(tzinfo=timezone.utc)
except ValueError:
continue
expected = "'YYYY-MM-DD' or 'YYYY-MM-DD HH:MM:SS'" if is_v2 else "'YYYY-MM-DD HH:MM:SS'"
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Invalid date format: {date_str}. Expected: {expected}",
)
def parse_date(date_str: str) -> datetime:
date_str = date_str.strip()
for fmt in formats:
try:
return datetime.strptime(date_str, fmt).replace(tzinfo=timezone.utc)
except ValueError:
continue
expected = "'YYYY-MM-DD' or 'YYYY-MM-DD HH:MM:SS'" if is_v2 else "'YYYY-MM-DD HH:MM:SS'"
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Invalid date format: {date_str}. Expected: {expected}",
)
start_date_obj = parse_date(start_date)
end_date_obj = parse_date(end_date)
# Convert to ISO format strings for Prisma
start_date_iso = start_date_obj.isoformat() # Already in UTC, no need to add Z
end_date_iso = end_date_obj.isoformat() # Already in UTC, no need to add Z
start_date_obj = parse_date(start_date)
end_date_obj = parse_date(end_date)
# Build where conditions
where_conditions: dict[str, Any] = {
"startTime": {"gte": start_date_iso, "lte": end_date_iso},
}
where_conditions: dict[str, Any] = {}
if start_date_obj is not None and end_date_obj is not None:
where_conditions["startTime"] = {
"gte": start_date_obj.isoformat(), # Already in UTC, no need to add Z
"lte": end_date_obj.isoformat(),
}
if team_id is not None:
where_conditions["team_id"] = team_id
@ -1827,9 +1838,19 @@ async def ui_view_spend_logs(
where_conditions["spend"]["gte"] = min_spend
if max_spend is not None:
where_conditions["spend"]["lte"] = max_spend
is_admin_view = _is_admin_view_safe(user_api_key_dict=user_api_key_dict)
# A request_id lookup drops the date window, so a non-admin could otherwise
# reach any single row by id; require they own it, mirroring the detail
# endpoint. That ownership check fully authorizes the one row, so the
# general scoping below is skipped for id lookups. Scoped to the UI route
# so the public v2 contract is unchanged.
if request_id is not None and not is_v2 and not is_admin_view:
await _assert_user_can_view_request_id(
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
request_id=request_id,
)
permitted_team_ids: List[str] | None = None
if not is_admin_view:
if not is_request_id_lookup and not is_admin_view:
if team_id is not None:
can_view_team = await _can_team_member_view_log(
prisma_client=prisma_client,
@ -1875,15 +1896,16 @@ async def ui_view_spend_logs(
sql_params: List[Any] = []
p = 1 # parameter index counter
# Date range (always present). Wrap the param side with
# `AT TIME ZONE 'UTC'` so comparison against the plain `timestamp`
# column does not depend on the DB session timezone (see #22529).
sql_conditions.append(f"\"startTime\" >= (${p}::timestamptz AT TIME ZONE 'UTC')")
sql_params.append(start_date_obj)
p += 1
sql_conditions.append(f"\"startTime\" <= (${p}::timestamptz AT TIME ZONE 'UTC')")
sql_params.append(end_date_obj)
p += 1
# Date range. Wrap the param side with `AT TIME ZONE 'UTC'` so comparison
# against the plain `timestamp` column does not depend on the DB session
# timezone (see #22529). Absent for a request_id-only lookup (see above).
if start_date_obj is not None and end_date_obj is not None:
sql_conditions.append(f"\"startTime\" >= (${p}::timestamptz AT TIME ZONE 'UTC')")
sql_params.append(start_date_obj)
p += 1
sql_conditions.append(f"\"startTime\" <= (${p}::timestamptz AT TIME ZONE 'UTC')")
sql_params.append(end_date_obj)
p += 1
# Equality filters - read effective values from where_conditions (post-authorization)
for sql_col, wc_key in [

View file

@ -374,12 +374,22 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
if isinstance(v, BaseModel):
v = v.model_dump()
additional_usage_values.update({k: v})
if "cache_read_input_tokens" not in additional_usage_values:
prompt_tokens_details = additional_usage_values.get("prompt_tokens_details")
if isinstance(prompt_tokens_details, dict):
prompt_tokens_details = additional_usage_values.get("prompt_tokens_details")
if not isinstance(prompt_tokens_details, dict):
usage_object = clean_metadata.get("usage_object")
if isinstance(usage_object, dict):
prompt_tokens_details = usage_object.get("prompt_tokens_details")
if isinstance(prompt_tokens_details, dict):
if "cache_read_input_tokens" not in additional_usage_values:
cached_tokens = prompt_tokens_details.get("cached_tokens")
if isinstance(cached_tokens, int) and cached_tokens > 0:
additional_usage_values["cache_read_input_tokens"] = cached_tokens
if "cache_creation_input_tokens" not in additional_usage_values:
cache_write_tokens = prompt_tokens_details.get("cache_write_tokens") or prompt_tokens_details.get(
"cache_creation_tokens"
)
if isinstance(cache_write_tokens, int) and cache_write_tokens > 0:
additional_usage_values["cache_creation_input_tokens"] = cache_write_tokens
clean_metadata["additional_usage_values"] = additional_usage_values
if litellm.cache is not None:

View file

@ -1042,21 +1042,14 @@ class ResponseAPILoggingUtils:
prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None
if response_api_usage.input_tokens_details:
if isinstance(response_api_usage.input_tokens_details, dict):
input_tokens_details = dict(response_api_usage.input_tokens_details)
cache_write_tokens = input_tokens_details.pop("cache_write_tokens", None)
if input_tokens_details.get("cache_creation_tokens") is None and cache_write_tokens is not None:
input_tokens_details["cache_creation_tokens"] = cache_write_tokens
prompt_tokens_details = PromptTokensDetailsWrapper(**input_tokens_details)
prompt_tokens_details = PromptTokensDetailsWrapper(**response_api_usage.input_tokens_details)
else:
prompt_tokens_details = PromptTokensDetailsWrapper(
cached_tokens=getattr(response_api_usage.input_tokens_details, "cached_tokens", None),
audio_tokens=getattr(response_api_usage.input_tokens_details, "audio_tokens", None),
text_tokens=getattr(response_api_usage.input_tokens_details, "text_tokens", None),
image_tokens=getattr(response_api_usage.input_tokens_details, "image_tokens", None),
cache_creation_tokens=getattr(
response_api_usage.input_tokens_details, "cache_creation_tokens", None
)
or getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
)
completion_tokens_details: Optional[CompletionTokensDetailsWrapper] = None
output_tokens_details = getattr(response_api_usage, "output_tokens_details", None)

View file

@ -261,12 +261,3 @@ class MCPServer(BaseModel):
if self.oauth_passthrough is not True:
return False
return any(h.lower() == "authorization" for h in self.extra_headers)
@property
def has_token_exchange_config(self) -> bool:
"""True if this server is configured for OAuth2 token exchange (OBO / RFC 8693)."""
return (
self.auth_type == MCPAuth.oauth2_token_exchange
and bool(self.client_id and self.client_secret)
and bool(self.token_exchange_endpoint or self.token_url)
)

View file

@ -98,3 +98,38 @@ class ToolUsageLogsResponse(BaseModel):
total: int
page: int
page_size: int
class ToolSpendEntry(BaseModel):
"""Total spend attributed to one tool over the requested window."""
tool_name: str
spend: float = Field(
0.0,
description="Attributed spend: a request that used several tools counts its full spend toward each of them",
)
call_count: int = 0
total_tokens: int = 0
class ToolSpendDailyEntry(BaseModel):
"""Spend attributed to one tool on one UTC day."""
date: str
tool_name: str
spend: float = 0.0
call_count: int = 0
class ToolSpendResponse(BaseModel):
by_tool: List[ToolSpendEntry] = Field(default_factory=list)
daily: List[ToolSpendDailyEntry] = Field(default_factory=list)
total_spend: float = Field(
0.0,
description=(
"Deduplicated spend of every request that called at least one tool in the window; "
"less than the sum of per-tool attributed spend whenever multi-tool requests exist"
),
)
start_date: str | None = None
end_date: str | None = None

View file

@ -1537,14 +1537,27 @@ class PromptTokensDetailsWrapper(
audio_length_seconds: Optional[float] = None
"""Length of audio sent to the model. Used for multimodal embeddings priced per audio-second."""
cache_write_tokens: Optional[int] = None
"""Number of cache write (creation) tokens sent to the model. OpenAI naming (prompt_tokens_details.cache_write_tokens); this is the canonical field."""
cache_creation_tokens: Optional[int] = None
"""Number of cache creation tokens sent to the model. Used for Anthropic prompt caching."""
"""Number of cache creation tokens sent to the model. Anthropic/Bedrock naming; kept in sync with cache_write_tokens (assigning either mirrors to the other)."""
cache_creation_token_details: Optional[CacheCreationTokenDetails] = None
"""Details of cache creation tokens sent to the model. Used for tracking 5m/1h cache creation tokens for Anthropic prompt caching."""
def __setattr__(self, name: str, value: object) -> None:
super().__setattr__(name, value)
if name == "cache_write_tokens":
super().__setattr__("cache_creation_tokens", value)
elif name == "cache_creation_tokens":
super().__setattr__("cache_write_tokens", value)
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.cache_write_tokens = (
self.cache_write_tokens if self.cache_write_tokens is not None else self.cache_creation_tokens
)
if self.character_count is None:
del self.character_count
if self.image_count is None:
@ -1557,6 +1570,8 @@ class PromptTokensDetailsWrapper(
del self.web_search_requests
if self.tool_use_tokens is None:
del self.tool_use_tokens
if self.cache_write_tokens is None:
del self.cache_write_tokens
if self.cache_creation_tokens is None:
del self.cache_creation_tokens
if self.cache_creation_token_details is None:
@ -1665,10 +1680,10 @@ class Usage(SafeAttributeModel, CompletionUsage):
if "cache_creation_input_tokens" in params and isinstance(params["cache_creation_input_tokens"], int):
if _prompt_tokens_details is None:
_prompt_tokens_details = PromptTokensDetailsWrapper(
cache_creation_tokens=params["cache_creation_input_tokens"]
cache_write_tokens=params["cache_creation_input_tokens"]
)
else:
_prompt_tokens_details.cache_creation_tokens = params["cache_creation_input_tokens"]
_prompt_tokens_details.cache_write_tokens = params["cache_creation_input_tokens"]
super().__init__(
prompt_tokens=prompt_tokens or 0,

View file

@ -198,6 +198,7 @@ e2e-dev = [
"playwright==1.61.0",
"websockets>=15.0.1,<16.0",
"locust==2.45.0",
"mcp>=1.28.1,<2.0",
]
proxy-dev = [
"prisma==0.11.0",

View file

@ -2,8 +2,10 @@
raw HTTP client imports (requests, urllib.request, httpx, aiohttp, http.client) are
banned in suite code. Importing requests' exception types for catching is fine
anywhere; a small allowlist grandfathers the files that legitimately make raw calls
(the transport itself, the root conftest liveness probe, and the claude_code version
resolver's constant registry URL fetch). Referenced by tests/e2e/CLAUDE.md."""
(the transport itself, the root conftest liveness probe, the claude_code version
resolver's constant registry URL fetch, and the mcp OAuth client, whose httpx
client is the object the official mcp SDK's streamable_http_client requires and so
cannot go through the sync requests transport). Referenced by tests/e2e/CLAUDE.md."""
from __future__ import annotations
@ -19,6 +21,7 @@ ALLOWED_RAW_CLIENT_FILES = {
"e2e_http.py": ("requests",),
"conftest.py": ("requests",),
"claude_code/pr_gate_version_resolver.py": ("urllib.request",),
"mcp/oauth_chat_client.py": ("httpx",),
}
EXCEPTION_ONLY_NAMES = frozenset({"RequestException", "ConnectionError", "Timeout", "HTTPError"})

View file

@ -14,7 +14,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`)
- `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials (API surface; not Playwright)
- `a2a/` - the A2A (agent-to-agent) surface: admin registration via `/v1/agents`, proxy-fronted card discovery at `/.well-known/agent-card.json`, and JSON-RPC `message/send` invocation, driving agents backed by the litellm completion bridge (a real provider) and asserting protocol-version normalization (0.3 vs 1.0)
- `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server only (see "MCP suite: real Datadog only" below)
- `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server (see "MCP suite: real Datadog only" below); plus the gateway-managed OAuth (authorization_code) path exercised through `/chat/completions`, the one behavior Datadog's static-header auth cannot reach, seeding the per-user upstream token via the interactive authorize dance driven with the mcp SDK's own OAuth client (headless-browser consent from a saved session) and asserting the completion lists and executes the server's tools with the stored per-user token
- `logging/` - logging-integration delivery (datadog and friends)
- `security/` - secret handling and log-leak protection
- `router/` - routing and reliability behavior (fallbacks, cooldowns)
@ -33,6 +33,7 @@ Every test under `tests/e2e/mcp/` must exercise the proxy against the real Datad
- Prefer calling real Datadog tools that prove the product path (e.g. `search_datadog_logs` for list/call and permission denials). Seed a unique marker (`e2e-datadog-mcp-*`) in a chat completion when you need a log the tool can find; dual-read with `dd_logs` from conftest when delivery matters
- Delete the MCP server (and any keys) through `resources.defer` the same way every other suite tears down
- If a new MCP behavior cannot be covered with Datadog's tool surface, say so in the PR and get agreement before inventing another upstream; the default is always Datadog
- The one standing exception is `test_mcp_chat_completion_oauth_e2e.py`. Datadog authenticates with the static `DD-API-KEY` / `DD-APPLICATION-KEY` headers and exposes no authorize/token dance at all, so it cannot exercise gateway-managed OAuth or per-user token seeding in any form. That test drives a real Linear MCP server instead; it is still a real remote upstream, so the no-mock, no-fixture rule above holds unchanged
## Lay the pattern down in a class

View file

@ -23,7 +23,7 @@ from typing import Callable
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from batch_client import (
UPLOAD_FILENAME,
@ -702,14 +702,7 @@ class TestBedrockBatchAssumeRole:
def test_unified_batch_create_with_assume_role(
self, client: BatchClient, resources: ResourceManager
) -> None:
(role_arn,) = require_env("AWS_ROLE_NAME")
require_env(
"AWS_ACCESS_KEY_ID",
"AWS_SECRET_ACCESS_KEY",
"AWS_REGION",
"AWS_BATCH_S3_BUCKET",
"AWS_BATCH_ROLE_ARN",
)
role_arn = os.environ["AWS_ROLE_NAME"]
session_name = f"e2e-batch-sts-{unique_marker()}"[:64]
model_name = batch_model_name("bedrock-sts-batch")
@ -819,7 +812,7 @@ class TestHostedVllmBatch:
def test_unified_file_and_batch_create(
self, client: BatchClient, resources: ResourceManager
) -> None:
(api_base,) = require_env("HOSTED_VLLM_API_BASE")
api_base = os.environ["HOSTED_VLLM_API_BASE"]
api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None
model_id = (
os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct"

View file

@ -38,6 +38,9 @@ UI_BASE_URL = os.environ.get("E2E_UI_BASE_URL", PROXY_BASE_URL).rstrip("/")
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")
LINEAR_MCP_URL = os.environ.get("E2E_LINEAR_MCP_URL", "https://mcp.linear.app/mcp")
LINEAR_STORAGE_STATE = os.environ.get("E2E_LINEAR_STORAGE_STATE", "")
# 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
# read exported spans back through it.
@ -99,22 +102,6 @@ ANOMALY_SPEND_SETTLE_SECONDS = float(
)
def require_env(*names: str) -> tuple[str, ...]:
"""Return the non-empty values for each env name, or hard-fail naming which are missing.
Live e2e never skips for missing credentials: a missing key is a red run so
ops knows the suite cannot prove the product path.
"""
missing = tuple(name for name in names if not (os.environ.get(name) or "").strip())
if missing:
joined = ", ".join(missing)
raise AssertionError(
f"missing required env for e2e: {joined}. "
"Add them to tests/e2e/.env locally and to litellm ops for stage/CI."
)
return tuple((os.environ.get(name) or "").strip() for name in names)
def datadog_mcp_url(*, toolsets: str = "core") -> str:
"""Regional Datadog remote MCP endpoint for this process's DD_SITE.

View file

@ -8,9 +8,11 @@ a 200 means the guardrail never ran.
from __future__ import annotations
import os
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import UnknownApiError
from guardrails_client import GuardrailsClient
from lifecycle import ResourceManager
@ -33,11 +35,8 @@ class TestBedrockGuardrail:
def test_bedrock_pre_call_blocks_harmful_prompt(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
(identifier, version) = require_env(
"BEDROCK_GUARDRAIL_IDENTIFIER",
"BEDROCK_GUARDRAIL_VERSION",
)
require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION")
identifier = os.environ["BEDROCK_GUARDRAIL_IDENTIFIER"]
version = os.environ["BEDROCK_GUARDRAIL_VERSION"]
name = f"e2e-bedrock-guard-{unique_marker()}"
guardrail_id = client.create_bedrock_guardrail(

View file

@ -16,7 +16,7 @@ from __future__ import annotations
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import unwrap
from guardrails_client import BlockCodeExecutionParamsBody, GuardrailsClient
from lifecycle import ResourceManager
@ -46,7 +46,6 @@ class TestBlockCodeExecutionGuardrail:
def test_blocks_execution_request_but_allows_explanation(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-blockcode-backend")
name = f"e2e-block-code-{unique_marker()}"

View file

@ -14,7 +14,7 @@ from __future__ import annotations
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import UnknownApiError, unwrap
from guardrails_client import GuardrailsClient, OpenAIModerationParamsBody
from lifecycle import ResourceManager
@ -34,7 +34,6 @@ class TestOpenAIModerationGuardrail:
def test_moderation_blocks_flagged_input(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("OPENAI_API_KEY", "GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-moderation-backend")
name = f"e2e-openai-moderation-{unique_marker()}"

View file

@ -25,11 +25,12 @@ The chat backend is a gemini deployment created for the test.
from __future__ import annotations
import os
import time
import pytest
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, require_env, unique_marker
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker
from e2e_http import NoBody, require_successful_call, unwrap
from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody
from lifecycle import ResourceManager
@ -88,9 +89,8 @@ def _poll_logged_prompt(reader: OtelReader, *, call_id: str, genai_span: str) ->
def _presidio_params(
mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False
) -> PresidioParamsBody:
analyzer, anonymizer = require_env(
"PRESIDIO_ANALYZER_API_BASE", "PRESIDIO_ANONYMIZER_API_BASE"
)
analyzer = os.environ["PRESIDIO_ANALYZER_API_BASE"]
anonymizer = os.environ["PRESIDIO_ANONYMIZER_API_BASE"]
return PresidioParamsBody(
mode=mode,
default_on=False,
@ -124,7 +124,6 @@ class TestPresidioGuardrail:
def test_pre_call_masks_pii_before_the_model_sees_it(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-presidio-pre")
name = f"e2e-presidio-pre-{unique_marker()}"
guardrail_id = client.register(name, _presidio_params("pre_call"))
@ -149,7 +148,6 @@ class TestPresidioGuardrail:
def test_post_call_masks_pii_in_model_output(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
model = client.create_backend_model(resources, prefix="e2e-presidio-post")
name = f"e2e-presidio-post-{unique_marker()}"
guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True))
@ -173,7 +171,6 @@ class TestPresidioGuardrail:
def test_logging_only_masks_the_logged_prompt(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
require_env("GEMINI_API_KEY")
_require_otel_v2_active(client)
reader = build_otel_reader()

View file

@ -21,7 +21,7 @@ import os
import pytest
from pydantic import BaseModel
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import StreamingResponse, unwrap
from lifecycle import ResourceManager
from models import (
@ -250,7 +250,7 @@ class TestCohereChat:
def test_cohere_chat_returns_content(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
(cohere_key,) = require_env("COHERE_API_KEY")
cohere_key = os.environ["COHERE_API_KEY"]
model = f"e2e-cohere-chat-{unique_marker()}"
model_id = client.proxy.create_model(
model,
@ -343,7 +343,7 @@ class TestHostedVllmChat:
def test_hosted_vllm_chat_returns_content(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
(api_base,) = require_env("HOSTED_VLLM_API_BASE")
api_base = os.environ["HOSTED_VLLM_API_BASE"]
api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None
backend = (
os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct"
@ -395,7 +395,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_streams_real_content(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-chat-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -423,7 +422,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_logs_cost(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-cost-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -457,7 +455,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_returns_tool_call(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-tool-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -488,7 +485,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_structured_output_conforms_to_schema(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-schema-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -522,7 +518,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_reasoning_reports_reasoning_tokens(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-reasoning-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -561,7 +556,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_vision_describes_image(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-vision-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_VISION_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -579,7 +573,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_prompt_cache_hits_on_repeat(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-cache-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -610,7 +603,6 @@ class TestOpenAIChatCompletions:
def test_openai_chat_streams_tool_call(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
require_env("OPENAI_API_KEY")
model = f"e2e-openai-tool-stream-{unique_marker()}"
model_id = client.proxy.create_model(
model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY")
@ -645,7 +637,6 @@ class TestBedrockConverseChatCompletions:
"""
def _register(self, client: PassthroughClient, resources: ResourceManager, prefix: str) -> str:
require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION")
model = f"{prefix}-{unique_marker()}"
model_id = client.proxy.create_model(model, _bedrock_params())
resources.defer(lambda: client.proxy.delete_model(model_id))

View file

@ -9,7 +9,7 @@ from __future__ import annotations
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import require_successful_call
from endpoints_client import EndpointsClient, ImagesResult
from lifecycle import ResourceManager
@ -50,7 +50,6 @@ class TestImageGeneration:
def test_bedrock_image_generation_returns_image(
self, endpoints_client: EndpointsClient, resources: ResourceManager
) -> None:
require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION")
model = f"e2e-bedrock-image-{unique_marker()}"
model_id = endpoints_client.create_model(
model,

View file

@ -10,7 +10,7 @@ from __future__ import annotations
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import require_successful_call, unwrap
from endpoints_client import EndpointsClient, MessagesResult
from lifecycle import ResourceManager
@ -73,7 +73,6 @@ class TestAnthropicMessages:
def test_messages_logs_cost_matching_the_response_header(
self, endpoints_client: EndpointsClient, resources: ResourceManager
) -> None:
require_env("ANTHROPIC_API_KEY")
model = f"e2e-messages-cost-{unique_marker()}"
model_id = endpoints_client.create_model(
model,

View file

@ -9,7 +9,7 @@ from __future__ import annotations
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import require_successful_call
from endpoints_client import EndpointsClient, RerankResult
from lifecycle import ResourceManager
@ -56,7 +56,6 @@ class TestRerank:
def test_bedrock_rerank_scores_top_n(
self, endpoints_client: EndpointsClient, resources: ResourceManager
) -> None:
require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION")
model = f"e2e-bedrock-rerank-{unique_marker()}"
model_id = endpoints_client.create_model(
model,

View file

@ -13,7 +13,7 @@ from typing import cast
import pytest
from pydantic import BaseModel, ValidationError
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import require_successful_call
from endpoints_client import (
EndpointsClient,
@ -255,7 +255,6 @@ class TestResponses:
def test_responses_bedrock_returns_completion(
self, endpoints_client: EndpointsClient, resources: ResourceManager
) -> None:
require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION")
model = f"e2e-responses-{unique_marker()}"
model_id = endpoints_client.create_model(model, _bedrock_params())
resources.defer(lambda: endpoints_client.delete_model(model_id))
@ -270,7 +269,6 @@ class TestResponses:
def test_responses_bedrock_returns_function_call(
self, endpoints_client: EndpointsClient, resources: ResourceManager
) -> None:
require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION")
model = f"e2e-responses-{unique_marker()}"
model_id = endpoints_client.create_model(model, _bedrock_params())
resources.defer(lambda: endpoints_client.delete_model(model_id))

View file

@ -14,7 +14,7 @@ import time
import pytest
from pydantic import BaseModel, ConfigDict
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import require_successful_call
from endpoints_client import EndpointsClient, ResponsesResult
from lifecycle import ResourceManager
@ -42,7 +42,7 @@ class RedisKeyInfo(BaseModel):
def _redis_scan(marker: str) -> tuple[RedisKeyInfo, ...]:
import redis
(host,) = require_env("REDIS_HOST")
host = os.environ["REDIS_HOST"]
port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379")
try:
with socket.create_connection((host, port), timeout=3):

View file

@ -0,0 +1,58 @@
"""One-time helper to capture a logged-in Linear browser session for the
real-Linear MCP e2e test.
The real-Linear test drives the genuine gateway-managed authorization_code
dance against ``mcp.linear.app``. The only step that cannot be scripted is
Linear's login (magic link / SSO), so a human authenticates once here and the
resulting session (cookies + local storage) is persisted to disk. The e2e test
then loads that session in a headless Playwright context and clicks Approve on
Linear's consent screen every run, with no human and no login automation.
Run it with the e2e venv, log into Linear in the window that opens, then return
to the terminal and press Enter:
LITELLM=~/litellm-mcpe2e
"$LITELLM"/.venv/bin/python "$LITELLM"/tests/e2e/mcp/linear_session_capture.py
The session is written to ``E2E_LINEAR_STORAGE_STATE`` (default
``~/.litellm-e2e/linear_storage_state.json``), outside the repo. It is a
secret: never commit it. Re-run this whenever Linear expires the session.
"""
from __future__ import annotations
import os
from pathlib import Path
from playwright.sync_api import sync_playwright
DEFAULT_STATE_PATH = Path.home() / ".litellm-e2e" / "linear_storage_state.json"
def capture(state_path: Path) -> None:
"""Open a headed browser at Linear, wait for the human to log in, then save
the authenticated session to ``state_path``."""
state_path.parent.mkdir(parents=True, exist_ok=True)
with sync_playwright() as playwright:
browser = playwright.chromium.launch(headless=False)
context = browser.new_context()
page = context.new_page()
page.goto("https://linear.app/login", wait_until="domcontentloaded")
print("\n" + "=" * 72)
print("Log into Linear in the browser window that just opened.")
print("If Linear emails you a magic link, paste the link into THIS window's")
print("address bar (opening it in your default browser won't capture the")
print("session). Google SSO works too as long as you complete it here.")
print("When your Linear workspace has loaded, come back and press Enter.")
print("=" * 72)
input("Press Enter once you are logged in... ")
page.goto("https://mcp.linear.app/", wait_until="domcontentloaded")
context.storage_state(path=str(state_path))
browser.close()
print(f"\nSaved Linear session to {state_path}")
print("Point the e2e test at it with:")
print(f' export E2E_LINEAR_STORAGE_STATE="{state_path}"')
if __name__ == "__main__":
capture(Path(os.environ.get("E2E_LINEAR_STORAGE_STATE", str(DEFAULT_STATE_PATH))))

View file

@ -0,0 +1,271 @@
"""Client for the mcp chat-completion OAuth e2e suite.
Registers a gateway-managed OAuth (authorization_code) MCP server, seeds the
per-user upstream token by driving the interactive authorize dance with the
official mcp SDK's OAuthClientProvider (the browser leg is a headless Chromium
primed with a human's saved Linear session), then exercises the server through
/chat/completions, where the gateway lists and executes its tools with the
stored per-user token.
Management routes (/v1/mcp/server CRUD, /chat/completions) go through the
shared ProxyClient transport. The MCP protocol used to seed the token goes through
the mcp SDK, the same library production MCP hosts run.
"""
from __future__ import annotations
import asyncio
import re
import time
from dataclasses import dataclass
from typing import TYPE_CHECKING
from urllib.parse import parse_qsl
import httpx
import pytest
from mcp import ClientSession
from mcp.client.auth import OAuthClientProvider
from mcp.client.streamable_http import streamable_http_client
from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken
from e2e_config import PROXY_BASE_URL, REQUEST_TIMEOUT
from proxy_client import ProxyClient
from e2e_http import AuthHeaders, NoBody, unwrap
from models import ChatBody, ChatResponse, McpServerCreateBody, McpServerInfo
if TYPE_CHECKING:
from playwright.async_api import Route
# Where the "browser" lands at the end of the authorize dance. Nothing listens
# here: the route interceptor short-circuits the final redirect and reads the
# code/state off its query string, exactly like a desktop MCP host intercepting
# its loopback redirect.
OAUTH_CLIENT_REDIRECT_URI = "http://127.0.0.1:53682/e2e/callback"
BROWSER_CONSENT_TIMEOUT = 60.0
def _mcp_url(alias: str) -> str:
return f"{PROXY_BASE_URL}/{alias}/mcp"
class InMemoryTokenStorage:
"""The mcp SDK's TokenStorage protocol, in memory for one dance: the
DCR-registered client and the gateway tokens minted for it."""
def __init__(self) -> None:
self._tokens: OAuthToken | None = None
self._client_info: OAuthClientInformationFull | None = None
async def get_tokens(self) -> OAuthToken | None:
return self._tokens
async def set_tokens(self, tokens: OAuthToken) -> None:
self._tokens = tokens
async def get_client_info(self) -> OAuthClientInformationFull | None:
return self._client_info
async def set_client_info(self, client_info: OAuthClientInformationFull) -> None:
self._client_info = client_info
async def _browser_follow_authorize(start_url: str, storage_state_path: str) -> tuple[str, str | None]:
"""Play the browser's role for a real upstream whose authorize endpoint
serves an interactive consent page (Linear). A headless Chromium primed
with a human's saved Linear session opens the gateway authorize URL and
clicks through Linear's consent screens (the mcp.linear.app Approve form,
then the linear.app workspace-selection page), riding the rest of the chain
(Linear -> gateway callback -> host redirect_uri). The final hop is
intercepted and short-circuited, since nothing listens there, and its
code/state are read off the query string."""
from playwright.async_api import async_playwright
captured: dict[str, str] = {} # mutable-ok: hand-off from the request listener
trail: list[str] = [] # mutable-ok: navigation diagnostics for a failed dance
def _note_request(request: object) -> None:
url = getattr(request, "url", "")
if url.startswith(OAUTH_CLIENT_REDIRECT_URI) and "url" not in captured:
captured["url"] = url
async def _swallow_redirect(route: "Route") -> None:
await route.fulfill(status=200, content_type="text/plain", body="ok")
async with async_playwright() as playwright:
browser = await playwright.chromium.launch(headless=True)
context = await browser.new_context(storage_state=storage_state_path)
await context.route(re.compile(re.escape(OAUTH_CLIENT_REDIRECT_URI) + r".*"), _swallow_redirect)
page = await context.new_page()
page.on("request", _note_request)
page.on("framenavigated", lambda frame: trail.append(frame.url.split("?", 1)[0]))
await page.goto(start_url, wait_until="domcontentloaded")
deadline = time.monotonic() + BROWSER_CONSENT_TIMEOUT
while "url" not in captured and time.monotonic() < deadline:
try:
await page.wait_for_load_state("networkidle", timeout=8000)
except Exception: # noqa: BLE001 - a busy consent page never idles; fall through and try to advance it
pass
if "url" in captured:
break
control = page.locator(
'button[name="action"][value="approve"], button:has-text("Authorize"), '
'button:has-text("Allow"), button:has-text("@"), a:has-text("@")'
).first
try:
await control.click(timeout=5000)
except Exception: # noqa: BLE001 - nothing to advance yet; loop and re-check
await asyncio.sleep(0.5)
final_url = page.url
await browser.close()
landing = captured.get("url")
assert landing is not None, (
f"consent flow never reached {OAUTH_CLIENT_REDIRECT_URI}; "
f"final={final_url.split('?', 1)[0]!r}; trail={trail[-6:]}"
)
params = dict(parse_qsl(httpx.URL(landing).query.decode()))
assert "code" in params, f"client redirect_uri carried no code: {landing}"
return params["code"], params.get("state")
def _oauth_provider(url: str, storage: InMemoryTokenStorage, storage_state_path: str) -> OAuthClientProvider:
"""The SDK's real OAuth machinery (RFC 9728/8414 discovery, RFC 7591 DCR,
PKCE, token exchange) with the browser leg driven by Playwright against the
upstream's consent screen."""
code_holder: dict[str, str | None] = {} # mutable-ok: hand-off between the two SDK callbacks
async def redirect_handler(authorize_url: str) -> None:
code, state = await _browser_follow_authorize(authorize_url, storage_state_path)
code_holder["code"] = code
code_holder["state"] = state
async def callback_handler() -> tuple[str, str | None]:
code = code_holder.get("code")
assert code is not None, "callback_handler ran before the authorize redirect completed"
return code, code_holder.get("state")
return OAuthClientProvider(
server_url=url,
client_metadata=OAuthClientMetadata.model_validate(
{
"redirect_uris": [OAUTH_CLIENT_REDIRECT_URI],
"token_endpoint_auth_method": "none",
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
"client_name": "e2e-mcp-host",
}
),
storage=storage,
redirect_handler=redirect_handler,
callback_handler=callback_handler,
)
class _HeaderInjectingTransport(httpx.AsyncBaseTransport):
"""Adds the caller's LiteLLM key header to every outgoing SDK request
(discovery, DCR, token exchange), so the gateway resolves which user to
store the upstream token for from the key on the token exchange, exactly
like a production MCP host configured with a LiteLLM key header."""
def __init__(self, inner: httpx.AsyncBaseTransport, headers: dict[str, str]) -> None:
self._inner = inner
self._headers = headers
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
for name, value in self._headers.items():
if name not in request.headers:
request.headers[name] = value
return await self._inner.handle_async_request(request)
def _oauth_http_client(headers: dict[str, str], auth: OAuthClientProvider) -> httpx.AsyncClient:
return httpx.AsyncClient(
headers=headers,
auth=auth,
timeout=httpx.Timeout(REQUEST_TIMEOUT),
follow_redirects=True,
transport=_HeaderInjectingTransport(httpx.AsyncHTTPTransport(), headers),
)
async def _seed_via_dance(
url: str, headers: dict[str, str], storage: InMemoryTokenStorage, storage_state_path: str
) -> tuple[str, ...]:
async with _oauth_http_client(headers, _oauth_provider(url, storage, storage_state_path)) as http_client:
async with streamable_http_client(url, http_client=http_client) as (read, write, _):
async with ClientSession(read, write) as session:
await session.initialize()
listed = await session.list_tools()
return tuple(sorted(tool.name for tool in listed.tools))
@dataclass(frozen=True, slots=True)
class ChatMcpClient:
proxy: ProxyClient
def create_server(self, body: McpServerCreateBody) -> McpServerInfo:
return unwrap(
self.proxy.transport.post(
"/v1/mcp/server",
headers=self.proxy.transport.master,
json=body,
response_type=McpServerInfo,
)
)
def server_info(self, server_id: str) -> McpServerInfo:
return unwrap(
self.proxy.transport.get(
f"/v1/mcp/server/{server_id}",
headers=self.proxy.transport.master,
params=NoBody(),
response_type=McpServerInfo,
)
)
def delete_server(self, server_id: str) -> None:
_ = self.proxy.transport.delete(
f"/v1/mcp/server/{server_id}",
headers=self.proxy.transport.master,
json=NoBody(),
response_type=NoBody,
)
def seed_user_token(self, alias: str, key: str, storage_state_path: str) -> tuple[str, ...]:
"""Drive the interactive authorize dance for `key`'s user so the gateway
stores their upstream token, retried to the shared deadline since the
just-created server and key propagate asynchronously. The LiteLLM key
rides x-litellm-api-key so the gateway binds the token to that user.
Returns the upstream tool names the dance listed, proof the token works."""
headers = {"x-litellm-api-key": f"Bearer {key}"}
storage = InMemoryTokenStorage()
deadline = time.monotonic() + self.proxy.poll_timeout
last_error: Exception | None = None
while time.monotonic() < deadline:
try:
return asyncio.run(_seed_via_dance(_mcp_url(alias), headers, storage, storage_state_path))
except Exception as exc: # noqa: BLE001 - retried to the deadline; the last error surfaces below
last_error = exc
time.sleep(self.proxy.poll_interval)
pytest.fail(
f"authorize dance for {alias!r} never completed within {self.proxy.poll_timeout}s; "
f"last error: {last_error!r}"
)
def chat_with_mcp(self, headers: AuthHeaders, body: ChatBody) -> ChatResponse:
"""POST /chat/completions carrying the LiteLLM key in `headers` (either
ingress form) with an MCP server attached in `body.tools`. The gateway
resolves the user from the key and lists/executes the server's tools
with that user's stored upstream token."""
return unwrap(
self.proxy.transport.post(
"/chat/completions",
headers=headers,
json=body,
response_type=ChatResponse,
)
)
def build_chat_client(proxy: ProxyClient) -> ChatMcpClient:
return ChatMcpClient(proxy=proxy)

View file

@ -0,0 +1,197 @@
"""On-demand e2e: a chat completion drives a gateway-managed OAuth MCP server.
The real end-user flow for MCP over an OAuth server: a user registers a Linear
authorization_code server, authorizes it once so the gateway stores their
upstream token, then sends a normal /chat/completions request with the Linear
MCP attached. The gateway resolves the user from the LiteLLM key, lists Linear's
tools with the stored per-user token, lets the model call one, executes it
upstream with that token, and returns the answer. This is proven against the
real Linear MCP server (mcp.linear.app) and a real Anthropic model, once per
documented ingress header (x-litellm-api-key and Authorization).
The authorize dance is seeded through the mcp SDK's OAuthClientProvider; the one
step Linear cannot auto-approve is the human consent, so it is captured once out
of band (mcp/linear_session_capture.py) into a saved browser session and a
headless Chromium clicks Approve every run. The test therefore skips unless
E2E_LINEAR_STORAGE_STATE points at that session, so it never runs on the per-PR
CI path; it is a nightly/on-demand real-server smoke test.
Fail-before-fix: without the stored per-user token the gateway lists no Linear
tools, so mcp_list_tools comes back empty, nothing is called, and the
assertions fail; a served, called, non-empty Linear tool proves the gateway
pulled and used the user's token.
"""
from __future__ import annotations
import os
import pytest
from e2e_config import CHEAP_ANTHROPIC_MODEL, LINEAR_MCP_URL, LINEAR_STORAGE_STATE, unique_marker
from e2e_http import AuthHeaders
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, KeyGenerateBody, McpChatTool, McpServerCreateBody, ObjectPermission
from proxy_client import ProxyClient
pytest.importorskip("mcp", reason="mcp SDK not installed; run `uv sync --inexact --group e2e-dev`")
pytest.importorskip(
"playwright.async_api",
reason="playwright not installed; run `uv pip install playwright` and `playwright install chromium`",
)
from oauth_chat_client import ChatMcpClient, build_chat_client # noqa: E402 # imports follow the importorskip guards
pytestmark = [
pytest.mark.e2e,
pytest.mark.skipif(
not LINEAR_STORAGE_STATE or not os.path.exists(LINEAR_STORAGE_STATE),
reason="set E2E_LINEAR_STORAGE_STATE to a Linear session captured via mcp/linear_session_capture.py",
),
]
# Pinned from a live dance during verification (never guessed); the gateway
# prefixes every upstream tool name with the server alias. list_teams is a
# read-only Linear tool that takes no arguments and returns the caller's teams.
LINEAR_READONLY_TOOL = "list_teams"
LINEAR_PROMPT = "Use the list_teams tool to list my Linear teams, then reply with the name of one of them."
@pytest.fixture(scope="session")
def chat_client(proxy: ProxyClient) -> ChatMcpClient:
return build_chat_client(proxy)
class TestMcpChatCompletionOauth:
"""A scoped internal-user key on a real Linear authorization_code server,
used through /chat/completions once per ingress header: the gateway pulls
the user's stored upstream token, lists and executes Linear's tools during
the completion, and returns the answer."""
@pytest.mark.covers("mcp.list_tools.oauth.succeeds")
@pytest.mark.covers("mcp.call_tool.oauth.succeeds")
def test_chat_completion_uses_linear_with_x_litellm_api_key_header(
self, chat_client: ChatMcpClient, resources: ResourceManager
) -> None:
marker = unique_marker()
alias = f"e2elinear{marker}"
created = chat_client.create_server(
McpServerCreateBody(
alias=alias,
url=LINEAR_MCP_URL,
allow_all_keys=False,
auth_type="oauth2",
oauth2_flow="authorization_code",
)
)
resources.defer(lambda: chat_client.delete_server(created.server_id))
stored = chat_client.server_info(created.server_id)
assert stored.auth_type == "oauth2"
assert stored.oauth2_flow == "authorization_code"
assert stored.allow_all_keys is False
key = chat_client.proxy.generate_key(
KeyGenerateBody(
user_id="e2e-test-user",
object_permission=ObjectPermission(mcp_servers=[created.server_id]),
)
)
resources.defer(lambda: chat_client.proxy.delete_key(key))
seeded = chat_client.seed_user_token(alias, key, LINEAR_STORAGE_STATE)
assert f"{alias}-{LINEAR_READONLY_TOOL}" in seeded, (
f"the authorize dance listed {seeded}, expected it to include {alias}-{LINEAR_READONLY_TOOL}"
)
response = chat_client.chat_with_mcp(
AuthHeaders.model_validate({"x-litellm-api-key": f"Bearer {key}"}),
ChatBody(
model=CHEAP_ANTHROPIC_MODEL,
messages=[ChatMessage(role="user", content=LINEAR_PROMPT)],
tools=[
McpChatTool(
server_url=f"litellm_proxy/mcp/{alias}",
server_label=alias,
require_approval="never",
)
],
),
)
message = response.choices[0].message
assert message is not None and message.content, f"completion returned no answer: {response}"
meta = message.provider_specific_fields
assert meta is not None, f"no MCP metadata on the completion: {response}"
listed = {t.function.name for t in (meta.mcp_list_tools or []) if t.function}
assert f"{alias}-{LINEAR_READONLY_TOOL}" in listed, (
f"the gateway listed {sorted(listed)}, expected the stored token to surface {alias}-{LINEAR_READONLY_TOOL}"
)
results = [r for r in (meta.mcp_call_results or []) if r.name == f"{alias}-{LINEAR_READONLY_TOOL}"]
assert results and results[0].result, (
f"Linear tool {alias}-{LINEAR_READONLY_TOOL} was not executed with a result: {meta.mcp_call_results}"
)
@pytest.mark.covers("mcp.list_tools.oauth.succeeds")
@pytest.mark.covers("mcp.call_tool.oauth.succeeds")
def test_chat_completion_uses_linear_with_authorization_bearer_header(
self, chat_client: ChatMcpClient, resources: ResourceManager
) -> None:
marker = unique_marker()
alias = f"e2elinear{marker}"
created = chat_client.create_server(
McpServerCreateBody(
alias=alias,
url=LINEAR_MCP_URL,
allow_all_keys=False,
auth_type="oauth2",
oauth2_flow="authorization_code",
)
)
resources.defer(lambda: chat_client.delete_server(created.server_id))
stored = chat_client.server_info(created.server_id)
assert stored.auth_type == "oauth2"
assert stored.oauth2_flow == "authorization_code"
assert stored.allow_all_keys is False
key = chat_client.proxy.generate_key(
KeyGenerateBody(
user_id="e2e-test-user",
object_permission=ObjectPermission(mcp_servers=[created.server_id]),
)
)
resources.defer(lambda: chat_client.proxy.delete_key(key))
seeded = chat_client.seed_user_token(alias, key, LINEAR_STORAGE_STATE)
assert f"{alias}-{LINEAR_READONLY_TOOL}" in seeded, (
f"the authorize dance listed {seeded}, expected it to include {alias}-{LINEAR_READONLY_TOOL}"
)
response = chat_client.chat_with_mcp(
AuthHeaders.model_validate({"authorization": f"Bearer {key}"}),
ChatBody(
model=CHEAP_ANTHROPIC_MODEL,
messages=[ChatMessage(role="user", content=LINEAR_PROMPT)],
tools=[
McpChatTool(
server_url=f"litellm_proxy/mcp/{alias}",
server_label=alias,
require_approval="never",
)
],
),
)
message = response.choices[0].message
assert message is not None and message.content, f"completion returned no answer: {response}"
meta = message.provider_specific_fields
assert meta is not None, f"no MCP metadata on the completion: {response}"
listed = {t.function.name for t in (meta.mcp_list_tools or []) if t.function}
assert f"{alias}-{LINEAR_READONLY_TOOL}" in listed, (
f"the gateway listed {sorted(listed)}, expected the stored token to surface {alias}-{LINEAR_READONLY_TOOL}"
)
results = [r for r in (meta.mcp_call_results or []) if r.name == f"{alias}-{LINEAR_READONLY_TOOL}"]
assert results and results[0].result, (
f"Linear tool {alias}-{LINEAR_READONLY_TOOL} was not executed with a result: {meta.mcp_call_results}"
)

View file

@ -6,6 +6,7 @@ response validates without mirroring every proxy field. No untyped dicts.
from __future__ import annotations
from collections.abc import Sequence
from datetime import datetime
from typing import Literal
@ -196,6 +197,19 @@ class ChatTool(BaseModel):
function: ChatToolFunction
class McpChatTool(BaseModel):
"""An MCP server attached to a chat completion (OpenAI `type: "mcp"` tool).
`server_url` selects the gateway-registered server by its alias suffix; with
`require_approval="never"` the gateway lists, calls, and feeds the server's
tools back to the model in one agentic turn."""
type: Literal["mcp"] = "mcp"
server_url: str
require_approval: str
server_label: str | None = None
allowed_tools: list[str] | None = None
class ChatBody(BaseModel):
model: str
messages: list[ChatMessage]
@ -206,7 +220,7 @@ class ChatBody(BaseModel):
reasoning_effort: str | None = None
thinking: ThinkingParam | None = None
service_tier: str | None = None
tools: list[ChatTool] | None = None
tools: Sequence[ChatTool | McpChatTool] | None = None
tool_choice: str | None = None
guardrails: list[str] | None = None
response_format: dict[str, object] | None = None
@ -242,10 +256,46 @@ class ToolCall(BaseModel):
function: ToolCallFunction = ToolCallFunction()
class McpToolFunctionRef(BaseModel):
name: str
class McpListedTool(BaseModel):
"""One entry of `mcp_list_tools`: a tool the gateway listed from the
attached MCP server and exposed to the model, in OpenAI function shape."""
function: McpToolFunctionRef | None = None
class McpToolCall(BaseModel):
"""One entry of `mcp_tool_calls`: a tool the model asked the gateway to run."""
function: McpToolFunctionRef | None = None
class McpCallResult(BaseModel):
"""One entry of `mcp_call_results`: what the gateway got back from executing
a tool upstream on the caller's behalf."""
name: str | None = None
result: str | None = None
class McpResponseMetadata(BaseModel):
"""`choices[].message.provider_specific_fields` MCP section: which tools the
gateway listed from the attached server, which the model called, and their
results. Populated only when the completion drove an MCP server."""
mcp_list_tools: list[McpListedTool] | None = None
mcp_tool_calls: list[McpToolCall] | None = None
mcp_call_results: list[McpCallResult] | None = None
class OutMessage(BaseModel):
content: str | None = None
reasoning_content: str | None = None
tool_calls: list[ToolCall] | None = None
provider_specific_fields: McpResponseMetadata | None = None
class ChatChoice(BaseModel):
@ -355,6 +405,36 @@ class CountTokensResponse(BaseModel):
input_tokens: int
# ---------- mcp servers ----------
class McpServerCreateBody(BaseModel):
"""POST /v1/mcp/server. For a gateway-managed OAuth server, `auth_type` is
`oauth2` and `oauth2_flow` is `authorization_code`; the upstream endpoints
are discovered and registered via DCR when left unset. `allow_all_keys`
false scopes the server to keys granted it through object_permission."""
alias: str
url: str
transport: str = "http"
allow_all_keys: bool = True
auth_type: str | None = None
oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None
authorization_url: str | None = None
token_url: str | None = None
class McpServerInfo(BaseModel):
"""Response of POST /v1/mcp/server and GET /v1/mcp/server/{server_id}."""
server_id: str
alias: str | None = None
url: str | None = None
auth_type: str | None = None
oauth2_flow: str | None = None
allow_all_keys: bool | None = None
class EmbedBody(BaseModel):
model: str
input: str

View file

@ -11,7 +11,7 @@ import socket
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from models import KeyGenerateBody, LiteLLMParamsBody
@ -23,7 +23,7 @@ BACKEND = "anthropic/claude-haiku-4-5-20251001"
def _require_redis_reachable() -> None:
(host,) = require_env("REDIS_HOST")
host = os.environ["REDIS_HOST"]
port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379")
try:
with socket.create_connection((host, port), timeout=3):

View file

@ -13,7 +13,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
import pytest
from e2e_config import require_env, unique_marker
from e2e_config import unique_marker
from e2e_http import require_successful_call
from lifecycle import ResourceManager
from models import KeyGenerateBody, LiteLLMParamsBody
@ -28,7 +28,7 @@ RECOVERY_TIMEOUT = float(
def _require_redis() -> None:
(host,) = require_env("REDIS_HOST")
host = os.environ["REDIS_HOST"]
port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379")
try:
with socket.create_connection((host, port), timeout=3):

View file

@ -2271,6 +2271,37 @@ def test_token_type_cost_breakdown_reads_cache_write_tokens():
)
def test_generic_cost_per_token_openai_cache_write_tokens_gpt_5_6():
"""
Regression: OpenAI gpt-5.6 reports cache-write tokens under
prompt_tokens_details.cache_write_tokens (not the Anthropic cache_creation_tokens
name). Those tokens must be billed at the cache-write rate rather than the plain
input rate. Customer report: cache creation tokens were never counted for the
GPT-5.6 series, so cost was undercounted on cache-write requests.
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "gpt-5.6"
usage = Usage(
prompt_tokens=1000,
completion_tokens=10,
total_tokens=1010,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0, cache_write_tokens=800),
)
assert usage.prompt_tokens_details.cache_write_tokens == 800
assert usage.prompt_tokens_details.cache_creation_tokens == 800
prompt_cost, _ = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="openai")
info = litellm.get_model_info(model=model, custom_llm_provider="openai")
expected_prompt = (1000 - 800) * info["input_cost_per_token"] + 800 * info["cache_creation_input_token_cost"]
assert prompt_cost == pytest.approx(expected_prompt)
assert info["cache_creation_input_token_cost"] > info["input_cost_per_token"]
assert prompt_cost > 1000 * info["input_cost_per_token"]
def test_token_type_cost_breakdown_reconciles_with_generic_total():
"""
Both-ways check: the reasoning subset must sum with the remaining (text) output
@ -2326,6 +2357,65 @@ def test_token_type_cost_breakdown_zero_without_special_tokens():
)
@pytest.mark.parametrize(
"raw_usage, expect_read, expect_write",
[
(
{
"input_tokens": 5000,
"output_tokens": 10,
"total_tokens": 5010,
"input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 4012},
},
False,
True,
),
(
{
"input_tokens": 5000,
"output_tokens": 10,
"total_tokens": 5010,
"input_tokens_details": {"cached_tokens": 4012, "cache_write_tokens": 0},
},
True,
False,
),
],
)
def test_token_type_cost_breakdown_openai_responses_api_cache_write_read(
raw_usage, expect_read, expect_write
):
"""Regression for #34309: OpenAI Responses API reports cache tokens under
input_tokens_details.{cached_tokens, cache_write_tokens}, not the Anthropic-style
top-level cache_creation_input_tokens. The itemized breakdown must still populate
cache_read_cost / cache_creation_cost from the transformed usage."""
from litellm.responses.utils import ResponseAPILoggingUtils
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "gpt-5.6"
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_usage)
breakdown = get_token_type_cost_breakdown(
model=model, custom_llm_provider="openai", usage=usage
)
info = litellm.get_model_info(model=model, custom_llm_provider="openai")
if expect_write:
assert breakdown.cache_creation_cost == pytest.approx(
4012 * info["cache_creation_input_token_cost"]
)
assert breakdown.cache_creation_cost > 0
assert breakdown.cache_read_cost == 0.0
if expect_read:
assert breakdown.cache_read_cost == pytest.approx(
4012 * info["cache_read_input_token_cost"]
)
assert breakdown.cache_read_cost > 0
assert breakdown.cache_creation_cost == 0.0
def test_token_type_cost_breakdown_handles_unknown_model_gracefully():
"""A model with no pricing must yield zeros, never raise."""
breakdown = get_token_type_cost_breakdown(

View file

@ -1,539 +0,0 @@
"""
Tests for OAuth 2.0 Token Exchange (RFC 8693) handler for MCP servers.
Covers: exchange flow, caching, error handling, resolve_mcp_auth integration,
bearer token extraction, and config loading.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.proxy._experimental.mcp_server.auth.token_exchange import (
TOKEN_EXCHANGE_GRANT_TYPE,
TokenExchangeHandler,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import (
resolve_mcp_auth,
)
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def _obo_server(**overrides) -> MCPServer:
defaults = dict(
server_id="srv-obo-1",
name="test-obo",
url="https://mcp.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
client_id="litellm-client-id",
client_secret="litellm-client-secret",
token_exchange_endpoint="https://idp.example.com/oauth2/token",
audience="api://mcp-server",
scopes=["mcp.tools.read", "mcp.tools.execute"],
)
defaults.update(overrides)
return MCPServer(**defaults)
def _exchange_response(token="exchanged-tok-abc", expires_in=3600):
resp = MagicMock()
resp.json.return_value = {
"access_token": token,
"token_type": "Bearer",
"expires_in": expires_in,
}
resp.raise_for_status = MagicMock()
resp.text = ""
return resp
# ── Exchange Flow ──
@pytest.mark.asyncio
async def test_exchange_token_success():
"""Token exchange sends correct RFC 8693 parameters and returns access_token."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("scoped-token-1")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
result = await handler.exchange_token("user-jwt-xyz", server)
assert result == "scoped-token-1"
mock_client.post.assert_called_once()
_, kwargs = mock_client.post.call_args
data = kwargs["data"]
assert data["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE
assert data["subject_token"] == "user-jwt-xyz"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:access_token"
assert data["audience"] == "api://mcp-server"
assert data["scope"] == "mcp.tools.read mcp.tools.execute"
assert data["client_id"] == "litellm-client-id"
assert data["client_secret"] == "litellm-client-secret"
@pytest.mark.asyncio
async def test_exchange_token_no_audience():
"""When audience is None, it is omitted from the request."""
handler = TokenExchangeHandler()
server = _obo_server(audience=None)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response()
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
await handler.exchange_token("user-jwt", server)
_, kwargs = mock_client.post.call_args
assert "audience" not in kwargs["data"]
@pytest.mark.asyncio
async def test_exchange_token_no_scopes():
"""When scopes is None, scope param is omitted from the request."""
handler = TokenExchangeHandler()
server = _obo_server(scopes=None)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response()
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
await handler.exchange_token("user-jwt", server)
_, kwargs = mock_client.post.call_args
assert "scope" not in kwargs["data"]
# ── Caching ──
@pytest.mark.asyncio
async def test_exchange_token_cached():
"""Second call with same user token uses cache — only 1 HTTP POST."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cached-exchange-tok")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
t1 = await handler.exchange_token("same-jwt", server)
t2 = await handler.exchange_token("same-jwt", server)
assert t1 == t2 == "cached-exchange-tok"
assert mock_client.post.call_count == 1
@pytest.mark.asyncio
async def test_different_user_tokens_not_shared():
"""Different user JWTs get different exchanged tokens."""
handler = TokenExchangeHandler()
server = _obo_server()
call_count = 0
async def mock_post(url, data=None):
nonlocal call_count
call_count += 1
resp = MagicMock()
resp.json.return_value = {
"access_token": f"exchanged-{call_count}",
"expires_in": 3600,
}
resp.raise_for_status = MagicMock()
return resp
mock_client = AsyncMock()
mock_client.post = mock_post
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
t1 = await handler.exchange_token("user-a-jwt", server)
t2 = await handler.exchange_token("user-b-jwt", server)
assert t1 == "exchanged-1"
assert t2 == "exchanged-2"
assert call_count == 2
# ── Error Handling ──
@pytest.mark.asyncio
async def test_exchange_token_http_error():
"""HTTP errors from the IDP are wrapped in a ValueError."""
handler = TokenExchangeHandler()
server = _obo_server()
mock_response = MagicMock()
mock_response.status_code = 400
mock_response.text = "invalid_grant"
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Bad Request",
request=MagicMock(),
response=mock_response,
)
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="failed with status 400"),
):
await handler.exchange_token("bad-jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_http_error_does_not_log_response_body():
"""Raw IDP error bodies are not logged because they can contain credentials."""
handler = TokenExchangeHandler()
server = _obo_server()
raw_response_body = "client_secret=do-not-log"
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.text = raw_response_body
mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
"Unauthorized",
request=MagicMock(),
response=mock_response,
)
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.verbose_logger.debug"
) as mock_debug,
pytest.raises(ValueError, match="failed with status 401"),
):
await handler.exchange_token("bad-jwt", server)
logged_values = " ".join(
str(value)
for call in mock_debug.call_args_list
for value in [*call.args, *call.kwargs.values()]
)
assert raw_response_body not in logged_values
@pytest.mark.asyncio
async def test_exchange_token_missing_access_token():
"""Response without access_token raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server()
resp = MagicMock()
resp.json.return_value = {"token_type": "Bearer"}
resp.raise_for_status = MagicMock()
mock_client = AsyncMock()
mock_client.post.return_value = resp
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
),
pytest.raises(ValueError, match="missing 'access_token'"),
):
await handler.exchange_token("jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_missing_endpoint():
"""Missing token_exchange_endpoint and token_url raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server(token_exchange_endpoint=None, token_url=None)
with pytest.raises(ValueError, match="no token_exchange_endpoint or token_url"):
await handler.exchange_token("jwt", server)
@pytest.mark.asyncio
async def test_exchange_token_missing_credentials():
"""Missing client_id or client_secret raises ValueError."""
handler = TokenExchangeHandler()
server = _obo_server(client_id=None, client_secret=None)
# has_token_exchange_config will be False, so we call _do_exchange directly
with pytest.raises(ValueError, match="missing client_id or client_secret"):
await handler._do_exchange("jwt", server)
# ── resolve_mcp_auth Integration ──
@pytest.mark.asyncio
async def test_resolve_mcp_auth_with_token_exchange():
"""resolve_mcp_auth delegates to token exchange when server has OBO config and subject_token provided."""
server = _obo_server()
mock_handler = AsyncMock()
mock_handler.exchange_token.return_value = "obo-scoped-token"
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.mcp_token_exchange_handler",
mock_handler,
):
result = await resolve_mcp_auth(server, subject_token="user-jwt")
assert result == "obo-scoped-token"
mock_handler.exchange_token.assert_called_once_with("user-jwt", server)
@pytest.mark.asyncio
async def test_resolve_mcp_auth_obo_without_subject_token_falls_through():
"""Without a subject_token, resolve_mcp_auth falls through to client_credentials."""
server = _obo_server(
token_url="https://auth.example.com/token",
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cc-token")
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
):
result = await resolve_mcp_auth(server, subject_token=None)
# Falls through to client_credentials since subject_token is None
# The server has client_id/client_secret/token_url so has_client_credentials is True
assert result == "cc-token"
@pytest.mark.asyncio
async def test_resolve_mcp_auth_obo_without_subject_token_uses_cached_client_credentials():
"""The M2M fallback for OBO servers reuses the client_credentials cache."""
server = _obo_server(
server_id="srv-obo-m2m-cache",
token_url="https://auth.example.com/token",
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("cached-cc-token")
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
):
first = await resolve_mcp_auth(server, subject_token=None)
second = await resolve_mcp_auth(server, subject_token=None)
assert first == second == "cached-cc-token"
mock_client.post.assert_called_once()
@pytest.mark.asyncio
async def test_resolve_mcp_auth_header_beats_obo():
"""An explicit mcp_auth_header takes priority over OBO token exchange."""
server = _obo_server()
result = await resolve_mcp_auth(
server, mcp_auth_header="Bearer override", subject_token="user-jwt"
)
assert result == "Bearer override"
# ── Bearer Token Extraction ──
def test_extract_bearer_token_from_oauth2_headers():
"""Extracts token from oauth2_headers Authorization header."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers={"Authorization": "Bearer my-jwt-token"},
raw_headers=None,
)
assert result == "my-jwt-token"
def test_extract_bearer_token_from_raw_headers():
"""Falls back to raw_headers when oauth2_headers missing."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers=None,
raw_headers={"authorization": "Bearer raw-jwt"},
)
assert result == "raw-jwt"
def test_extract_bearer_token_no_bearer_prefix():
"""Returns token as-is when no Bearer prefix."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers={"Authorization": "some-opaque-token"},
raw_headers=None,
)
assert result == "some-opaque-token"
def test_extract_bearer_token_none():
"""Returns None when no auth headers present."""
result = MCPServerManager._extract_bearer_token(
oauth2_headers=None,
raw_headers=None,
)
assert result is None
# ── MCPServer Properties ──
def test_has_token_exchange_config_true():
"""has_token_exchange_config is True for a fully configured OBO server."""
server = _obo_server()
assert server.has_token_exchange_config is True
def test_has_token_exchange_config_false_wrong_auth_type():
"""has_token_exchange_config is False when auth_type is not oauth2_token_exchange."""
server = _obo_server(auth_type=MCPAuth.oauth2)
assert server.has_token_exchange_config is False
def test_has_token_exchange_config_false_missing_creds():
"""has_token_exchange_config is False when client_id/client_secret missing."""
server = _obo_server(client_id=None)
assert server.has_token_exchange_config is False
def test_has_token_exchange_config_uses_token_url_fallback():
"""has_token_exchange_config is True when token_url is set instead of token_exchange_endpoint."""
server = _obo_server(
token_exchange_endpoint=None,
token_url="https://idp.example.com/token",
)
assert server.has_token_exchange_config is True
# ── Config Loading ──
@pytest.mark.asyncio
async def test_config_loading_token_exchange_fields():
"""load_servers_from_config correctly maps OBO config fields to MCPServer."""
manager = MCPServerManager()
config = {
"my_obo_server": {
"url": "https://mcp.example.com/mcp",
"transport": "http",
"auth_type": "oauth2_token_exchange",
"client_id": "my-client",
"client_secret": "my-secret",
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://my-mcp",
"scopes": ["read", "write"],
"subject_token_type": "urn:ietf:params:oauth:token-type:jwt",
}
}
await manager.load_servers_from_config(config)
servers = list(manager.config_mcp_servers.values())
assert len(servers) == 1
server = servers[0]
assert server.auth_type == MCPAuth.oauth2_token_exchange
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert server.audience == "api://my-mcp"
assert server.subject_token_type == "urn:ietf:params:oauth:token-type:jwt"
assert server.client_id == "my-client"
assert server.client_secret == "my-secret"
assert server.scopes == ["read", "write"]
assert server.has_token_exchange_config is True
@pytest.mark.asyncio
async def test_config_loading_default_subject_token_type():
"""subject_token_type defaults to access_token when not specified in config."""
manager = MCPServerManager()
config = {
"obo_defaults": {
"url": "https://mcp.example.com/mcp",
"transport": "http",
"auth_type": "oauth2_token_exchange",
"client_id": "cid",
"client_secret": "csec",
"token_exchange_endpoint": "https://idp.example.com/token",
}
}
await manager.load_servers_from_config(config)
server = list(manager.config_mcp_servers.values())[0]
assert server.subject_token_type == "urn:ietf:params:oauth:token-type:access_token"
@pytest.mark.asyncio
async def test_database_loading_token_exchange_scopes_from_credentials():
"""DB-loaded OBO server credentials retain configured scopes."""
manager = MCPServerManager()
db_server = LiteLLM_MCPServerTable(
server_id="srv-obo-db",
server_name="obo_db_server",
url="https://mcp.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
credentials={
"client_id": "db-client",
"client_secret": "db-secret",
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://db-mcp",
"scopes": ["db.read", "db.write"],
},
)
server = await manager.build_mcp_server_from_table(
db_server,
credentials_are_encrypted=False,
)
assert server.auth_type == MCPAuth.oauth2_token_exchange
assert server.client_id == "db-client"
assert server.client_secret == "db-secret"
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert server.audience == "api://db-mcp"
assert server.scopes == ["db.read", "db.write"]
@pytest.mark.asyncio
async def test_exchange_token_uses_client_secret_basic_when_configured():
"""LIT-4091: token exchange with token_endpoint_auth_method=client_secret_basic sends the
client credentials as HTTP Basic and omits client_secret from the body."""
import base64
handler = TokenExchangeHandler()
server = _obo_server(
server_id="srv-obo-basic", token_endpoint_auth_method="client_secret_basic"
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("scoped-basic")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
result = await handler.exchange_token("user-jwt-basic", server)
assert result == "scoped-basic"
_, kwargs = mock_client.post.call_args
expected = "Basic " + base64.b64encode(b"litellm-client-id:litellm-client-secret").decode()
assert kwargs["headers"]["Authorization"] == expected
assert "client_secret" not in kwargs["data"]
assert "client_id" not in kwargs["data"]
assert kwargs["data"]["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE

View file

@ -3302,19 +3302,20 @@ async def test_token_root_does_not_resolve_private_server_for_external_client():
@pytest.mark.asyncio
async def test_register_root_resolves_single_oauth2_server():
"""When /register is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
try:
from fastapi import Request
async def test_register_root_does_aggregate_dcr_not_single_server_resolution():
"""Root /register is the aggregate DCR endpoint: it mints a stateless llm_dcrc_ client
from the request's redirect_uris and does NOT resolve a single configured oauth2 server
(a single-server deployment registers at /{server}/register instead)."""
import json
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server()
@ -3325,33 +3326,37 @@ async def test_register_root_resolves_single_oauth2_server():
mock_request.headers = {}
try:
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value={}),
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}),
),
patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"),
):
result = await register_client(request=mock_request, mcp_server_name=None)
response = await register_client(request=mock_request, mcp_server_name=None)
# Should resolve to the single server and return its name as client_id
assert result["client_id"] == "test_oauth"
assert "redirect_uris" in result
body = json.loads(response.body)
assert body["client_id"].startswith("llm_dcrc_")
assert body["client_id"] != "test_oauth"
assert body["token_endpoint_auth_method"] == "none"
finally:
global_mcp_server_manager.registry.clear()
@pytest.mark.asyncio
async def test_register_root_does_not_resolve_private_server_for_external_client():
"""Root /register must not reveal or use a hidden MCP server."""
try:
from fastapi import Request
async def test_register_root_does_not_leak_a_private_server():
"""Root /register never resolves or reveals a configured server, so a private one cannot
leak to an external caller: it always mints the aggregate DCR client instead."""
import json
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.registry.clear()
oauth2_server = _create_oauth2_server(available_on_public_internet=False)
@ -3365,17 +3370,19 @@ async def test_register_root_does_not_resolve_private_server_for_external_client
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value={}),
new=AsyncMock(return_value={"redirect_uris": ["https://claude.ai/cb"]}),
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="198.51.100.10",
),
patch("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637"),
):
result = await register_client(request=mock_request, mcp_server_name=None)
response = await register_client(request=mock_request, mcp_server_name=None)
assert result["client_id"] == "dummy_client"
assert result["redirect_uris"] == ["https://llm.example.com/callback"]
body = json.loads(response.body)
assert body["client_id"].startswith("llm_dcrc_")
assert "test_oauth" not in body["client_id"]
finally:
global_mcp_server_manager.registry.clear()
@ -5155,7 +5162,10 @@ async def test_bridge_refresh_grant_with_non_envelope_is_invalid_grant_before_up
def _mint_test_refresh_envelope(
server_id="bridge_srv", key_hash="hashed-litellm-key-77", upstream_refresh="UPSTREAM-REFRESH", identity=None,
server_id="bridge_srv",
key_hash="hashed-litellm-key-77",
upstream_refresh="UPSTREAM-REFRESH",
identity=None,
scope=None,
):
"""Mint a refresh envelope the way the producer does, for driving the refresh_token grant in tests.
@ -5178,7 +5188,9 @@ def _mint_test_refresh_envelope(
keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY)
identity = identity if identity is not None else key_hash_identity(server_id=server_id, key_hash=key_hash)
sealed = build_bridge_refresh_token_response(
identity, RefreshCredential(refresh_token=SecretStr(upstream_refresh), scope=scope), keys,
identity,
RefreshCredential(refresh_token=SecretStr(upstream_refresh), scope=scope),
keys,
datetime.now(timezone.utc),
)
assert isinstance(sealed, SealedEnvelope)
@ -5413,7 +5425,10 @@ async def test_bridge_refresh_re_requests_the_sealed_scope_when_client_omits_it(
)
captured: dict = {}
response = await _refresh_for_bridge_server(
server, refresh_env, {"access_token": "NEW-ACCESS", "token_type": "Bearer", "expires_in": 3600}, None,
server,
refresh_env,
{"access_token": "NEW-ACCESS", "token_type": "Bearer", "expires_in": 3600},
None,
fake_client_out=captured,
)
@ -5553,7 +5568,9 @@ async def test_bridge_refresh_upstream_invalid_grant_maps_to_invalid_grant():
error_response = MagicMock()
error_response.status_code = 400
error_response.text = '{"error": "invalid_grant", "error_description": "refresh token expired"}'
error_response.json = MagicMock(return_value={"error": "invalid_grant", "error_description": "refresh token expired"})
error_response.json = MagicMock(
return_value={"error": "invalid_grant", "error_description": "refresh token expired"}
)
error_response.raise_for_status = MagicMock(
side_effect=httpx.HTTPStatusError("bad", request=MagicMock(), response=error_response)
)
@ -7183,7 +7200,9 @@ def _upstream_token_response(status_code: int, *, json_body: object = None, text
return httpx.Response(status_code, text=text_body, request=request)
async def _exchange_with_upstream_response(upstream_response, *, server_client_id="web-client.apps.googleusercontent.com"):
async def _exchange_with_upstream_response(
upstream_response, *, server_client_id="web-client.apps.googleusercontent.com"
):
"""Run the raw (non-bridge) authorization_code exchange against a canned upstream token-endpoint
response and return what the gateway would hand the client. ``server_client_id=None`` models the
caller-supplied-credentials flow (no stored client on the server)."""
@ -7334,9 +7353,7 @@ async def test_token_exchange_bounds_relayed_error_fields():
async def test_token_exchange_200_without_access_token_is_502_not_keyerror():
"""A 200 whose body has no usable access_token used to KeyError into a 500; the raw arm now
answers 502 with the same wording as the bridge arm's no_upstream_token rejection."""
response = await _exchange_with_upstream_response(
_upstream_token_response(200, json_body={"token_type": "Bearer"})
)
response = await _exchange_with_upstream_response(_upstream_token_response(200, json_body={"token_type": "Bearer"}))
assert response.status_code == 502
body = json.loads(response.body)
@ -7357,7 +7374,9 @@ async def test_token_exchange_relays_rejection_when_http_client_raises():
)
raising_client = MagicMock()
raising_client.post = AsyncMock(
side_effect=httpx.HTTPStatusError("Client error '401 Unauthorized'", request=rejection.request, response=rejection)
side_effect=httpx.HTTPStatusError(
"Client error '401 Unauthorized'", request=rejection.request, response=rejection
)
)
from fastapi import Request
@ -7422,7 +7441,9 @@ async def test_register_relays_rejection_when_http_client_raises():
)
raising_client = MagicMock()
raising_client.post = AsyncMock(
side_effect=httpx.HTTPStatusError("Client error '400 Bad Request'", request=rejection.request, response=rejection)
side_effect=httpx.HTTPStatusError(
"Client error '400 Bad Request'", request=rejection.request, response=rejection
)
)
oauth2_server = _bridge_server(auth_type=MCPAuth.oauth2, dcr_bridge=None)
@ -7800,9 +7821,7 @@ async def test_hydrate_does_not_overwrite_explicit_config_client_id():
auth_type=MCPAuth.oauth2,
client_id="explicit-from-config",
)
store_read = AsyncMock(
return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []}
)
store_read = AsyncMock(return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []})
with (
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
@ -8019,9 +8038,7 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate():
mock_request.headers = {}
try:
authorization_response = _build_oauth_authorization_server_response(
request=mock_request, mcp_server_name=None
)
authorization_response = _build_oauth_authorization_server_response(request=mock_request, mcp_server_name=None)
resource_response = await _build_oauth_protected_resource_response(
request=mock_request, mcp_server_name=None, use_standard_pattern=True
)
@ -8033,6 +8050,65 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate():
global_mcp_server_manager.registry.clear()
def test_gateway_dcr_flow_routing_engages_only_for_llm_dcrc_clients(monkeypatch):
"""The aggregate DCR arms engage for llm_dcrc_ client_ids (register always mints one,
authorize/token route into the aggregate flow); a non-gateway client_id keeps the
per-server behavior, and /authorize/complete exists but 400s without a valid flow."""
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import router
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-for-lit3637")
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-for-lit3637", raising=False)
global_mcp_server_manager.registry.clear()
app = FastAPI()
app.include_router(router)
client = TestClient(app)
registered = client.post("/register", json={"redirect_uris": ["https://claude.ai/cb"]})
assert registered.status_code == 201
assert registered.json()["client_id"].startswith("llm_dcrc_")
assert registered.json()["token_endpoint_auth_method"] == "none"
authorize_params = {
"client_id": "llm_dcrc_bogus",
"redirect_uri": "https://claude.ai/cb",
"response_type": "code",
"code_challenge": "c" * 43,
"code_challenge_method": "S256",
}
bogus_client = client.get("/authorize", params=authorize_params)
assert bogus_client.status_code == 400
assert bogus_client.json()["error"] == "invalid_client"
no_cookie = client.post("/authorize/complete", data={"flow": "h"})
assert no_cookie.status_code == 400
assert no_cookie.json()["error"] == "invalid_request"
token_response = client.post(
"/token",
data={
"grant_type": "authorization_code",
"client_id": "llm_dcrc_bogus",
"code": "x",
"redirect_uri": "https://claude.ai/cb",
"code_verifier": "v" * 43,
},
)
assert token_response.status_code == 400
assert token_response.json()["error"] == "invalid_grant"
upstream_shaped = client.post(
"/token",
data={"grant_type": "authorization_code", "client_id": "regular-upstream-client", "code": "x"},
)
assert upstream_shaped.status_code == 404
@pytest.mark.asyncio
async def test_authorize_wall_names_the_fix_for_urlless_servers():
"""LIT-4629: the authorize wall previously said only "authorization url is not set" with no
@ -8150,6 +8226,118 @@ async def test_register_wall_names_the_fix_for_urlless_servers():
assert "Issuer" in detail_text
@pytest.mark.asyncio
async def test_authorize_wall_points_at_discovery_failure_for_url_servers():
"""LIT-4658: a server WITH a url that still has no authorization_url got here because OAuth
discovery against that url failed (typically a misconfigured url); the old detail blamed
"servers with no url", sending the operator down the wrong path. The detail must now name the
discovery failure and point at the proxy logs where LIT-4658's warnings carry the reason."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
authorize_with_server,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="typo-url-wall",
name="typo_wall",
server_name="typo_wall",
url="https://typo-host.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
)
mock_request = MagicMock()
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await authorize_with_server(
request=mock_request,
mcp_server=server,
client_id="client",
redirect_uri="http://localhost/callback",
)
assert exc_info.value.status_code == 400
detail_text = str(exc_info.value.detail)
assert "may be misconfigured" in detail_text
assert "proxy logs" in detail_text
assert "Servers with no url" not in detail_text
assert "typo-host.example.com" not in detail_text
@pytest.mark.asyncio
async def test_token_wall_points_at_discovery_failure_for_url_servers():
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
exchange_token_with_server,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="typo-url-token-wall",
name="typo_token_wall",
server_name="typo_token_wall",
url="https://typo-host.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
authorization_url="https://idp.example.com/authorize",
)
mock_request = MagicMock()
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await exchange_token_with_server(
request=mock_request,
mcp_server=server,
grant_type="authorization_code",
code="auth-code",
redirect_uri="http://localhost/callback",
client_id="client",
client_secret=None,
code_verifier="verifier",
)
assert exc_info.value.status_code == 400
detail_text = str(exc_info.value.detail)
assert "token url is not configured" in detail_text
assert "may be misconfigured" in detail_text
assert "Servers with no url" not in detail_text
@pytest.mark.asyncio
async def test_authorize_wall_names_the_issuer_for_anchored_servers():
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
authorize_with_server,
)
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="anchored-wall",
name="anchored_wall",
server_name="anchored_wall",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
issuer="https://idp.example.com",
issuer_is_anchored=True,
)
mock_request = MagicMock()
mock_request.base_url = "https://litellm.example.com/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await authorize_with_server(
request=mock_request,
mcp_server=server,
client_id="client",
redirect_uri="http://localhost/callback",
)
assert exc_info.value.status_code == 400
detail_text = str(exc_info.value.detail)
assert "verify the Issuer" in detail_text
assert "Servers with no url" not in detail_text
assert "idp.example.com" not in detail_text
def test_passthrough_authorization_code_round_trips_and_rejects_hostile_input():
"""The passthrough gateway code seals and recovers the ephemeral DCR client and upstream code,
and is total over hostile input: a raw upstream code opens to None, and a tampered or

View file

@ -0,0 +1,590 @@
"""Tests for the aggregate gateway DCR flow (register, authorize, complete, token)."""
import hashlib
import json
from base64 import urlsafe_b64encode
from datetime import datetime, timedelta, timezone
from http.cookies import SimpleCookie
from urllib.parse import parse_qs, urlparse
import pytest
from starlette.requests import Request
from litellm.caching.caching import DualCache
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import (
CONNECT_FLOW_COOKIE_PREFIX,
GATEWAY_AUTH_CODE_PREFIX,
GATEWAY_AUTH_CODE_TTL_SECONDS,
GATEWAY_DCR_CLIENT_ID_PREFIX,
_GatewayAuthCode,
_seal,
aggregate_authorize,
aggregate_token,
complete_connect_flow,
is_gateway_dcr_client_id,
open_gateway_dcr_client,
register_aggregate_client,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import (
resolve_session_bearer,
session_keys_from_master_key,
SessionBearerAdmitted,
)
MASTER_KEY = "sk-gateway-dcr-flow-tests"
REDIRECT_URI = "https://claude.ai/api/mcp/auth_callback"
CODE_VERIFIER = "verifier-" + "v" * 43
CODE_CHALLENGE = urlsafe_b64encode(hashlib.sha256(CODE_VERIFIER.encode("ascii")).digest()).rstrip(b"=").decode("ascii")
@pytest.fixture(autouse=True)
def _salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", MASTER_KEY)
def _request(path="/authorize", query="", cookies=None, method="GET"):
cookie_header = []
if cookies:
cookie = SimpleCookie()
for name, value in cookies.items():
cookie[name] = value
cookie_header = [(b"cookie", cookie.output(header="", sep="; ").strip().encode())]
return Request(
{
"type": "http",
"method": method,
"scheme": "https",
"path": path,
"query_string": query.encode(),
"headers": [(b"host", b"llm.example.com"), *cookie_header],
}
)
async def _register(redirect_uris) -> dict:
response = await register_aggregate_client(
request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris}
)
return json.loads(response.body)
async def _reload_user_active(user_id: str):
return None
@pytest.mark.asyncio
async def test_register_mints_stateless_public_client():
body = await _register([REDIRECT_URI])
assert body["token_endpoint_auth_method"] == "none"
assert "client_secret" not in body
assert body["redirect_uris"] == [REDIRECT_URI]
assert is_gateway_dcr_client_id(body["client_id"])
record = open_gateway_dcr_client(body["client_id"])
assert record is not None
assert record.redirect_uris == (REDIRECT_URI,)
@pytest.mark.asyncio
async def test_register_allows_loopback_http_for_dev_clients():
body = await _register(["http://localhost:6274/oauth/callback"])
assert is_gateway_dcr_client_id(body["client_id"])
@pytest.mark.parametrize(
"code_challenge",
["short", "", "p" * 300, "ünïcode-challenge", "AAAA" * 20],
)
def test_pkce_mismatched_challenge_returns_false_never_raises(code_challenge):
"""A wrong-length or non-ASCII code_challenge must VERIFY FALSE, not raise.
Pins the reason this compares bytes rather than str: hmac.compare_digest raises TypeError on
two str with non-ASCII content, but on bytes of unequal length it simply returns False. A
review flagged this as an unhandled 500 on length mismatch; encoding both sides to bytes is
exactly what makes that impossible, so the claim is pinned here rather than in a comment."""
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _pkce_verifier_matches
assert _pkce_verifier_matches("a" * 43, code_challenge) is False
@pytest.mark.asyncio
async def test_register_allows_allowlisted_native_callback():
"""Native MCP clients register a private-use scheme, not https. Registration shares
the one redirect-URI shape owner with /authorize, so the callback the allowlist
already trusts there is registrable here rather than rejected as non-https."""
body = await _register(["cursor://anysphere.cursor-mcp/oauth/callback"])
assert is_gateway_dcr_client_id(body["client_id"])
record = open_gateway_dcr_client(body["client_id"])
assert record is not None
assert record.redirect_uris == ("cursor://anysphere.cursor-mcp/oauth/callback",)
@pytest.mark.asyncio
async def test_register_rejects_userinfo_spoofed_origin():
"""``https://claude.ai@attacker.example/cb`` parses with netloc
``claude.ai@attacker.example``, so a naive origin display on the consent screen reads
as claude.ai while the code would be delivered to attacker.example. Rejected at
registration, which is the only way such a URI could enter a sealed client."""
response = await register_aggregate_client(
request=_request(path="/register", method="POST"),
request_body={"redirect_uris": ["https://claude.ai@attacker.example/callback"]},
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_redirect_uri"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"redirect_uris",
[
[],
"not-a-list",
["http://evil.example.com/callback"],
["https://claude.ai/cb#fragment"],
["ftp://claude.ai/cb"],
["https://a.example.com/" + "p" * 300],
["https://a.example.com/1", "https://a.example.com/2", "https://a.example.com/3", "https://a.example.com/4"],
[12345],
],
)
async def test_register_rejects_bad_redirect_uris(redirect_uris):
response = await register_aggregate_client(
request=_request(path="/register", method="POST"), request_body={"redirect_uris": redirect_uris}
)
assert response.status_code == 400
assert json.loads(response.body)["error"] in ("invalid_redirect_uri", "invalid_client_metadata")
@pytest.mark.asyncio
async def test_tampered_client_id_does_not_open():
body = await _register([REDIRECT_URI])
tampered = body["client_id"][:-4] + "AAAA"
assert open_gateway_dcr_client(tampered) is None
assert open_gateway_dcr_client("llm_dcrc_garbage") is None
assert open_gateway_dcr_client("other_prefix") is None
def _authorize(
client_id, session_user_id, redirect_uri=REDIRECT_URI, challenge=CODE_CHALLENGE, method="S256", response_type="code"
):
return aggregate_authorize(
request=_request(query=f"client_id={client_id}"),
client_id=client_id,
redirect_uri=redirect_uri,
state="client-state-123",
code_challenge=challenge,
code_challenge_method=method,
response_type=response_type,
session_user_id=session_user_id,
)
@pytest.mark.asyncio
async def test_authorize_validation_failures_never_redirect_to_client():
client_id = (await _register([REDIRECT_URI]))["client_id"]
for response, expected_error in (
(_authorize("llm_dcrc_bogus", "u1"), "invalid_client"),
(_authorize(client_id, "u1", redirect_uri="https://attacker.example.com/cb"), "invalid_request"),
(_authorize(client_id, "u1", response_type="token"), "unsupported_response_type"),
(_authorize(client_id, "u1", challenge=None), "invalid_request"),
(_authorize(client_id, "u1", method="plain"), "invalid_request"),
):
assert response.status_code == 400
assert json.loads(response.body)["error"] == expected_error
@pytest.mark.asyncio
async def test_authorize_without_session_redirects_to_login_with_return_to():
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = _authorize(client_id, session_user_id=None)
assert response.status_code == 303
location = response.headers["location"]
assert location.startswith("https://llm.example.com/sso/key/generate?return_to=")
assert "return_to=%2Fauthorize" in location
@pytest.mark.asyncio
async def test_authorize_with_session_hands_browser_to_connect_page_with_flow_cookie():
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = _authorize(client_id, session_user_id="u1")
assert response.status_code == 303
location = urlparse(response.headers["location"])
assert location.path == "/ui/chat/integrations"
params = parse_qs(location.query)
handle = params["connect_flow"][0]
assert params["connect_client"] == ["https://claude.ai"]
set_cookie = response.headers["set-cookie"]
assert f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}" in set_cookie
assert "HttpOnly" in set_cookie
return handle, set_cookie
def _flow_cookie_from(response) -> tuple:
location = urlparse(response.headers["location"])
handle = parse_qs(location.query)["connect_flow"][0]
cookie = SimpleCookie()
cookie.load(response.headers["set-cookie"])
name = f"{CONNECT_FLOW_COOKIE_PREFIX}{handle}"
return handle, {name: cookie[name].value}
@pytest.mark.asyncio
async def test_full_walk_register_authorize_complete_token_and_replay():
"""The whole front door on one deterministic walk: register -> authorize ->
complete -> token, then the security edges on the same artifacts (user mismatch,
PKCE mismatch, single-use replay, refresh rotation, cross-client refresh)."""
client_id = (await _register([REDIRECT_URI]))["client_id"]
authorize_response = _authorize(client_id, session_user_id="u1")
handle, cookies = _flow_cookie_from(authorize_response)
denied = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="attacker",
cache=DualCache(),
)
assert denied.status_code == 403
anonymous = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id=None,
cache=DualCache(),
)
assert anonymous.status_code == 401
completed = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
cache=DualCache(),
)
assert completed.status_code == 303
redirect = urlparse(completed.headers["location"])
assert f"{redirect.scheme}://{redirect.netloc}{redirect.path}" == REDIRECT_URI
params = parse_qs(redirect.query)
assert params["state"] == ["client-state-123"]
code = params["code"][0]
assert code.startswith(GATEWAY_AUTH_CODE_PREFIX)
cache = DualCache()
async def _token(**overrides):
arguments = {
"request": _request("/token", method="POST"),
"grant_type": "authorization_code",
"code": code,
"redirect_uri": REDIRECT_URI,
"client_id": client_id,
"code_verifier": CODE_VERIFIER,
"refresh_token": None,
"master_key": MASTER_KEY,
"reload_user": _reload_user_active,
"cache": cache,
}
return await aggregate_token(**{**arguments, **overrides})
wrong_verifier = await _token(code_verifier="wrong-" + "w" * 43)
assert json.loads(wrong_verifier.body)["error"] == "invalid_grant"
wrong_client = await _token(client_id=(await _register([REDIRECT_URI]))["client_id"])
assert json.loads(wrong_client.body)["error"] == "invalid_grant"
token_response = await _token()
assert token_response.status_code == 200
payload = json.loads(token_response.body)
assert payload["token_type"] == "Bearer"
assert 0 < payload["expires_in"] <= 3600
keys = session_keys_from_master_key(MASTER_KEY)
admitted = resolve_session_bearer(f"Bearer {payload['access_token']}", keys, datetime.now(timezone.utc))
assert isinstance(admitted, SessionBearerAdmitted)
assert admitted.principal.user_id == "u1"
assert admitted.principal.client_id == client_id
replay = await _token()
assert json.loads(replay.body)["error"] == "invalid_grant"
refreshed = await _token(grant_type="refresh_token", code=None, refresh_token=payload["refresh_token"])
assert refreshed.status_code == 200
rotated = json.loads(refreshed.body)
assert rotated["refresh_token"] != payload["refresh_token"]
# Rotation is single-use: replaying the now-consumed refresh token cannot mint a second pair
# (a captured token is dead once the legitimate holder has rotated).
replayed = await _token(grant_type="refresh_token", code=None, refresh_token=payload["refresh_token"])
assert json.loads(replayed.body)["error"] == "invalid_grant"
assert "already used" in json.loads(replayed.body).get("error_description", "")
cross_client = await _token(
grant_type="refresh_token",
code=None,
refresh_token=payload["refresh_token"],
client_id=(await _register([REDIRECT_URI]))["client_id"],
)
assert json.loads(cross_client.body)["error"] == "invalid_grant"
@pytest.mark.asyncio
async def test_complete_rejects_missing_tampered_and_expired_flows():
missing = await complete_connect_flow(
request=_request("/authorize/complete", method="POST"),
flow_handle="nope",
session_user_id="u1",
cache=DualCache(),
)
assert missing.status_code == 400
tampered = await complete_connect_flow(
request=_request("/authorize/complete", cookies={f"{CONNECT_FLOW_COOKIE_PREFIX}h1": "garbage"}, method="POST"),
flow_handle="h1",
session_user_id="u1",
cache=DualCache(),
)
assert tampered.status_code == 400
@pytest.mark.asyncio
async def test_token_rejects_expired_code_and_missing_configuration():
expired_code = _seal(
GATEWAY_AUTH_CODE_PREFIX,
_GatewayAuthCode(
user_id="u1",
client_id="llm_dcrc_x",
redirect_uri=REDIRECT_URI,
code_challenge=CODE_CHALLENGE,
jti="jti-1",
iat=int((datetime.now(timezone.utc) - timedelta(seconds=500)).timestamp()),
exp=int((datetime.now(timezone.utc) - timedelta(seconds=500 - GATEWAY_AUTH_CODE_TTL_SECONDS)).timestamp()),
),
)
response = await aggregate_token(
request=_request("/token", method="POST"),
grant_type="authorization_code",
code=expired_code,
redirect_uri=REDIRECT_URI,
client_id="llm_dcrc_x",
code_verifier=CODE_VERIFIER,
refresh_token=None,
master_key=MASTER_KEY,
reload_user=_reload_user_active,
cache=DualCache(),
)
assert json.loads(response.body)["error"] == "invalid_grant"
no_master_key = await aggregate_token(
request=_request("/token", method="POST"),
grant_type="authorization_code",
code="llm_gcode_x",
redirect_uri=REDIRECT_URI,
client_id="llm_dcrc_x",
code_verifier=CODE_VERIFIER,
refresh_token=None,
master_key=None,
reload_user=_reload_user_active,
cache=DualCache(),
)
assert no_master_key.status_code == 500
assert json.loads(no_master_key.body)["error"] == "server_error"
unsupported = await aggregate_token(
request=_request("/token", method="POST"),
grant_type="password",
code=None,
redirect_uri=None,
client_id="llm_dcrc_x",
code_verifier=None,
refresh_token=None,
master_key=MASTER_KEY,
reload_user=_reload_user_active,
cache=DualCache(),
)
assert json.loads(unsupported.body)["error"] == "unsupported_grant_type"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"failure,expected_status,expected_error",
[
("no_active_key", 400, "invalid_grant"),
("unavailable", 503, "temporarily_unavailable"),
("unresolvable", 500, "server_error"),
],
)
async def test_token_gates_on_live_user_revalidation(failure, expected_status, expected_error):
client_id = (await _register([REDIRECT_URI]))["client_id"]
authorize_response = _authorize(client_id, session_user_id="deactivated-user")
handle, cookies = _flow_cookie_from(authorize_response)
completed = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="deactivated-user",
cache=DualCache(),
)
code = parse_qs(urlparse(completed.headers["location"]).query)["code"][0]
async def _reload_user_failing(user_id: str):
return failure
response = await aggregate_token(
request=_request("/token", method="POST"),
grant_type="authorization_code",
code=code,
redirect_uri=REDIRECT_URI,
client_id=client_id,
code_verifier=CODE_VERIFIER,
refresh_token=None,
master_key=MASTER_KEY,
reload_user=_reload_user_failing,
cache=DualCache(),
)
assert response.status_code == expected_status
assert json.loads(response.body)["error"] == expected_error
@pytest.mark.asyncio
async def test_flow_is_single_use_shared_cache_rejects_second_complete():
"""A double-submit of the finish step mints only ONE code: the second complete over the
same cache fails invalid_request (atomic flow claim), so one sign-in cannot yield two codes."""
cache = DualCache()
client_id = (await _register([REDIRECT_URI]))["client_id"]
handle, cookies = _flow_cookie_from(_authorize(client_id, session_user_id="u1"))
first = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
cache=cache,
)
assert first.status_code == 303
second = await complete_connect_flow(
request=_request("/authorize/complete", cookies=cookies, method="POST"),
flow_handle=handle,
session_user_id="u1",
cache=cache,
)
assert second.status_code == 400
assert json.loads(second.body)["error"] == "invalid_request"
@pytest.mark.asyncio
async def test_token_rejects_out_of_range_code_verifier():
"""RFC 7636: a code_verifier outside 43-128 chars is invalid_request, not a confusing
invalid_grant PKCE-mismatch."""
for bad in ["short", "x" * 200]:
response = await aggregate_token(
request=_request("/token", method="POST"),
grant_type="authorization_code",
code="llm_gcode_whatever",
redirect_uri=REDIRECT_URI,
client_id="llm_dcrc_x",
code_verifier=bad,
refresh_token=None,
master_key=MASTER_KEY,
reload_user=_reload_user_active,
cache=DualCache(),
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_request"
@pytest.mark.asyncio
async def test_authorize_rejects_over_long_state():
client_id = (await _register([REDIRECT_URI]))["client_id"]
response = aggregate_authorize(
request=_request(query=f"client_id={client_id}"),
client_id=client_id,
redirect_uri=REDIRECT_URI,
state="s" * 2000,
code_challenge=CODE_CHALLENGE,
code_challenge_method="S256",
response_type="code",
session_user_id="u1",
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_request"
@pytest.mark.asyncio
async def test_non_ascii_code_challenge_fails_grant_not_500():
"""A non-ASCII code_challenge (unvalidated from the client) must yield a clean
invalid_grant, never a TypeError-driven 500 (bytes comparison, not str)."""
client_id = (await _register([REDIRECT_URI]))["client_id"]
# Seal a code carrying a non-ASCII challenge directly (authorize requires S256 shape,
# but the challenge charset is not validated there, so this state is reachable).
from datetime import datetime, timezone
code = _seal(
GATEWAY_AUTH_CODE_PREFIX,
_GatewayAuthCode(
user_id="u1",
client_id=client_id,
redirect_uri=REDIRECT_URI,
code_challenge="challenge-with-€-non-ascii",
jti="jti-x",
iat=int(datetime.now(timezone.utc).timestamp()),
exp=int(datetime.now(timezone.utc).timestamp()) + 120,
),
)
response = await aggregate_token(
request=_request("/token", method="POST"),
grant_type="authorization_code",
code=code,
redirect_uri=REDIRECT_URI,
client_id=client_id,
code_verifier=CODE_VERIFIER,
refresh_token=None,
master_key=MASTER_KEY,
reload_user=_reload_user_active,
cache=DualCache(),
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_grant"
@pytest.mark.asyncio
async def test_single_use_guard_in_memory_is_single_use_within_process():
"""No Redis configured (single-replica): the in-memory increment is authoritative — the first claim
wins, a replay of the same id loses."""
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard
guard = _SingleUseGuard(DualCache()) # redis_cache is None
assert await guard.claim("jti-inmem", 60) is True
assert await guard.claim("jti-inmem", 60) is False # replay of the same id
@pytest.mark.asyncio
async def test_single_use_guard_uses_redis_as_sole_authority_when_configured():
"""With Redis configured it is the SOLE authority: the shared INCR result decides the claim (1 →
first caller, >1 → replay), and the per-worker in-memory count is never consulted."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard
cache = DualCache()
cache.redis_cache = MagicMock()
cache.redis_cache.async_increment = AsyncMock(return_value=1)
# in-memory must NOT be consulted when Redis is configured — poison it so any fallback is visible.
cache.async_increment_cache = AsyncMock(side_effect=AssertionError("must not fall back to in-memory"))
guard = _SingleUseGuard(cache)
assert await guard.claim("jti-redis", 60) is True
cache.redis_cache.async_increment = AsyncMock(return_value=2)
assert await guard.claim("jti-redis", 60) is False # Redis says 2 → replay
@pytest.mark.asyncio
async def test_single_use_guard_fails_closed_when_redis_errors():
"""A Redis fault must fail the claim CLOSED (refuse the id) rather than fall back to the per-worker
in-memory count — which would let each replica observe count==1 and replay the one-time id (the
Cursor/Veria replay-across-workers finding)."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import _SingleUseGuard
cache = DualCache()
cache.redis_cache = MagicMock()
cache.redis_cache.async_increment = AsyncMock(side_effect=ConnectionError("redis down"))
cache.async_increment_cache = AsyncMock(return_value=1) # would fail OPEN if the guard fell back
guard = _SingleUseGuard(cache)
assert await guard.claim("jti-fault", 60) is False # fail closed, not a fallback count of 1

View file

@ -7375,6 +7375,10 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow():
("", None),
("not a url", None),
("http://[::1", None),
# urlsplit validates the port lazily on attribute access, so a malformed port must not
# raise out of the helper: the server loaders call it while warning about exactly this
# kind of typo'd url (LIT-4658)
("https://example.com:bad/mcp", None),
],
)
def test_redact_mcp_resource_url_strips_credentials(url, expected):

View file

@ -2333,6 +2333,59 @@ class TestMCPServerManager:
assert emitted.headers["Authorization"] == "Bearer upstream-token"
assert not kwargs["extra_headers"] or "authorization" not in {k.lower() for k in kwargs["extra_headers"]}
@pytest.mark.asyncio
async def test_create_mcp_client_token_exchange_never_falls_back_to_v1(self):
"""A configured OBO server is owned end to end by the v2 token_exchange arm, even when the
caller supplies an x-mcp-* override. This is what makes the v1 OBO handler unreachable, so if
it ever defers to v1 again the deleted handler is silently needed back."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import (
UpstreamCredentialProvider,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
class _StubExchanger:
def __init__(self):
self.subject_tokens = []
async def exchange(self, subject_token, server, config, *, tenant_id=""):
self.subject_tokens.append(subject_token)
return Ok(OAuthToken(access_token="exchanged-token"))
async def invalidate(self, subject_token, server, config, *, tenant_id=""):
return None
exchanger = _StubExchanger()
manager = MCPServerManager()
server = MCPServer(
server_id="obo-egress",
name="obo",
url="https://example.com",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
client_id="gateway-client",
client_secret="gateway-secret",
token_exchange_endpoint="https://idp.example.com/oauth2/token",
)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.resolve_mcp_auth",
new_callable=AsyncMock,
) as mock_resolve,
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as mock_client_cls,
):
await manager._create_mcp_client(
server=server,
mcp_auth_header="Bearer caller-override",
subject_token="eyJ-subject-token",
cred_provider=UpstreamCredentialProvider(token_exchanger=exchanger),
)
mock_resolve.assert_not_awaited()
assert exchanger.subject_tokens == ["eyJ-subject-token"]
assert self._emitted_authorization(mock_client_cls) == "Bearer exchanged-token"
@staticmethod
def _emitted_authorization(mock_client_cls) -> str:
kwargs = mock_client_cls.call_args.kwargs
@ -3031,7 +3084,7 @@ class TestMCPServerManager:
registration_url="https://discovered.example.com/register",
)
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
assert server_url == "https://example.com/mcp"
# oauth2 (browser flow) keeps the origin fallback; only OBO disables it.
assert allow_origin_fallback is True
@ -5426,7 +5479,7 @@ class TestMCPServerTimestamps:
manager = MCPServerManager()
calls: list[bool] = []
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
calls.append(allow_origin_fallback)
return MCPOAuthMetadata(
scopes=None,
@ -5461,7 +5514,7 @@ class TestMCPServerTimestamps:
manager = MCPServerManager()
calls: list[str] = []
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
calls.append(server_url)
raise AssertionError("discovery must not run when token_exchange_endpoint is configured")
@ -5491,7 +5544,7 @@ class TestMCPServerTimestamps:
back to the row, so the next rebuild skips discovery instead of re-running it every time."""
manager = MCPServerManager()
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
assert server_url == "https://example.com/mcp"
assert allow_origin_fallback is False # OBO never guesses the origin
return MCPOAuthMetadata(
@ -5602,7 +5655,7 @@ class TestMCPServerTimestamps:
_dcr_bridge_relays_client_registration keys off that column."""
manager = MCPServerManager()
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
assert allow_origin_fallback is True
return MCPOAuthMetadata(
scopes=["mcp.read", "mcp.write"],
@ -5817,7 +5870,7 @@ class TestMCPServerTimestamps:
persist_discovered_endpoints=False neither the oauth2 nor the OBO write-back may fire."""
manager = MCPServerManager()
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
return MCPOAuthMetadata(
scopes=["s1"],
authorization_url="https://idp.example.com/authorize",
@ -8539,7 +8592,7 @@ class TestOBOEndpointDiscovery:
)
seen = []
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
seen.append((server_url, allow_origin_fallback))
return discovered
@ -8567,7 +8620,7 @@ class TestOBOEndpointDiscovery:
async def test_config_obo_with_configured_endpoint_skips_discovery(self):
manager = MCPServerManager()
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True):
async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False):
raise AssertionError("discovery must not run when the endpoint is configured")
manager._descovery_metadata = fake_discovery # type: ignore[attr-defined]
@ -9028,3 +9081,162 @@ class TestUrllessIssuerDiscovery:
anchored.assert_awaited_once_with("https://idp.example.com", None)
resource_rooted.assert_not_awaited()
assert built.token_url == "https://idp.example.com/token"
class TestDiscoveryFailureLogging:
"""LIT-4658: a misconfigured MCP server url must be diagnosable from default-level server logs.
Discovery failures used to die at debug level and the config-load path emitted no warning at
all, so the only operator-facing signal was the bare 400 at /authorize."""
def _connect_error_client(self, url: str) -> MagicMock:
client = MagicMock()
client.get = AsyncMock(
side_effect=httpx.ConnectError(f"[Errno 8] nodename nor servname provided for {url}")
)
return client
@pytest.mark.asyncio
async def test_descovery_metadata_warns_with_redacted_attempts_on_connect_error(self, caplog):
manager = MCPServerManager()
secret_url = "https://typo-host.example.com/mcp/s/PATHSECRET/mcp"
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=self._connect_error_client(secret_url),
),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
result = await manager._descovery_metadata(secret_url, warn_when_no_metadata=True)
assert result is None
assert "found no authorization server metadata" in caplog.text
assert "ConnectError" in caplog.text
assert "https://typo-host.example.com" in caplog.text
# hosted MCP urls embed credentials in the path; neither the url nor the exception
# text may leak it into warning-level logs
assert "PATHSECRET" not in caplog.text
@pytest.mark.asyncio
async def test_descovery_metadata_stays_silent_without_warn_flag(self, caplog):
manager = MCPServerManager()
url = "https://typo-host.example.com/mcp"
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=self._connect_error_client(url),
),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
result = await manager._descovery_metadata(url)
assert result is None
assert "found no authorization server metadata" not in caplog.text
@pytest.mark.asyncio
async def test_descovery_metadata_attempt_trail_names_each_failed_step(self, caplog):
manager = MCPServerManager()
url = "https://real-host.example.com/mcp-typo"
client = MagicMock()
client.get = AsyncMock(
return_value=httpx.Response(404, request=httpx.Request("GET", url))
)
with (
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
return_value=client,
),
caplog.at_level(logging.WARNING, logger="LiteLLM"),
):
result = await manager._descovery_metadata(url, warn_when_no_metadata=True)
assert result is None
assert "HTTP 404" in caplog.text
assert "well-known protected-resource lookup found no authorization servers" in caplog.text
assert "origin fallback" in caplog.text
@pytest.mark.asyncio
async def test_load_servers_from_config_warns_when_endpoints_unresolved(self, caplog):
manager = MCPServerManager()
manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined]
config = {
"typo_server": {
"url": "https://typo.example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2,
"oauth2_flow": "authorization_code",
}
}
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
await manager.load_servers_from_config(config)
assert "typo_server" in caplog.text
assert "authorization_url, token_url" in caplog.text
assert "unresolved" in caplog.text
assert "verify the configured server url" in caplog.text
@pytest.mark.asyncio
@pytest.mark.parametrize(
"extra_config",
[
{
"authorization_url": "https://idp.example.com/auth",
"token_url": "https://idp.example.com/token",
},
{
"oauth2_flow": "client_credentials",
"token_url": "https://idp.example.com/token",
"client_id": "cid",
"client_secret": "csec",
},
],
)
async def test_load_servers_from_config_silent_when_flow_needs_covered(self, caplog, extra_config):
"""Manually covered endpoints and M2M servers (which never need authorization_url) must not
warn on every reload; the warning is a misconfiguration signal, not discovery telemetry."""
manager = MCPServerManager()
manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined]
config = {
"covered_server": {
"url": "https://up.example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2,
"oauth2_flow": "authorization_code",
**extra_config,
}
}
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
await manager.load_servers_from_config(config)
assert "unresolved" not in caplog.text
assert "no discovery source" not in caplog.text
@pytest.mark.asyncio
async def test_config_server_without_discovery_source_warns_about_missing_endpoints(self, caplog):
manager = MCPServerManager()
manager._register_openapi_tools = AsyncMock() # type: ignore[attr-defined]
config = {
"spec_only": {
"spec_path": "https://example.com/openapi.yaml",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2,
"oauth2_flow": "authorization_code",
}
}
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
await manager.load_servers_from_config(config)
assert "no discovery source" in caplog.text
assert "authorization_url and token_url are not set manually" in caplog.text
@pytest.mark.asyncio
async def test_db_build_warns_when_discovery_fails_for_oauth2_row(self, caplog):
manager = MCPServerManager()
manager._descovery_metadata = AsyncMock(return_value=None) # type: ignore[attr-defined]
record = LiteLLM_MCPServerTable(
server_id="typo-row-1",
server_name="typo_row",
url="https://typo.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
oauth2_flow="authorization_code",
)
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False)
assert "typo_row" in caplog.text
assert "authorization_url, token_url" in caplog.text
assert "unresolved" in caplog.text

View file

@ -237,7 +237,7 @@ class TestListToolRestApiWithToolSearch:
return_value={},
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids",
"litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids",
return_value=[],
),
patch(
@ -316,7 +316,7 @@ class TestListToolRestApiWithToolSearch:
return_value={},
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids",
"litellm.proxy._experimental.mcp_server.rest_endpoints._v1_resolved_oauth2_server_ids",
return_value=[],
),
patch(

View file

@ -2783,3 +2783,68 @@ class TestRestListToolsetFiltering:
)
assert [tool.name for tool in result] == ["lookup_status"]
class TestV1ResolvedOauth2Gate:
"""The REST surface must stop resolving per-user OAuth2 tokens for servers the v2 resolver owns.
``_resolve_v2_auth`` drops any Authorization built here for an ``authorization_code`` server and
injects the resolver's own token, so the v1 lookup was a DB round-trip whose result was discarded.
A server that still defers to v1 (upstream-delegated oauth2) must keep resolving, which is what
makes these assertions non-vacuous.
"""
@staticmethod
def _oauth2_server(*, delegate_auth_to_upstream: bool) -> Any:
from litellm.proxy._experimental.mcp_server.server import MCPServer
from litellm.types.mcp import MCPTransport
return MCPServer(
server_id="oauth2-srv",
name="oauth2-srv",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
delegate_auth_to_upstream=delegate_auth_to_upstream,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"delegate_auth_to_upstream, expected_headers, expected_lookups",
[
(False, None, 0),
(True, {"Authorization": "Bearer stored-token"}, 1),
],
)
async def test_user_oauth_headers_skip_v2_owned_servers(
self, delegate_auth_to_upstream, expected_headers, expected_lookups, monkeypatch
):
from litellm.proxy._experimental.mcp_server import db as mcp_db
server = self._oauth2_server(delegate_auth_to_upstream=delegate_auth_to_upstream)
resolve_token = AsyncMock(return_value={"access_token": "stored-token"})
monkeypatch.setattr(mcp_db, "resolve_valid_user_oauth_token", resolve_token)
headers = await rest_endpoints._get_user_oauth_extra_headers(
server,
UserAPIKeyAuth(user_id="alice", api_key="sk-1234"),
prefetched_creds={"oauth2-srv": {"access_token": "stored-token"}},
)
assert headers == expected_headers
assert resolve_token.await_count == expected_lookups
def test_prefetch_preflight_only_counts_v1_resolved_servers(self, monkeypatch):
v2_owned = self._oauth2_server(delegate_auth_to_upstream=False)
v1_resolved = self._oauth2_server(delegate_auth_to_upstream=True)
v1_resolved.server_id = "delegate-srv"
registry = {"oauth2-srv": v2_owned, "delegate-srv": v1_resolved}
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: registry.get(server_id),
)
assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv"]) == set()
assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv", "delegate-srv"]) == {"delegate-srv"}

View file

@ -176,9 +176,7 @@ async def test_authenticate_user_invalid_credentials():
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
with patch.dict(
os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": "correct-password"}
):
with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": "correct-password"}):
with pytest.raises(ProxyException) as exc_info:
await authenticate_user(
username=ui_username,
@ -227,9 +225,7 @@ async def test_authenticate_user_wrong_password():
)
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(
return_value=mock_user
)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=mock_user)
with patch.dict(
os.environ,
@ -279,9 +275,7 @@ async def test_authenticate_user_email_case_insensitive_login():
return None
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(
side_effect=mock_find_first
)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first)
with patch.dict(
os.environ,
@ -334,9 +328,7 @@ async def test_authenticate_user_database_required_for_admin():
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
with patch.dict(
os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password}
):
with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password}):
with patch(
"litellm.proxy.auth.login_utils.user_update",
new_callable=AsyncMock,
@ -429,9 +421,7 @@ def test_authenticate_user_non_ascii_direct_comparison():
assert result is True
# And correctly returns False for different passwords
result = secrets.compare_digest(
password.encode("utf-8"), "different£pass".encode("utf-8")
)
result = secrets.compare_digest(password.encode("utf-8"), "different£pass".encode("utf-8"))
assert result is False
@ -531,9 +521,7 @@ async def test_authenticate_user_database_login_with_non_ascii_password():
return None
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(
side_effect=mock_find_first
)
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first)
with patch.dict(
os.environ,
@ -559,3 +547,58 @@ async def test_authenticate_user_database_login_with_non_ascii_password():
assert isinstance(result, LoginResult)
assert result.user_id == "test-user-123"
assert result.user_email == user_email
class TestEncodeUiSessionJwt:
"""The UI session cookie must carry a bounded exp so it does not stay
signature-valid until the master key rotates, and so the session-cookie readers
that require a bounded lifetime (the MCP interactive sign-in) accept it."""
def _decode(self, token: str) -> dict:
import jwt
return jwt.decode(token, "sk-master-for-tests", algorithms=["HS256"])
def test_encoded_cookie_carries_bounded_exp(self):
import time
from litellm.proxy.auth.login_utils import encode_ui_session_jwt
token_object = {"user_id": "u1", "key": "sk-abc", "login_method": "username_password"}
with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "24h"):
token = encode_ui_session_jwt(token_object, "sk-master-for-tests")
claims = self._decode(token)
assert claims["user_id"] == "u1"
assert claims["login_method"] == "username_password"
remaining = claims["exp"] - int(time.time())
assert 23 * 3600 < remaining <= 24 * 3600
def test_duration_is_honored_from_env(self):
import time
from litellm.proxy.auth.login_utils import encode_ui_session_jwt
with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "1h"):
token = encode_ui_session_jwt({"user_id": "u1"}, "sk-master-for-tests")
remaining = self._decode(token)["exp"] - int(time.time())
assert 0 < remaining <= 3600
def test_cookie_is_accepted_by_the_exp_requiring_session_reader(self):
"""The regression this change exists for: before it, the UI cookie carried no
exp and _user_id_from_session_cookie (require=["exp"]) rejected every real login,
so the MCP interactive sign-in could never capture identity. A cookie minted by
this helper must now be accepted."""
from unittest.mock import MagicMock
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import (
_user_id_from_session_cookie,
)
from litellm.proxy.auth.login_utils import encode_ui_session_jwt
token_object = {"user_id": "cornell-user", "key": "sk-abc", "login_method": "sso"}
with patch("litellm.proxy.auth.login_utils.LITELLM_UI_SESSION_DURATION", "24h"):
token = encode_ui_session_jwt(token_object, "sk-master-for-tests")
request = MagicMock()
request.cookies = {"token": token}
with patch("litellm.proxy.proxy_server.master_key", "sk-master-for-tests"):
assert _user_id_from_session_cookie(request) == "cornell-user"

View file

@ -179,6 +179,45 @@ async def test_budget_reservation_runs_when_not_disabled():
assert user_api_key_auth_obj.budget_reservation == reservation
@pytest.mark.asyncio
@pytest.mark.parametrize(
"general_settings,expected_flag",
[
({"fail_closed_budget_enforcement": True}, True),
({}, False),
],
)
async def test_fail_closed_budget_enforcement_reaches_reservation(
general_settings, expected_flag
):
"""#33923: the strict flag must be threaded into reserve_budget_for_request so a
failed reservation write can reject instead of failing open."""
user_api_key_auth_obj = UserAPIKeyAuth(token="test_token")
with patch(
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
new=AsyncMock(return_value=None),
) as mock_reserve:
await _reserve_budget_after_common_checks(
user_api_key_auth_obj=user_api_key_auth_obj,
request_data={"model": "gpt-4o"},
route="/v1/chat/completions",
llm_router=None,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=MagicMock(),
proxy_logging_obj=MagicMock(),
skip_budget_checks=False,
general_settings=general_settings,
)
assert (
mock_reserve.await_args.kwargs["fail_closed_budget_enforcement"]
is expected_flag
)
@pytest.mark.asyncio
async def test_should_not_reuse_cached_key_object_for_request_state():
key_cache = DualCache()
@ -1290,6 +1329,250 @@ async def test_scim_deactivated_user_key_is_rejected():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
async def test_cached_proxy_admin_key_sets_via_virtual_key_marker():
"""Cached PROXY_ADMIN auth objects early-return before the marked DB and
master-key returns, and cache serialization drops the exclude=True marker;
the cache-hit boundary must restore it or cached admin traffic silently
bypasses overwrite_user_with_key_hash stamping."""
from fastapi import Request
from starlette.datastructures import URL
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
from litellm.proxy.proxy_server import hash_token
api_key = "sk-cached-admin-marker-test"
hashed_key = hash_token(api_key)
cached_token = UserAPIKeyAuth(
api_key=api_key,
token=hashed_key,
user_id="cached-admin-user",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
assert cached_token.via_virtual_key is False
mock_cache = AsyncMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.delete_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
_attrs_to_set = {
"prisma_client": MagicMock(),
"user_api_key_cache": mock_cache,
"proxy_logging_obj": mock_proxy_logging_obj,
"master_key": "sk-master-key",
"general_settings": {},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with patch(
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
return_value=cached_token,
):
result = await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
assert isinstance(result, UserAPIKeyAuth)
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
assert result.via_virtual_key is True
assert result.api_key == hashed_key
finally:
for attr, val in _original_values.items():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
async def test_master_key_auth_sets_via_virtual_key_marker():
"""Master-key requests must also be stamped by overwrite_user_with_key_hash;
the auth path substitutes the stable alias for api_key and must mark the
result as proxy-validated."""
from fastapi import Request
from starlette.datastructures import URL
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
master_key = "sk-master-key"
mock_cache = AsyncMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.delete_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
import litellm.proxy.proxy_server as _proxy_server_mod
_attrs_to_set = {
"prisma_client": MagicMock(),
"user_api_key_cache": mock_cache,
"proxy_logging_obj": mock_proxy_logging_obj,
"master_key": master_key,
"general_settings": {},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
result = await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {master_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
assert isinstance(result, UserAPIKeyAuth)
assert result.via_virtual_key is True
assert result.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS
finally:
for attr, val in _original_values.items():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
async def test_db_virtual_key_auth_sets_via_virtual_key_marker():
"""via_virtual_key gates overwrite_user_with_key_hash stamping and is
forge-stripped from validated input, so the DB auth path setting it by
post-construction assignment is the only thing that turns stamping on."""
from fastapi import Request
from starlette.datastructures import URL
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
from litellm.proxy.proxy_server import hash_token
api_key = "sk-via-virtual-key-marker-test"
hashed_key = hash_token(api_key)
valid_token = UserAPIKeyAuth(
api_key=api_key,
token=hashed_key,
user_id="marker-test-user",
)
mock_cache = AsyncMock()
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.delete_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
mock_proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = (
AsyncMock()
)
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
mock_prisma_client = MagicMock()
import litellm.proxy.proxy_server as _proxy_server_mod
_attrs_to_set = {
"prisma_client": mock_prisma_client,
"user_api_key_cache": mock_cache,
"proxy_logging_obj": mock_proxy_logging_obj,
"master_key": "sk-master-key",
"general_settings": {},
"llm_model_list": [],
"llm_router": None,
"open_telemetry_logger": None,
"model_max_budget_limiter": MagicMock(),
"user_custom_auth": None,
"jwt_handler": None,
"litellm_proxy_admin_name": "admin",
}
_original_values = {
attr: getattr(_proxy_server_mod, attr, None) for attr in _attrs_to_set
}
try:
for attr, val in _attrs_to_set.items():
setattr(_proxy_server_mod, attr, val)
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
with (
patch(
"litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key",
new_callable=AsyncMock,
return_value=valid_token,
),
patch(
"litellm.proxy.auth.user_api_key_auth.get_user_object",
new_callable=AsyncMock,
return_value=None,
),
):
result = await _user_api_key_auth_builder(
request=request,
api_key=f"Bearer {api_key}",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={},
)
assert isinstance(result, UserAPIKeyAuth)
assert result.via_virtual_key is True
assert result.api_key == hashed_key
finally:
for attr, val in _original_values.items():
setattr(_proxy_server_mod, attr, val)
@pytest.mark.asyncio
async def test_return_user_api_key_auth_obj_user_spend_and_budget():
"""

View file

@ -13,12 +13,17 @@ from datetime import datetime, timezone
from typing import Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.management_endpoints.tool_management_endpoints import router
from litellm.proxy.management_endpoints.tool_management_endpoints import (
_build_tool_spend_response,
_ToolSpendRow,
router,
)
from litellm.types.tool_management import LiteLLM_ToolTableRow
# --- helpers ---
@ -50,9 +55,9 @@ def _make_app() -> FastAPI:
# Stub the auth dependency so we don't need a real proxy running.
def _override_auth():
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
return UserAPIKeyAuth(api_key="sk-test", user_id="admin")
return UserAPIKeyAuth(api_key="sk-test", user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
# A real (non-None) prisma stub for truthiness checks.
@ -147,3 +152,117 @@ class TestToolManagementEndpoints:
json={"tool_name": "my_tool", "input_policy": "invalid_value"},
)
assert resp.status_code == 422
def test_tool_spend_route_not_shadowed_by_get_tool(self):
prisma = MagicMock()
prisma.db.query_raw = AsyncMock(return_value=[])
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = self.client.get("/v1/tool/spend")
assert resp.status_code == 200
assert resp.json()["by_tool"] == []
def test_tool_spend_aggregates_and_sorts(self):
rows = [
{"date": "2026-07-01", "tool_name": "search", "call_count": 2, "spend": 1.0, "total_tokens": 100},
{"date": "2026-07-02", "tool_name": "search", "call_count": 1, "spend": 4.0, "total_tokens": 50},
{"date": "2026-07-01", "tool_name": "read_file", "call_count": 3, "spend": 2.0, "total_tokens": 300},
]
prisma = MagicMock()
prisma.db.query_raw = AsyncMock(side_effect=[rows, [{"total_spend": 5.5}]])
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = self.client.get("/v1/tool/spend?start_date=2026-07-01&end_date=2026-07-02")
assert resp.status_code == 200
body = resp.json()
assert [t["tool_name"] for t in body["by_tool"]] == ["search", "read_file"]
search = body["by_tool"][0]
assert search["spend"] == 5.0
assert search["call_count"] == 3
assert search["total_tokens"] == 150
assert len(body["daily"]) == 3
assert body["start_date"] == "2026-07-01"
assert body["end_date"] == "2026-07-02"
assert body["total_spend"] == 5.5
@patch("litellm.proxy.proxy_server.prisma_client", None)
def test_tool_spend_no_db_returns_500(self):
resp = self.client.get("/v1/tool/spend")
assert resp.status_code == 500
def test_tool_spend_end_date_is_inclusive_via_exclusive_next_day_bound(self):
prisma = MagicMock()
prisma.db.query_raw = AsyncMock(return_value=[])
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = self.client.get("/v1/tool/spend?start_date=2026-07-01&end_date=2026-07-02")
assert resp.status_code == 200
expected_binds = (
datetime(2026, 7, 1, tzinfo=timezone.utc).isoformat(),
datetime(2026, 7, 3, tzinfo=timezone.utc).isoformat(),
)
assert prisma.db.query_raw.await_count == 2
for call in prisma.db.query_raw.await_args_list:
assert tuple(call.args[1:]) == expected_binds
assert resp.json()["end_date"] == "2026-07-02"
@pytest.mark.parametrize(
"query",
[
"start_date=not-a-date",
"start_date=2026-02-30",
"start_date=07/01/2026",
"end_date=2026-13-01",
"end_date=20260701",
],
)
def test_tool_spend_malformed_date_returns_400(self, query: str):
prisma = MagicMock()
prisma.db.query_raw = AsyncMock(return_value=[])
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = self.client.get(f"/v1/tool/spend?{query}")
assert resp.status_code == 400
assert "Invalid date format" in resp.json()["detail"]
prisma.db.query_raw.assert_not_awaited()
def test_tool_spend_non_admin_returns_403(self):
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
app = _make_app()
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
api_key="sk-user", user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
)
client = TestClient(app, raise_server_exceptions=True)
prisma = MagicMock()
prisma.db.query_raw = AsyncMock(return_value=[])
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
resp = client.get("/v1/tool/spend")
assert resp.status_code == 403
prisma.db.query_raw.assert_not_awaited()
def _spend_row(date: str, tool_name: str, spend: float, call_count: int = 1, total_tokens: int = 10) -> _ToolSpendRow:
return _ToolSpendRow(date=date, tool_name=tool_name, call_count=call_count, spend=spend, total_tokens=total_tokens)
class TestBuildToolSpendResponse:
def test_multi_tool_attribution_double_counts_per_tool_but_not_total(self):
rows = [
_spend_row("2026-07-01", "a", spend=3.0),
_spend_row("2026-07-01", "b", spend=3.0),
]
resp = _build_tool_spend_response(rows, total_spend=3.0, start_date="2026-07-01", end_date="2026-07-01")
by_tool = {t.tool_name: t.spend for t in resp.by_tool}
assert by_tool == {"a": 3.0, "b": 3.0}
assert resp.total_spend == 3.0
def test_groups_across_days_and_sorts_by_spend(self):
rows = [
_spend_row("2026-07-01", "b", spend=1.0, call_count=2, total_tokens=100),
_spend_row("2026-07-02", "b", spend=4.0, call_count=1, total_tokens=50),
_spend_row("2026-07-01", "a", spend=2.0, call_count=3, total_tokens=300),
]
resp = _build_tool_spend_response(rows, total_spend=7.0, start_date="2026-07-01", end_date="2026-07-02")
assert [(t.tool_name, t.spend, t.call_count, t.total_tokens) for t in resp.by_tool] == [
("b", 5.0, 3, 150),
("a", 2.0, 3, 300),
]
assert len(resp.daily) == 3

View file

@ -7763,3 +7763,80 @@ async def test_cli_completion_persists_assertion_under_db_user_id():
retain_mock.assert_awaited_once_with(user_id="cli-user-id", assertion=assertion)
assert response.status_code == 200
class TestSameOriginReturnPath:
"""The same-origin relative return_to arm added for the MCP gateway DCR authorize
round-trip: only strictly relative paths qualify, so login can never redirect the
browser off the gateway origin."""
def test_accepts_relative_paths(self):
from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path
assert _is_same_origin_return_path("/authorize?client_id=llm_dcrc_x&state=s") is True
assert _is_same_origin_return_path("/some_server/authorize") is True
def test_rejects_absolute_protocol_relative_and_backslash_paths(self):
from litellm.proxy.management_endpoints.ui_sso import _is_same_origin_return_path
assert _is_same_origin_return_path("https://evil.example.com/authorize") is False
assert _is_same_origin_return_path("//evil.example.com/authorize") is False
assert _is_same_origin_return_path("/\\evil.example.com") is False
assert _is_same_origin_return_path("javascript:alert(1)") is False
assert _is_same_origin_return_path("") is False
class TestPersistReturnToCookieSharedHelper:
"""The single shared return_to helper used by EVERY sign-in branch (SSO / Okta / generic AND the
username/password form). It must be best-effort and NEVER raise — a bad return_to can never block
sign-in. Regression: the password form previously 400'd because it called _validate_return_to
directly (which raises for a non-matching absolute return_to when control_plane_url is set)."""
@staticmethod
def _cookie(resp) -> str:
return resp.headers.get("set-cookie", "")
def test_sets_cookie_for_same_origin_relative_path(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
resp = Response()
_persist_return_to_cookie(resp, "/mcp/authorize?client_id=llm_dcrc_abc")
assert "litellm_cp_return_to=" in self._cookie(resp)
def test_bad_absolute_with_control_plane_configured_does_not_raise_and_is_not_stored(self, monkeypatch):
"""THE regression: a non-matching absolute return_to with control_plane_url set must NOT raise
(it did, blocking the login form) and must NOT be stored — sign-in proceeds."""
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"}
)
resp = Response()
_persist_return_to_cookie(resp, "https://evil.example.com/steal") # must not raise
assert "litellm_cp_return_to=" not in self._cookie(resp)
def test_none_return_to_is_a_noop(self):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
resp = Response()
_persist_return_to_cookie(resp, None)
assert "litellm_cp_return_to=" not in self._cookie(resp)
def test_control_plane_matching_absolute_is_stored(self, monkeypatch):
from fastapi import Response
from litellm.proxy.management_endpoints.ui_sso import _persist_return_to_cookie
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", {"control_plane_url": "https://cp.example.com"}
)
resp = Response()
_persist_return_to_cookie(resp, "https://cp.example.com/ui?page=models")
assert "litellm_cp_return_to=" in self._cookie(resp)

View file

@ -49,9 +49,7 @@ def _install_login_mocks(monkeypatch, raise_on_auth: bool = False) -> None:
}
monkeypatch.setattr("litellm.proxy.auth.login_utils.authenticate_user", _fake_auth)
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.create_ui_token_object", _fake_token_object
)
monkeypatch.setattr("litellm.proxy.auth.login_utils.create_ui_token_object", _fake_token_object)
monkeypatch.setattr(ps, "master_key", "sk-test-master")
monkeypatch.setattr(ps, "general_settings", {})
monkeypatch.setattr(ps, "premium_user", False)
@ -69,9 +67,7 @@ def test_fallback_login_returns_html_form(client, monkeypatch):
body_lower = response.text.lower()
shape = {
"status": response.status_code,
"content_type_html": response.headers.get("content-type", "").startswith(
"text/html"
),
"content_type_html": response.headers.get("content-type", "").startswith("text/html"),
"has_form": "<form" in body_lower or "username" in body_lower,
}
assert shape == {
@ -88,9 +84,7 @@ def test_fallback_login_returns_html_form_with_ui_username_set(client, monkeypat
body_lower = response.text.lower()
shape = {
"status": response.status_code,
"content_type_html": response.headers.get("content-type", "").startswith(
"text/html"
),
"content_type_html": response.headers.get("content-type", "").startswith("text/html"),
"has_form_or_username": "<form" in body_lower or "username" in body_lower,
}
assert shape == {
@ -126,11 +120,7 @@ def test_fallback_login_invalid_method_405(client):
"""POST against the GET-only /fallback/login is rejected (error path)."""
response = client.post("/fallback/login")
assert response.status_code == 405
body = (
response.json()
if response.headers.get("content-type", "").startswith("application/json")
else {}
)
body = response.json() if response.headers.get("content-type", "").startswith("application/json") else {}
assert isinstance(body, dict)
@ -193,15 +183,15 @@ def test_v2_login_success_returns_token_and_redirect(client, monkeypatch):
json={"username": "admin", "password": "password"},
)
assert response.status_code == 200
assert normalize(
response.json(), volatile=frozenset({"token", "redirect_url"})
) == {"redirect_url": "<VOLATILE>", "token": "<VOLATILE>"}
assert normalize(response.json(), volatile=frozenset({"token", "redirect_url"})) == {
"redirect_url": "<VOLATILE>",
"token": "<VOLATILE>",
}
body = response.json()
set_cookie = response.headers.get("set-cookie", "")
shape = {
"redirect_url_has_ui": "/ui/" in body.get("redirect_url", ""),
"redirect_url_has_login_success": "login=success"
in body.get("redirect_url", ""),
"redirect_url_has_login_success": "login=success" in body.get("redirect_url", ""),
"token_in_body": bool(body.get("token")),
"token_cookie_set": "token=" in set_cookie,
}
@ -264,9 +254,7 @@ def test_v3_login_success_returns_code(client, monkeypatch):
from litellm.proxy import proxy_server as ps
_install_login_mocks(monkeypatch)
monkeypatch.setattr(
ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}
)
monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"})
# Force the local (non-redis) cache path
monkeypatch.setattr(ps, "redis_usage_cache", None)
fake_cache = MagicMock()
@ -301,9 +289,7 @@ def test_v3_login_authenticate_failure_500(client, monkeypatch):
from litellm.proxy import proxy_server as ps
_install_login_mocks(monkeypatch, raise_on_auth=True)
monkeypatch.setattr(
ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}
)
monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"})
response = client.post(
"/v3/login",
@ -337,9 +323,7 @@ def test_v3_login_exchange_missing_code_400(client, monkeypatch):
"""Error path: missing 'code' in body -> 400 with 'Missing' message."""
from litellm.proxy import proxy_server as ps
monkeypatch.setattr(
ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}
)
monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"})
response = client.post("/v3/login/exchange", json={})
assert response.status_code == 400
@ -352,9 +336,7 @@ def test_v3_login_exchange_invalid_code_401(client, monkeypatch):
"""Error path: code that isn't in cache -> 401 'Invalid or expired'."""
from litellm.proxy import proxy_server as ps
monkeypatch.setattr(
ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}
)
monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"})
monkeypatch.setattr(ps, "redis_usage_cache", None)
fake_cache = MagicMock()
fake_cache.async_get_cache = AsyncMock(return_value=None)
@ -372,9 +354,7 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc
"""Pin: valid code -> JSON {token, redirect_url} + token cookie + cache deleted (single-use)."""
from litellm.proxy import proxy_server as ps
monkeypatch.setattr(
ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"}
)
monkeypatch.setattr(ps, "general_settings", {"control_plane_url": "https://cp.example.invalid"})
monkeypatch.setattr(ps, "redis_usage_cache", None)
cached_payload = {
@ -388,9 +368,10 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc
response = client.post("/v3/login/exchange", json={"code": "valid-code"})
assert response.status_code == 200
assert normalize(
response.json(), volatile=frozenset({"token", "redirect_url"})
) == {"token": "<VOLATILE>", "redirect_url": "<VOLATILE>"}
assert normalize(response.json(), volatile=frozenset({"token", "redirect_url"})) == {
"token": "<VOLATILE>",
"redirect_url": "<VOLATILE>",
}
body = response.json()
set_cookie = response.headers.get("set-cookie", "")
shape = {
@ -405,3 +386,77 @@ def test_v3_login_exchange_success_returns_token_and_redirect(client, monkeypatc
"token_cookie_set": True,
"cache_deleted_once": True,
}
def test_login_form_honors_same_origin_return_to_cookie(client, monkeypatch):
"""The aggregate DCR connect flow preserves a same-origin return_to in the litellm_cp_return_to
cookie; /login must RESUME there after password sign-in instead of dead-ending at the dashboard."""
_install_login_mocks(monkeypatch)
return_to = "/mcp/authorize?client_id=llm_dcrc_abc&response_type=code"
response = client.post(
"/login",
data={"username": "admin", "password": "password"},
cookies={"litellm_cp_return_to": return_to},
follow_redirects=False,
)
assert response.status_code == 303
assert response.headers.get("location", "") == return_to # resumed the connect flow, not the dashboard
assert "token=" in response.headers.get("set-cookie", "")
def test_login_form_honors_control_plane_return_to_cookie(client, monkeypatch):
"""/login resumes through the SAME resumer the SSO callback uses, so it honors BOTH shapes
_persist_return_to_cookie is willing to store. Honoring only the relative one silently dropped
a control-plane return_to and landed the user on the dashboard."""
import litellm.proxy.proxy_server as ps
_install_login_mocks(monkeypatch)
monkeypatch.setitem(ps.general_settings, "control_plane_url", "https://cp.example.com")
response = client.post(
"/login",
data={"username": "admin", "password": "password"},
cookies={"litellm_cp_return_to": "https://cp.example.com/console"},
follow_redirects=False,
)
location = response.headers.get("location", "")
assert response.status_code == 303
assert location.startswith("https://cp.example.com/console")
# Cross-origin arm hands the JWT off via a one-time code rather than a cookie.
assert "code=" in location and "login=success" in location
assert "token=" not in response.headers.get("set-cookie", "")
def test_login_form_survives_stale_control_plane_return_to(client, monkeypatch):
"""A stale one-shot cookie must NEVER fail a completed sign-in. The resumer rejects a return_to
that no longer matches control_plane_url (a config change between the cookie's write and this
read); the user has already authenticated, so land on the dashboard instead of erroring."""
import litellm.proxy.proxy_server as ps
_install_login_mocks(monkeypatch)
monkeypatch.setitem(ps.general_settings, "control_plane_url", "https://new-cp.example.com")
response = client.post(
"/login",
data={"username": "admin", "password": "password"},
cookies={"litellm_cp_return_to": "https://old-cp.example.com/console"},
follow_redirects=False,
)
assert response.status_code == 303, "login must not break on a stale return_to cookie"
location = response.headers.get("location", "")
assert "old-cp.example.com" not in location
assert "/ui/" in location
def test_login_form_ignores_open_redirect_return_to(client, monkeypatch):
"""A non-same-origin return_to (open-redirect attempt) is rejected — /login falls back to the
dashboard rather than honoring an absolute/foreign URL."""
_install_login_mocks(monkeypatch)
response = client.post(
"/login",
data={"username": "admin", "password": "password"},
cookies={"litellm_cp_return_to": "https://evil.example.com/steal"},
follow_redirects=False,
)
assert response.status_code == 303
location = response.headers.get("location", "")
assert "evil.example.com" not in location
assert "/ui/" in location # dashboard fallback

View file

@ -1628,6 +1628,226 @@ async def test_ui_view_spend_logs_date_range_filter(client, monkeypatch):
assert data["data"][0]["id"] == "log2"
@pytest.mark.asyncio
async def test_ui_view_spend_logs_request_id_lookup_ignores_date_window(
client, monkeypatch
):
"""
LIT-3981: a request_id lookup on the UI route resolves across all time even
when the caller sends a date window that excludes the log (the dashboard
always sends a window). The window is dropped and request_id alone scopes
the query. Pre-fix the window was always applied, so an id from an older
page returned nothing.
"""
today = datetime.datetime.now(timezone.utc)
mock_spend_logs = [
{
"id": "log_old",
"request_id": "req-old",
"api_key": "sk-test-key",
"user": "test_user_1",
"team_id": "team1",
"spend": 0.05,
"startTime": (today - datetime.timedelta(days=90)).isoformat(),
"model": "gpt-4",
},
]
captured: dict = {}
def filter_fn(where):
captured["where"] = where
rows = _filter_logs_by_date_range(mock_spend_logs, where)
if where.get("request_id"):
rows = [r for r in rows if r["request_id"] == where["request_id"]]
return rows
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
make_ui_spend_logs_mock_prisma(mock_spend_logs, filter_fn),
)
# A 5-day window that EXCLUDES the 90-day-old log, as the dashboard sends.
start_date = (today - datetime.timedelta(days=5)).strftime("%Y-%m-%d %H:%M:%S")
end_date = today.strftime("%Y-%m-%d %H:%M:%S")
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
response = client.get(
"/spend/logs/ui",
params={
"request_id": "req-old",
"start_date": start_date,
"end_date": end_date,
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
assert data["total"] == 1
assert data["data"][0]["request_id"] == "req-old"
# Query dropped the time window and scoped solely by the primary key.
assert "startTime" not in captured["where"]
assert captured["where"]["request_id"] == "req-old"
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_requires_dates_without_request_id(
client, monkeypatch
):
"""The date window stays mandatory on the UI route when no request_id is set."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
make_ui_spend_logs_mock_prisma([], lambda where: []),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
response = client.get(
"/spend/logs/ui", headers={"Authorization": "Bearer sk-test"}
)
assert response.status_code == 400
assert "date" in response.text.lower()
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_logs_v2_still_requires_dates_with_request_id(client, monkeypatch):
"""The public /spend/logs/v2 contract is unchanged: dates remain required even
when request_id is supplied. Only the internal UI route relaxes the window."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
make_ui_spend_logs_mock_prisma([], lambda where: []),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
response = client.get(
"/spend/logs/v2",
params={"request_id": "req-old"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 400
assert "date" in response.text.lower()
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_request_id_blocks_non_owner(client, monkeypatch):
"""A non-admin looking up a request_id they do not own is rejected (403), so
the relaxed date window cannot read another tenant's log by id."""
class _ForeignRow:
user = "other_user"
team_id = None
class _SpendLogs:
async def find_unique(self, where, include=None):
return _ForeignRow()
class _DB:
def __init__(self):
self.litellm_spendlogs = _SpendLogs()
class _Prisma:
def __init__(self):
self.db = _DB()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", _Prisma())
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1"
)
try:
response = client.get(
"/spend/logs/ui",
params={"request_id": "foreign-req"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 403
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_request_id_owner_scoped_by_id_only(
client, monkeypatch
):
"""A non-admin owner looking up their own request_id resolves across all time.
The ownership check authorizes the single row, so the query drops both the date
window and the general user/team scoping and filters by the primary key alone;
without that skip an internal user would have a `user`/`OR` clause added."""
today = datetime.datetime.now(timezone.utc)
mock_spend_logs = [
{
"id": "log_old",
"request_id": "req-old",
"api_key": "sk-test-key",
"user": "user_1",
"team_id": "team1",
"spend": 0.05,
"startTime": (today - datetime.timedelta(days=90)).isoformat(),
"model": "gpt-4",
},
]
captured: dict = {}
def filter_fn(where):
captured["where"] = where
rows = _filter_logs_by_date_range(mock_spend_logs, where)
if where.get("request_id"):
rows = [r for r in rows if r["request_id"] == where["request_id"]]
return rows
mock_prisma = make_ui_spend_logs_mock_prisma(mock_spend_logs, filter_fn)
class _OwnedRow:
user = "user_1"
team_id = "team1"
async def _find_unique(where, include=None):
return _OwnedRow()
mock_prisma.db.find_unique = _find_unique
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
# A 5-day window that EXCLUDES the 90-day-old log, as the dashboard sends.
start_date = (today - datetime.timedelta(days=5)).strftime("%Y-%m-%d %H:%M:%S")
end_date = today.strftime("%Y-%m-%d %H:%M:%S")
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1"
)
try:
response = client.get(
"/spend/logs/ui",
params={
"request_id": "req-old",
"start_date": start_date,
"end_date": end_date,
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
assert data["total"] == 1
assert data["data"][0]["request_id"] == "req-old"
assert "startTime" not in captured["where"]
assert captured["where"]["request_id"] == "req-old"
assert "user" not in captured["where"]
assert "OR" not in captured["where"]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_unauthorized(client):
# Test without authorization header

View file

@ -109,6 +109,143 @@ def test_get_logging_payload_does_not_map_missing_or_zero_cached_tokens(prompt_t
assert "cache_read_input_tokens" not in additional_usage_values
def test_get_logging_payload_maps_openai_cache_write_tokens_to_cache_creation_input_tokens():
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=1000,
completion_tokens=2,
total_tokens=1002,
prompt_tokens_details={"cached_tokens": 0, "cache_write_tokens": 800},
)
)
assert additional_usage_values["cache_creation_input_tokens"] == 800
assert additional_usage_values["prompt_tokens_details"]["cache_write_tokens"] == 800
def test_get_logging_payload_preserves_anthropic_cache_creation_input_tokens():
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=1000,
completion_tokens=2,
total_tokens=1002,
cache_creation_input_tokens=300,
)
)
assert additional_usage_values["cache_creation_input_tokens"] == 300
@pytest.mark.parametrize(
"prompt_tokens_details",
[None, {"cached_tokens": 100}, {"cached_tokens": 100, "cache_write_tokens": 0}],
)
def test_get_logging_payload_does_not_map_missing_or_zero_cache_write_tokens(prompt_tokens_details):
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=10,
completion_tokens=2,
total_tokens=12,
prompt_tokens_details=prompt_tokens_details,
)
)
assert "cache_creation_input_tokens" not in additional_usage_values
def _make_standard_logging_payload_with_usage_object(usage_object: dict) -> StandardLoggingPayload:
return StandardLoggingPayload(
id="test-id-responses",
call_type="responses",
stream=False,
response_cost=0.02,
status="success",
total_tokens=1010,
prompt_tokens=1000,
completion_tokens=10,
startTime=1234567890.0,
endTime=1234567891.0,
completionStartTime=None,
model_map_information=StandardLoggingModelInformation(model_map_key="gpt-5.6", model_map_value=None),
model="gpt-5.6",
model_id="model-123",
model_group="openai",
custom_llm_provider="openai",
api_base="https://api.openai.com",
metadata=StandardLoggingMetadata(
user_api_key_hash="test_hash",
user_api_key_alias=None,
user_api_key_team_id=None,
user_api_key_org_id=None,
user_api_key_user_id=None,
user_api_key_team_alias=None,
spend_logs_metadata=None,
requester_ip_address=None,
requester_metadata=None,
user_api_key_end_user_id=None,
usage_object=usage_object,
),
cache_hit=False,
cache_key=None,
saved_cache_cost=0.0,
request_tags=[],
end_user=None,
requester_ip_address=None,
messages=[],
response={},
error_str=None,
model_parameters={},
hidden_params=StandardLoggingHiddenParams(
model_id="model-123",
cache_key=None,
api_base="https://api.openai.com",
response_cost="0.02",
litellm_overhead_time_ms=None,
additional_headers=None,
batch_models=None,
litellm_model_name=None,
usage_object=None,
),
)
def test_get_logging_payload_maps_responses_api_cache_write_tokens_from_usage_object():
"""Responses API (/v1/responses) usage is not chat-Usage-shaped, so
additional_usage_values can't derive cache tokens from response_obj.usage.
The Admin UI Logs "Cache Creation Tokens" row reads
additional_usage_values.cache_creation_input_tokens, so it must be filled
from the normalized standard_logging usage_object (LIT-4633)."""
standard_logging_payload = _make_standard_logging_payload_with_usage_object(
usage_object={
"prompt_tokens": 1000,
"completion_tokens": 10,
"total_tokens": 1010,
"prompt_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 800, "cache_creation_tokens": 800},
}
)
payload = get_logging_payload(
kwargs={
"model": "gpt-5.6",
"call_type": "responses",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
"standard_logging_object": standard_logging_payload,
},
response_obj={
"id": "resp-test",
"usage": {
"input_tokens": 1000,
"output_tokens": 10,
"total_tokens": 1010,
"input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 800},
},
},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
additional_usage_values = json.loads(payload["metadata"])["additional_usage_values"]
assert additional_usage_values["cache_creation_input_tokens"] == 800
def test_sanitize_request_body_for_spend_logs_payload_basic():
request_body = {
"messages": [{"role": "user", "content": "Hello, how are you?"}],

View file

@ -3,6 +3,7 @@ from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
import litellm
from litellm.caching.dual_cache import DualCache
@ -1562,6 +1563,100 @@ async def test_should_skip_reservation_when_counter_increment_fails(
)
@pytest.mark.asyncio
async def test_should_raise_503_when_counter_increment_fails_and_fail_closed(
spend_counter_state,
monkeypatch,
):
"""#33923: with fail_closed_budget_enforcement on, a failed reservation write
must reject instead of silently degrading to read-time-only enforcement."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-budget-reserve-fail-closed",
spend=0.0,
max_budget=1.0,
)
async def fail_increment_cache(*args, **kwargs):
raise RuntimeError("counter unavailable")
monkeypatch.setattr(counter_cache, "async_increment_cache", fail_increment_cache)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=0.5,
):
with pytest.raises(HTTPException) as exc_info:
await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
fail_closed_budget_enforcement=True,
)
assert exc_info.value.status_code == 503
assert (
counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-reserve-fail-closed"
)
is None
)
@pytest.mark.asyncio
async def test_fail_closed_releases_earlier_counters_before_503(
spend_counter_state,
):
"""#33923: when a later counter's reservation write fails in strict mode, the
counters that already reserved must be released before the 503 propagates."""
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
valid_token = UserAPIKeyAuth(
token="key-budget-fail-closed-release",
spend=0.0,
max_budget=1.0,
budget_limits=[
{
"budget_duration": "1h",
"max_budget": 1.0,
}
],
)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
return_value=0.5,
):
with pytest.raises(HTTPException) as exc_info:
await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=valid_token,
team_object=None,
user_object=None,
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
fail_closed_budget_enforcement=True,
)
assert exc_info.value.status_code == 503
assert (
counter_cache.in_memory_cache.get_cache(
key="spend:key:key-budget-fail-closed-release"
)
== 0.0
)
@pytest.mark.asyncio
async def test_should_skip_reservation_when_counter_initialization_fails(
spend_counter_state,

View file

@ -5226,3 +5226,198 @@ async def test_add_litellm_data_to_request_unions_metadata_tags_with_header_tags
tags = updated["litellm_metadata"]["tags"]
assert "header-tag" in tags
assert "body-tag" in tags
def _make_chat_request_mock() -> MagicMock:
return _make_request_mock("/v1/chat/completions", {"Content-Type": "application/json"})
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_clobbers_caller_supplied_user(monkeypatch):
"""The flag exists so providers can ban by a tamper-proof id; a caller-chosen
`user` must never survive, and the raw sk- key must never be forwarded."""
from litellm.proxy._types import hash_token
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
raw_key = "sk-overwrite-user-test-1234"
user_api_key_dict = UserAPIKeyAuth(api_key=raw_key)
user_api_key_dict.via_virtual_key = True
data = {"model": "gpt-4o", "user": "attacker-chosen-id"}
updated_data = await add_litellm_data_to_request(
data=data,
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == hash_token(raw_key)
assert updated_data["user"] != "attacker-chosen-id"
assert raw_key not in updated_data["user"]
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_sets_user_when_absent(monkeypatch):
from litellm.proxy._types import hash_token
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
raw_key = "sk-overwrite-user-test-5678"
user_api_key_dict = UserAPIKeyAuth(api_key=raw_key)
user_api_key_dict.via_virtual_key = True
data = {"model": "gpt-4o"}
updated_data = await add_litellm_data_to_request(
data=data,
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == hash_token(raw_key)
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_disabled_preserves_caller_user():
assert litellm.overwrite_user_with_key_hash is False
user_api_key_dict = UserAPIKeyAuth(api_key="sk-overwrite-user-test-9999")
user_api_key_dict.via_virtual_key = True
data = {"model": "gpt-4o", "user": "caller-chosen-id"}
updated_data = await add_litellm_data_to_request(
data=data,
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_skips_custom_auth_credential(monkeypatch):
"""Custom-auth credentials are not sk-prefixed or JWTs, so UserAPIKeyAuth stores
them raw; the stamp must skip them entirely so auth material never leaks."""
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
raw_credential = "my-custom-auth-credential-abc123"
user_api_key_dict = UserAPIKeyAuth(api_key=raw_credential)
assert user_api_key_dict.api_key == raw_credential
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_skips_jwt_auth(monkeypatch):
"""A hashed JWT rotates on every token re-issue, so it is useless as a stable
ban id; JWT-authenticated requests are not stamped."""
from litellm.proxy._types import hash_token
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
hashed_jwt = f"hashed-jwt-{hash_token('some-jwt-token')}"
user_api_key_dict = UserAPIKeyAuth(api_key=hashed_jwt)
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_skips_hex_shaped_custom_credential(monkeypatch):
"""A custom-auth credential that happens to be 64 hex chars is indistinguishable
from a key hash by shape alone; only the server-set via_virtual_key marker may
authorize stamping, so this raw credential must never be forwarded."""
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
hex_shaped_credential = "a" * 64
user_api_key_dict = UserAPIKeyAuth(api_key=hex_shaped_credential)
assert user_api_key_dict.api_key == hex_shaped_credential
assert user_api_key_dict.via_virtual_key is False
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"
def test_via_virtual_key_cannot_be_forged_from_validated_input():
from_kwargs = UserAPIKeyAuth(api_key="b" * 64, via_virtual_key=True)
assert from_kwargs.via_virtual_key is False
from_dict = UserAPIKeyAuth.model_validate({"api_key": "b" * 64, "via_virtual_key": True})
assert from_dict.via_virtual_key is False
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_stamps_master_key_alias(monkeypatch):
"""Master-key requests carry the stable alias instead of a hash (so the master
key never propagates anywhere); the alias is the stampable id for them."""
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS)
user_api_key_dict.via_virtual_key = True
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "attacker-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == LITELLM_PROXY_MASTER_KEY_ALIAS
@pytest.mark.asyncio
async def test_overwrite_user_with_key_hash_rejects_alias_without_marker(monkeypatch):
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
monkeypatch.setattr(litellm, "overwrite_user_with_key_hash", True)
user_api_key_dict = UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS)
assert user_api_key_dict.via_virtual_key is False
updated_data = await add_litellm_data_to_request(
data={"model": "gpt-4o", "user": "caller-chosen-id"},
request=_make_chat_request_mock(),
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
assert updated_data["user"] == "caller-chosen-id"

View file

@ -127,11 +127,15 @@ def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch):
general_settings={},
premium_user=False,
)
mock_jwt_encode.assert_called_once_with(
{"user_id": "test-user"},
"test-master-key",
algorithm="HS256",
)
mock_jwt_encode.assert_called_once()
payload, secret = mock_jwt_encode.call_args.args
# The UI session token carries a bounded-lifetime `exp` claim (dynamic timestamp), alongside
# the user_id; assert its presence rather than an exact expiry value.
assert payload["user_id"] == "test-user"
assert isinstance(payload.get("exp"), int) and payload["exp"] > 0
assert set(payload.keys()) == {"user_id", "exp"}
assert secret == "test-master-key"
assert mock_jwt_encode.call_args.kwargs == {"algorithm": "HS256"}
def test_login_v2_returns_json_on_proxy_exception(monkeypatch):

View file

@ -369,6 +369,32 @@ class TestResponseAPILoggingUtils:
assert result.completion_tokens_details.image_tokens == 272
assert result.completion_tokens_details.text_tokens == 100
def test_transform_response_api_usage_maps_cache_write_tokens(self):
"""Responses API (/v1/responses) cache-write tokens must survive the usage transform.
gpt-5.6 returns usage.input_tokens_details.cache_write_tokens (an extra field
not typed on InputTokensDetails). Before the fix the transform rebuilt the token
details and dropped it, leaving the cache-creation metric empty (LIT-4633).
"""
usage = {
"input_tokens": 10062,
"output_tokens": 16,
"total_tokens": 10078,
"input_tokens_details": {
"cached_tokens": 0,
"cache_write_tokens": 10059,
},
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cache_write_tokens == 10059
assert result.prompt_tokens_details.cache_creation_tokens == 10059
assert result.prompt_tokens_details.cached_tokens == 0
def test_transform_response_api_usage_mixed_details(self):
"""Test transformation handles mixed token details (cached + image + audio)."""
# Setup - hypothetical usage with mixed token types
@ -461,64 +487,6 @@ class TestResponseAPILoggingUtils:
assert result.completion_tokens_details.text_tokens == 20
assert result.completion_tokens_details.audio_tokens is None
def test_transform_response_api_usage_maps_cache_write_tokens_dict(self):
"""Regression for LIT-4725 / #33772: the Responses API reports cache writes under
input_tokens_details.cache_write_tokens. The chat-shaped usage must carry them on
cache_creation_tokens so cost is computed identically to /chat/completions.
model_construct keeps input_tokens_details a raw dict so the dict branch runs."""
from litellm.types.llms.openai import ResponseAPIUsage
usage = ResponseAPIUsage.model_construct(
input_tokens=10_000,
output_tokens=20,
total_tokens=10_020,
input_tokens_details={"cached_tokens": 2_000, "cache_write_tokens": 8_000},
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cached_tokens == 2_000
assert result.prompt_tokens_details.cache_creation_tokens == 8_000
assert getattr(result.prompt_tokens_details, "cache_write_tokens", None) is None
def test_transform_response_api_usage_cache_creation_tokens_precedence_dict(self):
"""When both cache_creation_tokens and cache_write_tokens are present, the explicit
cache_creation_tokens wins (they describe the same tokens under different names)."""
from litellm.types.llms.openai import ResponseAPIUsage
usage = ResponseAPIUsage.model_construct(
input_tokens=10_000,
output_tokens=20,
total_tokens=10_020,
input_tokens_details={"cache_creation_tokens": 5_000, "cache_write_tokens": 8_000},
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cache_creation_tokens == 5_000
def test_transform_response_api_usage_maps_cache_write_tokens_object(self):
"""Object-path counterpart: a ResponseAPIUsage whose input_tokens_details object
carries cache_write_tokens must still land on cache_creation_tokens."""
from litellm.types.llms.openai import InputTokensDetails, ResponseAPIUsage
input_tokens_details = InputTokensDetails(cached_tokens=2_000)
input_tokens_details.cache_write_tokens = 8_000
usage = ResponseAPIUsage(
input_tokens=10_000,
output_tokens=20,
total_tokens=10_020,
input_tokens_details=input_tokens_details,
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cached_tokens == 2_000
assert result.prompt_tokens_details.cache_creation_tokens == 8_000
class TestResponsesAPIProviderSpecificParams:
"""

View file

@ -12,6 +12,7 @@ from pydantic import BaseModel
import litellm
from litellm.cost_calculator import (
BaseTokenUsageProcessor,
RealtimeAPITokenUsageProcessor,
completion_cost,
cost_per_token,
@ -3479,3 +3480,32 @@ def test_batch_cost_calculator_cache_creation_falls_back_to_input_rate():
)
assert prompt_cost == pytest.approx((1000 * 3e-6 + 8000 * 3e-7 + 2000 * 3e-6) / 2)
def test_combine_usage_objects_sums_mirrored_cache_write_fields_once():
"""
cache_write_tokens and cache_creation_tokens mirror each other on
PromptTokensDetailsWrapper, so field-iterating aggregation must sum the pair
once: a single 50-token usage stays 50 and two combine to 100, not double.
"""
single = Usage(
prompt_tokens=100,
completion_tokens=10,
total_tokens=110,
prompt_tokens_details=PromptTokensDetailsWrapper(cache_write_tokens=50),
)
combined = BaseTokenUsageProcessor.combine_usage_objects([single])
assert combined.prompt_tokens_details is not None
assert combined.prompt_tokens_details.cache_write_tokens == 50
assert combined.prompt_tokens_details.cache_creation_tokens == 50
anthropic_style = Usage(
prompt_tokens=100,
completion_tokens=10,
total_tokens=110,
cache_creation_input_tokens=50,
)
combined_pair = BaseTokenUsageProcessor.combine_usage_objects([anthropic_style, anthropic_style])
assert combined_pair.prompt_tokens_details is not None
assert combined_pair.prompt_tokens_details.cache_write_tokens == 100
assert combined_pair.prompt_tokens_details.cache_creation_tokens == 100

View file

@ -17,7 +17,9 @@ from litellm.types.utils import (
Delta,
LlmProviders,
ModelResponseStream,
PromptTokensDetailsWrapper,
StreamingChoices,
Usage,
)
from litellm.utils import (
ProviderConfigManager,
@ -34,6 +36,57 @@ from litellm.utils import (
# Adds the parent directory to the system path
def test_usage_openai_cache_write_tokens_populates_both_names():
"""OpenAI reports cache-write tokens as prompt_tokens_details.cache_write_tokens.
The Usage constructor must expose it under both cache_write_tokens (canonical,
OpenAI naming) and cache_creation_tokens (legacy, Anthropic naming)."""
usage = Usage(
prompt_tokens=1000,
completion_tokens=10,
total_tokens=1010,
prompt_tokens_details={"cached_tokens": 0, "cache_write_tokens": 800},
)
assert usage.prompt_tokens_details.cache_write_tokens == 800
assert usage.prompt_tokens_details.cache_creation_tokens == 800
def test_usage_anthropic_cache_creation_maps_to_cache_write_tokens():
"""Anthropic/Bedrock report the top-level cache_creation_input_tokens field.
It must be normalized onto the OpenAI cache_write_tokens name as well as the
legacy cache_creation_tokens name."""
usage = Usage(
prompt_tokens=500,
completion_tokens=50,
total_tokens=550,
cache_creation_input_tokens=300,
cache_read_input_tokens=120,
)
assert usage.prompt_tokens_details.cache_write_tokens == 300
assert usage.prompt_tokens_details.cache_creation_tokens == 300
assert usage.prompt_tokens_details.cached_tokens == 120
def test_prompt_tokens_details_no_cache_write_tokens_when_absent():
"""A read-only cache hit (no cache write) must not surface cache-write fields."""
details = PromptTokensDetailsWrapper(cached_tokens=800)
assert details.cached_tokens == 800
assert not hasattr(details, "cache_write_tokens")
assert not hasattr(details, "cache_creation_tokens")
def test_prompt_tokens_details_cache_write_creation_stay_in_sync_on_assignment():
"""Assigning either name after construction must mirror to the other, so a
caller that sets only one field can't leave the pair silently out of sync."""
details = PromptTokensDetailsWrapper(cache_write_tokens=100)
assert details.cache_write_tokens == details.cache_creation_tokens == 100
details.cache_write_tokens = 250
assert details.cache_write_tokens == details.cache_creation_tokens == 250
details.cache_creation_tokens = 375
assert details.cache_write_tokens == details.cache_creation_tokens == 375
@pytest.fixture
def local_model_cost_map(monkeypatch):
original_model_cost = litellm.model_cost

View file

@ -37,16 +37,6 @@
"count": 1
}
},
"src/app/(dashboard)/agents/_components/AgentsPanel.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/agents/_components/AgentsTable.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/agents/_components/add_agent_form.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -71,9 +61,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/refs": {
"count": 3
},
@ -84,9 +71,6 @@
"src/app/(dashboard)/agents/_components/agent_cost_view.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/app/(dashboard)/agents/_components/agent_form_fields.tsx": {
@ -123,9 +107,6 @@
},
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/agents/_components/cost_config_fields.tsx": {
@ -1143,25 +1124,6 @@
"count": 1
}
},
"src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx": {
"max-params": {
"count": 1
},
"unused-imports/no-unused-imports": {
"count": 1
}
},
"src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx": {
"local/no-complex-jsx-arrow": {
"count": 2
},
"no-restricted-imports": {
"count": 2
},
"react-hooks/set-state-in-effect": {
"count": 3
}
},
"src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.test.tsx": {
"react/display-name": {
"count": 1
@ -1180,6 +1142,11 @@
"count": 1
}
},
"src/app/(dashboard)/models-and-endpoints/layout.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/models-and-endpoints/utils/modelDataTransformer.ts": {
"prefer-const": {
"count": 6
@ -2134,9 +2101,6 @@
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"no-restricted-syntax": {
"count": 3
},
@ -2371,11 +2335,6 @@
"count": 1
}
},
"src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/LicenseExpiryBanner.tsx": {
"no-restricted-imports": {
"count": 1
@ -2386,14 +2345,6 @@
"count": 1
}
},
"src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx": {
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/Navbar/BlogDropdown/BlogDropdown.test.tsx": {
"max-nested-callbacks": {
"count": 12
@ -2972,7 +2923,7 @@
"count": 1
},
"no-nested-ternary": {
"count": 7
"count": 6
}
},
"src/components/chat/MCPConnectPicker.tsx": {
@ -3486,17 +3437,6 @@
"count": 1
}
},
"src/components/model_dashboard/all_models_table.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/model_filters.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -3544,17 +3484,6 @@
"count": 1
}
},
"src/components/molecules/filter.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-nested-ternary": {
"count": 2
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/molecules/message_manager.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -3563,28 +3492,6 @@
"count": 2
}
},
"src/components/molecules/models/columns.test.tsx": {
"no-restricted-imports": {
"count": 1
},
"react/display-name": {
"count": 1
}
},
"src/components/molecules/models/columns.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"max-params": {
"count": 1
},
"no-nested-ternary": {
"count": 2
},
"no-restricted-imports": {
"count": 2
}
},
"src/components/molecules/notifications_manager.test.tsx": {
"no-restricted-imports": {
"count": 1
@ -3689,9 +3596,6 @@
},
"no-restricted-imports": {
"count": 3
},
"unused-imports/no-unused-imports": {
"count": 1
}
},
"src/components/page_utils.test.ts": {
@ -4182,6 +4086,11 @@
"count": 1
}
},
"src/components/ui/hover-card.tsx": {
"local/filename-pascal-case": {
"count": 1
}
},
"src/components/ui/input-group.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -4454,17 +4363,6 @@
"count": 2
}
},
"src/components/view_logs/LogsTableToolbar.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-nested-ternary": {
"count": 4
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/view_logs/ToolsSection/FormattedToolView.tsx": {
"no-restricted-imports": {
"count": 1
@ -4493,9 +4391,6 @@
"src/components/view_logs/columns.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/view_logs/index.tsx": {
@ -4504,9 +4399,6 @@
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
},
"src/components/view_logs/log_filter_logic.tsx": {

View file

@ -140,8 +140,9 @@ describe("AgentsPanel", () => {
await user.click(await screen.findByTestId("agent-actions-agent-9"));
await user.click(await screen.findByTestId("agent-action-delete"));
const modal = await screen.findByRole("dialog");
await user.click(within(modal).getByRole("button", { name: /^delete$/i }));
const confirmPrompt = await screen.findByText(/are you sure you want to delete agent: Doomed Agent\?/i);
const confirmDialog = confirmPrompt.closest('[role="dialog"],[role="alertdialog"]') as HTMLElement;
await user.click(within(confirmDialog).getByRole("button", { name: /^delete$/i }));
await waitFor(() => {
expect(networking.deleteAgentCall).toHaveBeenCalledWith("test-token", "agent-9");

View file

@ -1,6 +1,5 @@
import React, { useState, useEffect } from "react";
import { Modal, Alert } from "antd";
import { Plus } from "lucide-react";
import { Info, Plus } from "lucide-react";
import { getAgentsList, deleteAgentCall } from "@/components/networking";
import AddAgentForm from "./add_agent_form";
import { isAdminRole } from "@/utils/roles";
@ -9,6 +8,16 @@ import AgentsTable from "./AgentsTable";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { Agent } from "@/components/agents/types";
import { Team } from "@/components/key_team_helpers/key_list";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import {
AlertDialog,
AlertDialogCancel,
AlertDialogContent,
AlertDialogDescription,
AlertDialogFooter,
AlertDialogHeader,
AlertDialogTitle,
} from "@/components/ui/alert-dialog";
import { Button } from "@/components/ui/button";
interface AgentsPanelProps {
@ -130,17 +139,18 @@ const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole, teams
<div className="w-full mx-auto flex-auto overflow-y-auto m-8 p-2">
<div className="flex flex-col gap-2 mb-4">
<h1 className="text-2xl font-bold">Agents</h1>
<p className="text-sm text-gray-600">
<p className="text-sm text-muted-foreground">
List of A2A-spec agents that are available to be used in your organization. Go to AI Hub, to make agents
public.
</p>
<Alert
message="Why do agents need keys?"
description="Keys scope access to an agent and allow it to call MCP tools. Assign a key when creating an agent or from the Virtual Keys page."
type="info"
showIcon
className="mb-3"
/>
<Alert className="mb-3">
<Info />
<AlertTitle>Why do agents need keys?</AlertTitle>
<AlertDescription>
Keys scope access to an agent and allow it to call MCP tools. Assign a key when creating an agent or from
the Virtual Keys page.
</AlertDescription>
</Alert>
{isAdmin && (
<div className="mt-2 flex items-center gap-4">
<Button onClick={handleAddAgent} disabled={!accessToken}>
@ -180,18 +190,27 @@ const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole, teams
/>
{agentToDelete && (
<Modal
title="Delete Agent"
open={agentToDelete !== null}
onOk={handleDeleteConfirm}
onCancel={handleDeleteCancel}
confirmLoading={isDeleting}
okText="Delete"
okButtonProps={{ danger: true }}
<AlertDialog
open
onOpenChange={(open) => {
if (!open) handleDeleteCancel();
}}
>
<p>Are you sure you want to delete agent: {agentToDelete.name}?</p>
<p>This action cannot be undone.</p>
</Modal>
<AlertDialogContent>
<AlertDialogHeader>
<AlertDialogTitle>Delete Agent</AlertDialogTitle>
<AlertDialogDescription>
Are you sure you want to delete agent: {agentToDelete.name}? This action cannot be undone.
</AlertDialogDescription>
</AlertDialogHeader>
<AlertDialogFooter>
<AlertDialogCancel>Cancel</AlertDialogCancel>
<Button variant="destructive" onClick={handleDeleteConfirm} disabled={isDeleting}>
Delete
</Button>
</AlertDialogFooter>
</AlertDialogContent>
</AlertDialog>
)}
</div>
);

View file

@ -1,13 +1,13 @@
"use client";
import { SortingState } from "@tanstack/react-table";
import { Tooltip, Switch } from "antd";
import { CheckCircleOutlined } from "@ant-design/icons";
import { Bot } from "lucide-react";
import { Bot, CircleCheck } from "lucide-react";
import React, { useMemo, useState } from "react";
import { Agent } from "@/components/agents/types";
import { DataTable } from "@/components/shared/DataTable";
import { Switch } from "@/components/ui/switch";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { getAgentsTableColumns } from "./AgentsTableColumns";
@ -67,18 +67,27 @@ const AgentsTable: React.FC<AgentsTableProps> = ({
size="compact"
toolbar={() => (
<div className="flex items-center justify-end">
<Tooltip title="When enabled, only agents with reachable URLs are shown">
<div className="flex items-center gap-2">
<CheckCircleOutlined className={healthCheckEnabled ? "text-green-500" : "text-muted-foreground"} />
<span className="text-sm text-muted-foreground">Health Check</span>
<Switch
size="small"
checked={healthCheckEnabled}
onChange={onHealthCheckToggle}
loading={isHealthCheckLoading}
<TooltipProvider delay={300}>
<Tooltip>
<TooltipTrigger
render={
<div className="flex items-center gap-2">
<CircleCheck
className={healthCheckEnabled ? "size-4 text-green-500" : "size-4 text-muted-foreground"}
/>
<span className="text-sm text-muted-foreground">Health Check</span>
<Switch
size="sm"
checked={healthCheckEnabled}
onCheckedChange={onHealthCheckToggle}
disabled={isHealthCheckLoading}
/>
</div>
}
/>
</div>
</Tooltip>
<TooltipContent>When enabled, only agents with reachable URLs are shown</TooltipContent>
</Tooltip>
</TooltipProvider>
</div>
)}
/>

View file

@ -126,10 +126,7 @@ describe("AgentCardDiscovery", () => {
expect(initialSelection.upstream_url).toBe("https://upstream.example.com");
expect(initialSelection.selected_card.skills).toHaveLength(2);
const summarizeLabel = screen.getByText("Summarize").closest("label");
expect(summarizeLabel).toBeTruthy();
const summarizeCheckbox = summarizeLabel!.querySelector("input[type='checkbox']") as HTMLInputElement;
await user.click(summarizeCheckbox);
await user.click(screen.getByRole("checkbox", { name: /Summarize/i }));
await waitFor(() => {
const latest = onApply.mock.calls.at(-1)?.[0];

View file

@ -1,18 +1,20 @@
"use client";
import React, { useCallback, useEffect, useMemo, useRef, useState } from "react";
import { Alert, Button, Checkbox, Collapse, Empty, Input, Space, Spin, Switch, Tag, Tooltip, Typography } from "antd";
// Empty is used in the skills panel below.
import {
CheckCircleTwoTone,
InfoCircleOutlined,
LinkOutlined,
ReloadOutlined,
SearchOutlined,
} from "@ant-design/icons";
import { ChevronDown, CircleAlert, CircleCheck, Info, Link as LinkIcon, RotateCw, Search, X } from "lucide-react";
import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer";
import { DiscoveredAgentCard, discoverAgentCardCall } from "@/components/networking";
import { Alert, AlertAction, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Checkbox } from "@/components/ui/checkbox";
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
import { Input } from "@/components/ui/input";
import { Switch } from "@/components/ui/switch";
import { Textarea } from "@/components/ui/textarea";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import {
ALLOWED_CAPABILITY_KEYS,
selectionsFromSavedAgentCard,
@ -20,9 +22,6 @@ import {
skillId,
} from "./agent_discovery_utils";
const { Text, Paragraph } = Typography;
const { Panel } = Collapse;
const DISCOVERY_DEBOUNCE_WAIT_MS = 400;
export interface DiscoveredAgentCardSelection {
@ -243,102 +242,115 @@ const AgentCardDiscovery: React.FC<AgentCardDiscoveryProps> = ({
const skillCount = card?.skills?.length ?? 0;
const selectedSkillCount = selectedSkillIds.size;
const renderDiscoverIcon = () => {
if (loading) return <UiLoadingSpinner className="size-4" />;
if (card) return <RotateCw />;
return <Search />;
};
const discoverLabel = card ? "Re-discover" : "Discover";
return (
<div className="border border-gray-200 rounded-lg p-4 bg-gray-50 mb-4">
<div className="flex items-center gap-2 mb-2">
<LinkOutlined className="text-indigo-600" />
<Text strong>Discover from agent URL</Text>
<Tooltip title="LiteLLM will fetch /.well-known/agent-card.json from this URL and let you pick which skills and capabilities to expose through the proxy.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
<div className="mb-4 rounded-lg border border-border bg-muted/50 p-4">
<div className="mb-2 flex items-center gap-2">
<LinkIcon className="size-4 text-primary" />
<span className="text-sm font-medium text-foreground">Discover from agent URL</span>
<TooltipProvider delay={300}>
<Tooltip>
<TooltipTrigger
render={
<span className="inline-flex text-muted-foreground">
<Info className="size-4" />
</span>
}
/>
<TooltipContent>
LiteLLM will fetch /.well-known/agent-card.json from this URL and let you pick which skills and
capabilities to expose through the proxy.
</TooltipContent>
</Tooltip>
</TooltipProvider>
</div>
{isParentDriven ? (
<>
<Paragraph className="text-xs text-gray-500 mb-2">
<p className="mb-2 text-xs text-muted-foreground">
Using the connection details you entered above. We&apos;ll fetch:
</Paragraph>
<div className="bg-white border border-gray-200 rounded-sm px-3 py-2 mb-3 font-mono text-xs text-gray-700 break-all">
</p>
<div className="mb-3 rounded-sm border border-border bg-background px-3 py-2 font-mono text-xs break-all text-foreground">
{discoveryRequest!.display_url || effectiveUrl || (
<span className="text-gray-400 italic">Fill in the fields above first</span>
<span className="text-muted-foreground italic">Fill in the fields above first</span>
)}
</div>
<div className="flex justify-end">
<Button
type="primary"
icon={card ? <ReloadOutlined /> : <SearchOutlined />}
loading={loading}
onClick={handleDiscover}
disabled={!effectiveUrl.trim()}
>
{card ? "Re-discover" : "Discover"}
<Button onClick={handleDiscover} disabled={loading || !effectiveUrl.trim()}>
{renderDiscoverIcon()}
{discoverLabel}
</Button>
</div>
</>
) : (
<>
<Paragraph className="text-xs text-gray-500 mb-3">
<p className="mb-3 text-xs text-muted-foreground">
Paste the upstream agent&apos;s base URL. We&apos;ll try <code>/.well-known/agent-card.json</code>,{" "}
<code>/.well-known/agent.json</code>, and <code>/agent.json</code> in order.
</Paragraph>
</p>
<Space.Compact style={{ width: "100%" }}>
<div className="flex w-full items-center gap-2">
<Input
placeholder="https://upstream-agent.example.com"
value={manualUrl}
onChange={(e) => setManualUrl(e.target.value)}
onPressEnter={handleDiscover}
allowClear
onKeyDown={(e) => {
if (e.key === "Enter") handleDiscover();
}}
disabled={loading}
/>
<Button
type="primary"
icon={card ? <ReloadOutlined /> : <SearchOutlined />}
loading={loading}
onClick={handleDiscover}
>
{card ? "Re-discover" : "Discover"}
<Button onClick={handleDiscover} disabled={loading}>
{renderDiscoverIcon()}
{discoverLabel}
</Button>
</Space.Compact>
</div>
</>
)}
{error && (
<Alert
className="mt-3"
type="error"
message="Discovery failed"
description={error}
showIcon
closable
onClose={() => setError(null)}
/>
<Alert variant="destructive" className="mt-3">
<CircleAlert />
<AlertTitle>Discovery failed</AlertTitle>
<AlertDescription>{error}</AlertDescription>
<AlertAction>
<Button variant="ghost" size="icon-xs" aria-label="Dismiss error" onClick={() => setError(null)}>
<X />
</Button>
</AlertAction>
</Alert>
)}
{loading && !card && (
<div className="flex items-center justify-center py-8">
<Spin />
<UiLoadingSpinner className="size-6 text-muted-foreground" />
</div>
)}
{card && (
<div className="mt-4 bg-white border border-gray-200 rounded-lg p-4">
<div className="flex items-center justify-between mb-3">
<Space>
<CheckCircleTwoTone twoToneColor="#52c41a" />
<Text strong>Upstream card loaded</Text>
{card.version && <Tag color="blue">v{card.version}</Tag>}
{card.provider?.organization && <Tag color="purple">{card.provider.organization}</Tag>}
</Space>
<div className="mt-4 rounded-lg border border-border bg-background p-4">
<div className="mb-3 flex flex-wrap items-center gap-2">
<CircleCheck className="size-4 text-green-600" />
<span className="text-sm font-medium text-foreground">Upstream card loaded</span>
{card.version && <Badge variant="secondary">v{card.version}</Badge>}
{card.provider?.organization && <Badge variant="secondary">{card.provider.organization}</Badge>}
</div>
<div className="grid grid-cols-1 md:grid-cols-2 gap-3 mb-4">
<div className="mb-4 grid grid-cols-1 gap-3 md:grid-cols-2">
<div>
<label className="text-xs font-medium text-gray-600 block mb-1">Name (shown to API clients)</label>
<label className="mb-1 block text-xs font-medium text-muted-foreground">
Name (shown to API clients)
</label>
<Input value={editedName} onChange={(e) => setEditedName(e.target.value)} placeholder="Agent name" />
</div>
<div>
<label className="text-xs font-medium text-gray-600 block mb-1">Description</label>
<Input.TextArea
<label className="mb-1 block text-xs font-medium text-muted-foreground">Description</label>
<Textarea
className="field-sizing-fixed min-h-0"
value={editedDescription}
onChange={(e) => setEditedDescription(e.target.value)}
rows={2}
@ -347,103 +359,114 @@ const AgentCardDiscovery: React.FC<AgentCardDiscoveryProps> = ({
</div>
</div>
<Collapse defaultActiveKey={["skills", "capabilities"]} ghost className="bg-transparent">
<Panel
key="skills"
header={
<Space>
<Text strong>Skills</Text>
<Tag>
{selectedSkillCount} / {skillCount} selected
</Tag>
</Space>
}
>
{skillCount === 0 ? (
<Empty image={Empty.PRESENTED_IMAGE_SIMPLE} description="Upstream card has no skills" />
) : (
<div className="space-y-2">
{(card.skills ?? []).map((skill, idx) => {
const id = skillId(skill, idx);
const checked = selectedSkillIds.has(id);
return (
<label
key={id}
className={`flex items-start gap-3 p-3 border rounded cursor-pointer transition-colors ${
checked ? "border-indigo-300 bg-indigo-50" : "border-gray-200 bg-white hover:border-gray-300"
}`}
>
<Checkbox checked={checked} onChange={(e) => toggleSkill(id, e.target.checked)} />
<div className="flex-1">
<div className="flex items-center gap-2 flex-wrap">
<Text strong>{skill.name || id}</Text>
{skill.id && <Tag style={{ marginLeft: 0 }}>{skill.id}</Tag>}
{(skill.tags ?? []).map((t: string) => (
<Tag key={t} color="geekblue">
{t}
</Tag>
))}
<div className="flex flex-col gap-4">
<Collapsible defaultOpen>
<div className="flex items-center gap-2">
<CollapsibleTrigger
render={
<button type="button" className="group flex items-center gap-2">
<ChevronDown className="size-4 text-muted-foreground transition-transform group-data-[panel-open]:rotate-180" />
<span className="text-sm font-medium text-foreground">Skills</span>
</button>
}
/>
<Badge variant="secondary">
{selectedSkillCount} / {skillCount} selected
</Badge>
</div>
<CollapsibleContent className="pt-2">
{skillCount === 0 ? (
<div className="py-6 text-center text-sm text-muted-foreground">Upstream card has no skills</div>
) : (
<div className="space-y-2">
{(card.skills ?? []).map((skill, idx) => {
const id = skillId(skill, idx);
const checked = selectedSkillIds.has(id);
return (
<label
key={id}
className={`flex cursor-pointer items-start gap-3 rounded border p-3 transition-colors ${
checked ? "border-primary/40 bg-primary/5" : "border-border bg-background hover:border-ring"
}`}
>
<Checkbox checked={checked} onCheckedChange={(next) => toggleSkill(id, next)} />
<div className="flex-1">
<div className="flex flex-wrap items-center gap-2">
<span className="text-sm font-medium text-foreground">{skill.name || id}</span>
{skill.id && <Badge variant="secondary">{skill.id}</Badge>}
{(skill.tags ?? []).map((t: string) => (
<Badge key={t} variant="outline">
{t}
</Badge>
))}
</div>
{skill.description && (
<p className="mt-1 line-clamp-2 text-xs text-muted-foreground">{skill.description}</p>
)}
</div>
{skill.description && (
<Paragraph
className="text-xs text-gray-500 mt-1 mb-0"
ellipsis={{ rows: 2, expandable: true, symbol: "more" }}
>
{skill.description}
</Paragraph>
)}
</label>
);
})}
</div>
)}
</CollapsibleContent>
</Collapsible>
<Collapsible defaultOpen>
<div className="flex items-center gap-2">
<CollapsibleTrigger
render={
<button type="button" className="group flex items-center gap-2">
<ChevronDown className="size-4 text-muted-foreground transition-transform group-data-[panel-open]:rotate-180" />
<span className="text-sm font-medium text-foreground">Capabilities</span>
</button>
}
/>
<TooltipProvider delay={300}>
<Tooltip>
<TooltipTrigger
render={
<span className="inline-flex text-muted-foreground">
<Info className="size-4" />
</span>
}
/>
<TooltipContent>
Only capabilities LiteLLM can faithfully proxy today are listed. Others (push notifications,
extensions) are coming soon.
</TooltipContent>
</Tooltip>
</TooltipProvider>
</div>
<CollapsibleContent className="pt-2">
<div className="space-y-2">
{ALLOWED_CAPABILITY_KEYS.map((key) => {
const upstreamHas = Boolean(card.capabilities?.[key]);
return (
<div
key={key}
className="flex items-center justify-between rounded-sm border border-border bg-background p-2"
>
<div className="flex items-center gap-2">
<span className="text-sm font-medium text-foreground capitalize">{key}</span>
{!upstreamHas && <Badge variant="outline">not advertised upstream</Badge>}
</div>
</label>
<Switch
checked={Boolean(selectedCapabilities[key])}
onCheckedChange={(checked) =>
setSelectedCapabilities((prev) => ({
...prev,
[key]: checked,
}))
}
/>
</div>
);
})}
</div>
)}
</Panel>
<Panel
key="capabilities"
header={
<Space>
<Text strong>Capabilities</Text>
<Tooltip title="Only capabilities LiteLLM can faithfully proxy today are listed. Others (push notifications, extensions) are coming soon.">
<InfoCircleOutlined className="text-gray-400" />
</Tooltip>
</Space>
}
>
<div className="space-y-2">
{ALLOWED_CAPABILITY_KEYS.map((key) => {
const upstreamHas = Boolean(card.capabilities?.[key]);
return (
<div
key={key}
className="flex items-center justify-between p-2 border border-gray-200 rounded-sm bg-white"
>
<div>
<Text strong className="capitalize">
{key}
</Text>
{!upstreamHas && (
<Tag className="ml-2" color="default">
not advertised upstream
</Tag>
)}
</div>
<Switch
checked={Boolean(selectedCapabilities[key])}
onChange={(checked) =>
setSelectedCapabilities((prev) => ({
...prev,
[key]: checked,
}))
}
/>
</div>
);
})}
</div>
</Panel>
</Collapse>
</CollapsibleContent>
</Collapsible>
</div>
</div>
)}
</div>

View file

@ -0,0 +1,54 @@
import React from "react";
import { describe, it, expect } from "vitest";
import { screen } from "@testing-library/react";
import { renderWithProviders } from "@/../tests/test-utils";
import AgentCostView from "./agent_cost_view";
import type { Agent } from "@/components/agents/types";
const makeAgent = (litellmParams: Agent["litellm_params"]): Agent => ({
agent_id: "agent-1",
agent_name: "Test Agent",
litellm_params: litellmParams,
});
describe("AgentCostView", () => {
it("renders nothing when the agent has no cost configuration at all", () => {
const { container } = renderWithProviders(<AgentCostView agent={makeAgent({ model: "gpt-4" })} />);
expect(container).toBeEmptyDOMElement();
});
it("renders every configured cost with a dollar-prefixed value", () => {
const fullyPricedParams = {
model: "gpt-4",
cost_per_query: 0.05,
input_cost_per_token: 0.000012,
output_cost_per_token: 0.000034,
};
renderWithProviders(<AgentCostView agent={makeAgent(fullyPricedParams)} />);
expect(screen.getByText("Cost Configuration")).toBeInTheDocument();
expect(screen.getByText("Cost Per Query")).toBeInTheDocument();
expect(screen.getByText("$0.05")).toBeInTheDocument();
expect(screen.getByText("Input Cost Per Token")).toBeInTheDocument();
expect(screen.getByText("$0.000012")).toBeInTheDocument();
expect(screen.getByText("Output Cost Per Token")).toBeInTheDocument();
expect(screen.getByText("$0.000034")).toBeInTheDocument();
});
it("omits the rows whose cost is not configured", () => {
renderWithProviders(<AgentCostView agent={makeAgent({ model: "gpt-4", cost_per_query: 0.25 })} />);
expect(screen.getByText("Cost Per Query")).toBeInTheDocument();
expect(screen.getByText("$0.25")).toBeInTheDocument();
expect(screen.queryByText("Input Cost Per Token")).not.toBeInTheDocument();
expect(screen.queryByText("Output Cost Per Token")).not.toBeInTheDocument();
});
it("still renders a zero cost rather than treating it as unset", () => {
renderWithProviders(<AgentCostView agent={makeAgent({ model: "gpt-4", cost_per_query: 0 })} />);
expect(screen.getByText("Cost Configuration")).toBeInTheDocument();
expect(screen.getByText("Cost Per Query")).toBeInTheDocument();
expect(screen.getByText("$0")).toBeInTheDocument();
});
});

View file

@ -1,6 +1,4 @@
import React from "react";
import { Title } from "@tremor/react";
import { Descriptions } from "antd";
import { Agent } from "@/components/agents/types";
interface AgentCostViewProps {
@ -18,20 +16,25 @@ const AgentCostView: React.FC<AgentCostViewProps> = ({ agent }) => {
return null;
}
const rows = (
[
["Cost Per Query", params.cost_per_query],
["Input Cost Per Token", params.input_cost_per_token],
["Output Cost Per Token", params.output_cost_per_token],
] as const
).filter(([, value]) => value !== undefined);
return (
<div style={{ marginTop: 24 }}>
<Title>Cost Configuration</Title>
<Descriptions bordered column={1} style={{ marginTop: 16 }}>
{params.cost_per_query !== undefined && (
<Descriptions.Item label="Cost Per Query">${params.cost_per_query}</Descriptions.Item>
)}
{params.input_cost_per_token !== undefined && (
<Descriptions.Item label="Input Cost Per Token">${params.input_cost_per_token}</Descriptions.Item>
)}
{params.output_cost_per_token !== undefined && (
<Descriptions.Item label="Output Cost Per Token">${params.output_cost_per_token}</Descriptions.Item>
)}
</Descriptions>
<div className="mt-6">
<h3 className="text-lg font-semibold text-foreground">Cost Configuration</h3>
<dl className="mt-4 divide-y divide-border overflow-hidden rounded-lg border border-border">
{rows.map(([label, value]) => (
<div key={label} className="grid grid-cols-1 sm:grid-cols-3">
<dt className="bg-muted/50 px-4 py-3 text-sm font-medium text-foreground">{label}</dt>
<dd className="px-4 py-3 text-sm text-foreground sm:col-span-2">${value}</dd>
</div>
))}
</dl>
</div>
);
};

View file

@ -1,9 +1,8 @@
import React from "react";
import { Button, Tooltip, Typography } from "antd";
import { KeyOutlined } from "@ant-design/icons";
import { KeyRound } from "lucide-react";
import { KeyResponse } from "@/components/key_team_helpers/key_list";
const { Title, Text } = Typography;
import { Button } from "@/components/ui/button";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
interface AgentVirtualKeysProps {
keys: KeyResponse[];
@ -13,29 +12,31 @@ interface AgentVirtualKeysProps {
const AgentVirtualKeys: React.FC<AgentVirtualKeysProps> = ({ keys, isLoading, onKeyClick }) => {
return (
<div style={{ marginTop: 24 }}>
<Title level={4}>Virtual Keys</Title>
<div className="mt-6">
<h4 className="text-base font-semibold text-foreground">Virtual Keys</h4>
{isLoading ? (
<Text className="mt-2 block">Loading keys...</Text>
<p className="mt-2 text-sm text-muted-foreground">Loading keys...</p>
) : keys.length === 0 ? (
<Text className="mt-2 block text-gray-500">No virtual key assigned to this agent.</Text>
<p className="mt-2 text-sm text-muted-foreground">No virtual key assigned to this agent.</p>
) : (
<div className="mt-3 flex flex-col gap-2">
{keys.map((key) => (
<div key={key.token} className="flex items-center gap-3 border border-gray-100 rounded-sm px-3 py-2">
<KeyOutlined className="text-gray-400" />
<span className="font-medium">{key.key_alias || "Unnamed key"}</span>
{key.key_name && <span className="font-mono text-xs text-gray-500">{key.key_name}</span>}
<Tooltip title={key.token}>
<Button
size="small"
type="link"
className="font-mono text-blue-500 ml-auto"
onClick={() => onKeyClick(key)}
>
{key.token?.slice(0, 12)}...
</Button>
</Tooltip>
<div key={key.token} className="flex items-center gap-3 rounded-sm border border-border px-3 py-2">
<KeyRound className="size-4 text-muted-foreground" />
<span className="text-sm font-medium text-foreground">{key.key_alias || "Unnamed key"}</span>
{key.key_name && <span className="font-mono text-xs text-muted-foreground">{key.key_name}</span>}
<TooltipProvider delay={300}>
<Tooltip>
<TooltipTrigger
render={
<Button variant="link" size="sm" className="ml-auto font-mono" onClick={() => onKeyClick(key)}>
{key.token?.slice(0, 12)}...
</Button>
}
/>
<TooltipContent>{key.token}</TooltipContent>
</Tooltip>
</TooltipProvider>
</div>
))}
</div>

View file

@ -0,0 +1,141 @@
import { fireEvent, render } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import type { DailyData, KeyMetricWithMetadata, SpendMetrics } from "@/components/UsagePage/types";
vi.mock("@/components/shared/advanced_date_picker", () => ({
__esModule: true,
default: () => <div data-testid="date-picker" />,
}));
import CacheLeakageCard from "./CacheLeakageCard";
const baseMetrics = (overrides: Partial<SpendMetrics>): SpendMetrics => ({
spend: 0,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
api_requests: 0,
successful_requests: 0,
failed_requests: 0,
cache_read_input_tokens: 0,
cache_creation_input_tokens: 0,
...overrides,
});
const key = (alias: string, metrics: Partial<SpendMetrics>): KeyMetricWithMetadata => ({
metrics: baseMetrics(metrics),
metadata: { key_alias: alias, team_id: null },
});
const dayWithKeys = (date: string, apiKeys: Record<string, KeyMetricWithMetadata>): DailyData => ({
date,
metrics: baseMetrics({}),
breakdown: {
models: {},
model_groups: {},
mcp_servers: {},
providers: {},
api_keys: apiKeys,
entities: {},
},
});
const dayWithModels = (date: string, models: Record<string, Partial<SpendMetrics>>): DailyData => ({
date,
metrics: baseMetrics({}),
breakdown: {
models: Object.fromEntries(
Object.entries(models).map(([name, m]) => [
name,
{ metrics: baseMetrics(m), metadata: {}, api_key_breakdown: {} },
]),
),
model_groups: {},
mcp_servers: {},
providers: {},
api_keys: {},
entities: {},
},
});
const renderWith = (results: DailyData[]) =>
render(
<CacheLeakageCard
activity={{
dateValue: {},
onDateChange: vi.fn(),
results,
loading: false,
isFetchingMore: false,
}}
/>,
);
describe("CacheLeakageCard", () => {
it("ranks leaking keys by uncached prompt tokens and shows cache hit ratio", () => {
const { getByText, getByLabelText } = renderWith([
dayWithKeys("2026-07-12", {
"hash-caching": key("caching-key", { prompt_tokens: 1000, cache_read_input_tokens: 900 }),
"hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }),
}),
]);
expect(getByText("leaky-key")).toBeInTheDocument();
expect(getByText("0.0%")).toBeInTheDocument();
expect(getByText("90.0%")).toBeInTheDocument();
[
"Input tokens you sent in this range that weren't served from or written to the cache",
"Share of your input tokens that were served from the cache",
"About how much you'd save if this uncached input used prompt caching. Estimated as uncached input tokens times the per-token discount your cached traffic already gets (realized cache savings ÷ cache-read tokens).",
].forEach((info) => expect(getByLabelText(info)).toBeInTheDocument());
});
it("sorts by the clicked column, worst cache hit rate first", () => {
const { getAllByRole, getByText } = renderWith([
dayWithKeys("2026-07-12", {
"hash-a": key("alpha", {
prompt_tokens: 10000,
cache_read_input_tokens: 9000,
prompt_caching_savings_spend: 9.0,
}),
"hash-b": key("bravo", {
prompt_tokens: 500,
cache_read_input_tokens: 50,
prompt_caching_savings_spend: 0.05,
}),
}),
]);
const firstDataRow = () => getAllByRole("row")[1];
expect(firstDataRow()).toHaveTextContent("alpha");
fireEvent.click(getByText("Cache hit rate"));
expect(firstDataRow()).toHaveTextContent("bravo");
fireEvent.click(getByText("Cache hit rate"));
expect(firstDataRow()).toHaveTextContent("alpha");
});
it("switches to the model view and lists only Anthropic models", () => {
const { getByText, queryByText } = renderWith([
dayWithModels("2026-07-12", {
"claude-sonnet-5": { prompt_tokens: 5000, cache_read_input_tokens: 0 },
"gpt-4o": { prompt_tokens: 8000, cache_read_input_tokens: 0 },
}),
]);
fireEvent.click(getByText("By model"));
expect(getByText("Cache leakage by model")).toBeInTheDocument();
expect(getByText("claude-sonnet-5")).toBeInTheDocument();
expect(queryByText("gpt-4o")).not.toBeInTheDocument();
});
it("shows an empty state when no key used tokens in the range", () => {
const { getByText, queryByRole } = renderWith([dayWithKeys("2026-07-12", {})]);
expect(getByText("No key usage in this range.")).toBeInTheDocument();
expect(queryByRole("table")).not.toBeInTheDocument();
});
});

View file

@ -0,0 +1,183 @@
"use client";
import React, { useMemo, useState } from "react";
import { ArrowDown, ArrowUp, ArrowUpDown, Info } from "lucide-react";
import AdvancedDatePicker from "@/components/shared/advanced_date_picker";
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { CacheLeakageDimension, CacheLeakageRow, computeCacheLeakage, pct, usd } from "./costOptimizationUtils";
import { DailyActivityRange } from "./useDailyActivityRange";
interface CacheLeakageCardProps {
activity: DailyActivityRange;
}
type SortColumn = "uncachedPromptTokens" | "cacheHitRatio" | "potentialSavings";
interface SortState {
column: SortColumn;
dir: "asc" | "desc";
}
const NATURAL_DIR: Record<SortColumn, "asc" | "desc"> = {
uncachedPromptTokens: "desc",
cacheHitRatio: "asc",
potentialSavings: "desc",
};
const compareRows = (a: CacheLeakageRow, b: CacheLeakageRow, sort: SortState): number => {
const av = a[sort.column];
const bv = b[sort.column];
if (av == null && bv == null) return 0;
if (av == null) return 1;
if (bv == null) return -1;
return sort.dir === "asc" ? av - bv : bv - av;
};
const InfoTooltip = ({ info }: { info: string }) => (
<Tooltip>
<TooltipTrigger render={<span className="inline-flex" aria-label={info} />}>
<Info className="h-3 w-3 text-gray-400" />
</TooltipTrigger>
<TooltipContent className="max-w-xs">{info}</TooltipContent>
</Tooltip>
);
const SortableHead = ({
column,
label,
info,
sort,
onSort,
}: {
column: SortColumn;
label: string;
info: string;
sort: SortState;
onSort: (column: SortColumn) => void;
}) => {
const active = sort.column === column;
const ActiveArrow = sort.dir === "asc" ? ArrowUp : ArrowDown;
const Arrow = active ? ActiveArrow : ArrowUpDown;
return (
<TableHead className="text-right">
<span className="inline-flex items-center justify-end gap-1">
<button
type="button"
onClick={() => onSort(column)}
aria-label={`Sort by ${label}`}
className="inline-flex items-center gap-1 font-medium hover:text-foreground"
>
{label}
<Arrow className={`h-3 w-3 ${active ? "text-foreground" : "text-gray-400"}`} />
</button>
<InfoTooltip info={info} />
</span>
</TableHead>
);
};
const CacheLeakageCard: React.FC<CacheLeakageCardProps> = ({ activity }) => {
const { dateValue, onDateChange, results, loading, isFetchingMore } = activity;
const [dimension, setDimension] = useState<CacheLeakageDimension>("key");
const [sort, setSort] = useState<SortState>({ column: "potentialSavings", dir: "desc" });
const leakage = useMemo(() => computeCacheLeakage(results, dimension), [results, dimension]);
const rows = useMemo(() => [...leakage.rows].sort((a, b) => compareRows(a, b, sort)), [leakage.rows, sort]);
const onSort = (column: SortColumn) =>
setSort((prev) =>
prev.column === column
? { column, dir: prev.dir === "asc" ? "desc" : "asc" }
: { column, dir: NATURAL_DIR[column] },
);
const subject = dimension === "model" ? "Models" : "Keys";
const firstColumn = dimension === "model" ? "Model" : "Key";
const emptyNoun = dimension === "model" ? "model" : "key";
return (
<TooltipProvider delay={300}>
<Card>
<CardHeader>
<div className="flex flex-wrap items-start justify-between gap-4">
<div>
<CardTitle>Cache leakage by {dimension === "model" ? "model" : "virtual key"}</CardTitle>
<p className="mt-1 text-sm text-muted-foreground">
{subject} sending large volumes of uncached input with a low cache hit rate are likely missing prompt
caching. Potential savings is approximate: uncached input priced at the realized cache-read discount.
{dimension === "model" ? " Limited to Anthropic (Claude) models, which support prompt caching." : ""}
</p>
</div>
<AdvancedDatePicker value={dateValue} onValueChange={onDateChange} />
</div>
<Tabs
value={dimension}
onValueChange={(value) => setDimension(value === "model" ? "model" : "key")}
className="mt-3"
>
<TabsList>
<TabsTrigger value="key">By virtual key</TabsTrigger>
<TabsTrigger value="model">By model</TabsTrigger>
</TabsList>
</Tabs>
</CardHeader>
<CardContent>
{rows.length === 0 ? (
<p className="py-8 text-center text-sm text-muted-foreground">
{loading || isFetchingMore ? "Loading..." : `No ${emptyNoun} usage in this range.`}
</p>
) : (
<Table>
<TableHeader>
<TableRow>
<TableHead>{firstColumn}</TableHead>
<SortableHead
column="uncachedPromptTokens"
label="Uncached input tokens"
info="Input tokens you sent in this range that weren't served from or written to the cache"
sort={sort}
onSort={onSort}
/>
<SortableHead
column="cacheHitRatio"
label="Cache hit rate"
info="Share of your input tokens that were served from the cache"
sort={sort}
onSort={onSort}
/>
<SortableHead
column="potentialSavings"
label="Potential savings"
info="About how much you'd save if this uncached input used prompt caching. Estimated as uncached input tokens times the per-token discount your cached traffic already gets (realized cache savings ÷ cache-read tokens)."
sort={sort}
onSort={onSort}
/>
</TableRow>
</TableHeader>
<TableBody>
{rows.map((row) => (
<TableRow key={row.id}>
<TableCell className="font-medium">
{row.label}
{row.sublabel && <span className="ml-1 text-xs text-muted-foreground">({row.sublabel})</span>}
</TableCell>
<TableCell className="text-right">{formatNumberWithCommas(row.uncachedPromptTokens)}</TableCell>
<TableCell className="text-right">{pct(row.cacheHitRatio)}</TableCell>
<TableCell className="text-right">
{row.potentialSavings == null ? "—" : usd(row.potentialSavings)}
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
)}
</CardContent>
</Card>
</TooltipProvider>
);
};
export default CacheLeakageCard;

View file

@ -0,0 +1,53 @@
import { fireEvent, render, waitFor } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
const mockUserDailyActivityCall = vi.fn();
vi.mock("@/components/networking", () => ({
userDailyActivityCall: (...args: unknown[]) => mockUserDailyActivityCall(...args),
getToolSpend: vi.fn().mockResolvedValue({ by_tool: [], daily: [], total_spend: 0, start_date: null, end_date: null }),
getGeneralSettingsCall: vi.fn().mockResolvedValue([]),
}));
vi.mock("@/components/shared/advanced_date_picker", () => ({
__esModule: true,
default: () => <div data-testid="date-picker" />,
}));
vi.mock("@/components/shared/charts", () => ({
AreaChart: () => <div />,
DonutChart: () => <div />,
BarChart: () => <div />,
DEFAULT_COLOR_CYCLE: ["emerald"],
}));
vi.mock("@/app/(dashboard)/router-settings/_components/general_settings", () => ({
PromptCachingPanel: () => <div data-testid="caching-settings" />,
}));
vi.mock("./PromptCompressionTab", () => ({ __esModule: true, default: () => <div /> }));
vi.mock("./AutorouterTab", () => ({ __esModule: true, default: () => <div /> }));
import CostOptimizationView from "./CostOptimizationView";
const singlePage = {
results: [],
metadata: { total_pages: 1, has_more: false, page: 1 },
};
describe("CostOptimizationView daily activity", () => {
it("fetches daily activity once for the page and shares it with every tab that needs it", async () => {
mockUserDailyActivityCall.mockResolvedValue(singlePage);
const { getByRole, getByTestId } = render(
<CostOptimizationView accessToken="test-token" userId="u1" userRole="proxy_admin" />,
);
await waitFor(() => expect(mockUserDailyActivityCall).toHaveBeenCalledTimes(1));
fireEvent.click(getByRole("tab", { name: "Prompt Caching" }));
await waitFor(() => expect(getByTestId("caching-settings")).toBeInTheDocument());
expect(mockUserDailyActivityCall).toHaveBeenCalledTimes(1);
});
});

View file

@ -8,6 +8,7 @@ import UsageTab from "./UsageTab";
import PromptCompressionTab from "./PromptCompressionTab";
import AutorouterTab from "./AutorouterTab";
import PromptCachingTab from "./PromptCachingTab";
import { useDailyActivityRange } from "./useDailyActivityRange";
interface CostOptimizationViewProps {
accessToken: string | null;
@ -16,11 +17,13 @@ interface CostOptimizationViewProps {
}
const CostOptimizationView: React.FC<CostOptimizationViewProps> = ({ accessToken, userId, userRole }) => {
const activity = useDailyActivityRange(accessToken, userId, userRole);
const items = [
{
key: "usage",
label: "Usage",
children: <UsageTab accessToken={accessToken} userId={userId} userRole={userRole} />,
children: <UsageTab accessToken={accessToken} activity={activity} />,
},
{
key: "compression",
@ -35,7 +38,7 @@ const CostOptimizationView: React.FC<CostOptimizationViewProps> = ({ accessToken
{
key: "caching",
label: "Prompt Caching",
children: <PromptCachingTab accessToken={accessToken} />,
children: <PromptCachingTab accessToken={accessToken} activity={activity} />,
},
];

View file

@ -0,0 +1,43 @@
import { render, waitFor } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
const mockGetGeneralSettingsCall = vi.fn();
vi.mock("@/components/networking", () => ({
getGeneralSettingsCall: (...args: unknown[]) => mockGetGeneralSettingsCall(...args),
}));
vi.mock("@/app/(dashboard)/router-settings/_components/general_settings", () => ({
PromptCachingPanel: () => <div data-testid="caching-settings" />,
}));
const mockCacheLeakageCard = vi.fn();
vi.mock("./CacheLeakageCard", () => ({
__esModule: true,
default: (props: unknown) => {
mockCacheLeakageCard(props);
return <div data-testid="cache-leakage-card" />;
},
}));
import PromptCachingTab from "./PromptCachingTab";
describe("PromptCachingTab", () => {
it("renders the cache leakage table alongside the caching settings", async () => {
mockGetGeneralSettingsCall.mockResolvedValue([]);
const activity = {
dateValue: {},
onDateChange: vi.fn(),
results: [],
loading: false,
isFetchingMore: false,
};
const { getByTestId } = render(<PromptCachingTab accessToken="test-token" activity={activity} />);
expect(getByTestId("caching-settings")).toBeInTheDocument();
expect(getByTestId("cache-leakage-card")).toBeInTheDocument();
await waitFor(() => expect(mockCacheLeakageCard).toHaveBeenCalledWith(expect.objectContaining({ activity })));
});
});

View file

@ -8,12 +8,15 @@ import {
PromptCachingPanel,
generalSettingsItem,
} from "@/app/(dashboard)/router-settings/_components/general_settings";
import CacheLeakageCard from "./CacheLeakageCard";
import { DailyActivityRange } from "./useDailyActivityRange";
interface PromptCachingTabProps {
accessToken: string | null;
activity: DailyActivityRange;
}
const PromptCachingTab: React.FC<PromptCachingTabProps> = ({ accessToken }) => {
const PromptCachingTab: React.FC<PromptCachingTabProps> = ({ accessToken, activity }) => {
const [settings, setSettings] = useState<generalSettingsItem[]>([]);
const loadSettings = useCallback(() => {
@ -43,8 +46,9 @@ const PromptCachingTab: React.FC<PromptCachingTabProps> = ({ accessToken }) => {
}
return (
<div className="w-full">
<div className="w-full space-y-6">
<PromptCachingPanel accessToken={accessToken} settings={settings} onChange={handleChange} />
<CacheLeakageCard activity={activity} />
</div>
);
};

View file

@ -1,16 +1,13 @@
import { render } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import type { ToolSpendResponse } from "@/components/networking";
import type { DailyData, SpendMetrics } from "@/components/UsagePage/types";
const mockUsePaginatedDailyActivity = vi.fn();
vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({
usePaginatedDailyActivity: (args: unknown) => mockUsePaginatedDailyActivity(args),
}));
const mockGetToolSpend = vi.fn();
vi.mock("@/components/networking", () => ({
userDailyActivityCall: vi.fn(),
getToolSpend: (...args: unknown[]) => mockGetToolSpend(...args),
}));
vi.mock("@/components/shared/advanced_date_picker", () => ({
@ -25,10 +22,16 @@ vi.mock("@/components/shared/charts", () => ({
DonutChart: ({ data, label }: { data: unknown; label: string }) => (
<div data-testid="donut-chart" data-label={label} data-slices={JSON.stringify(data)} />
),
BarChart: ({ data, categories }: { data: unknown; categories: string[] }) => (
<div data-testid="bar-chart" data-categories={categories.join(",")} data-series={JSON.stringify(data)} />
),
DEFAULT_COLOR_CYCLE: ["emerald", "blue", "violet", "amber"],
}));
import UsageTab from "./UsageTab";
const emptyToolSpend: ToolSpendResponse = { by_tool: [], daily: [], total_spend: 0, start_date: null, end_date: null };
const baseMetrics = (overrides: Partial<SpendMetrics>): SpendMetrics => ({
spend: 0,
prompt_tokens: 0,
@ -55,9 +58,20 @@ const day = (date: string, metrics: Partial<SpendMetrics>): DailyData => ({
},
});
const renderWith = (results: DailyData[]) => {
mockUsePaginatedDailyActivity.mockReturnValue({ data: { results }, loading: false, isFetchingMore: false });
return render(<UsageTab accessToken="test-token" userId="u1" userRole="proxy_admin" />);
const renderWith = (results: DailyData[], toolSpend = emptyToolSpend) => {
mockGetToolSpend.mockResolvedValue(toolSpend);
return render(
<UsageTab
accessToken="test-token"
activity={{
dateValue: { from: new Date("2026-07-01"), to: new Date("2026-07-14") },
onDateChange: vi.fn(),
results,
loading: false,
isFetchingMore: false,
}}
/>,
);
};
describe("UsageTab", () => {
@ -105,4 +119,22 @@ describe("UsageTab", () => {
const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
expect(slices).toEqual([{ driver: "Compression", usd: expect.closeTo(0.04, 5) }]);
});
it("renders spend-by-tool bars from the tool spend endpoint", async () => {
const toolSpend = {
by_tool: [
{ tool_name: "search", spend: 4.0, call_count: 3, total_tokens: 150 },
{ tool_name: "read_file", spend: 1.0, call_count: 2, total_tokens: 50 },
],
daily: [{ date: "2026-07-12", tool_name: "search", spend: 4.0, call_count: 3 }],
total_spend: 5.0,
start_date: "2026-07-12",
end_date: "2026-07-12",
};
const { findAllByTestId } = renderWith([day("2026-07-12", {})], toolSpend);
const bars = await findAllByTestId("bar-chart");
const series = JSON.parse(bars[0].getAttribute("data-series") ?? "[]");
expect(series[0]).toMatchObject({ tool_name: "search", spend: 4.0 });
});
});

View file

@ -1,35 +1,35 @@
"use client";
import React, { useMemo, useState } from "react";
import React, { useEffect, useMemo, useState } from "react";
import { Collapse } from "antd";
import { AreaChart, DonutChart } from "@/components/shared/charts";
import { AreaChart, BarChart, DonutChart, DEFAULT_COLOR_CYCLE } from "@/components/shared/charts";
import AdvancedDatePicker from "@/components/shared/advanced_date_picker";
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { userDailyActivityCall } from "@/components/networking";
import { DailyData, SpendMetrics } from "@/components/UsagePage/types";
import { getToolSpend, ToolSpendResponse } from "@/components/networking";
import { SpendMetrics } from "@/components/UsagePage/types";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { all_admin_roles } from "@/utils/roles";
import { usePaginatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity";
import { buildDailyToolSeries, topToolsBySpend, usd } from "./costOptimizationUtils";
import { DailyActivityRange } from "./useDailyActivityRange";
interface UsageTabProps {
accessToken: string | null;
userId: string | null;
userRole: string;
activity: DailyActivityRange;
}
type DateRange = { from?: Date; to?: Date };
const THIRTY_DAYS_MS = 30 * 24 * 60 * 60 * 1000;
const usd = (value: number): string => {
const decimals = value > 0 && value < 1 ? 4 : 2;
return `$${formatNumberWithCommas(value, decimals)}`;
const EMPTY_TOOL_SPEND: ToolSpendResponse = {
by_tool: [],
daily: [],
total_spend: 0,
start_date: null,
end_date: null,
};
const shortDate = (iso: string): string =>
new Date(`${iso}T00:00:00`).toLocaleDateString("en-US", { month: "short", day: "numeric" });
const isoDay = (d: Date): string => d.toISOString().slice(0, 10);
const compressionOf = (m: SpendMetrics): number => m.compression_savings_spend ?? 0;
const cachingOf = (m: SpendMetrics): number => m.prompt_caching_savings_spend ?? 0;
const savedTokensOf = (m: SpendMetrics): number => m.compression_saved_tokens ?? 0;
@ -81,23 +81,33 @@ const SummaryCard = ({ label, value, hint }: { label: string; value: string; hin
</Card>
);
const UsageTab: React.FC<UsageTabProps> = ({ accessToken, userId, userRole }) => {
const initialFrom = useMemo(() => new Date(new Date().getTime() - THIRTY_DAYS_MS), []);
const initialTo = useMemo(() => new Date(), []);
const [dateValue, setDateValue] = useState<DateRange>({ from: initialFrom, to: initialTo });
const UsageTab: React.FC<UsageTabProps> = ({ accessToken, activity }) => {
const { dateValue, onDateChange, results, loading, isFetchingMore } = activity;
const startTime = dateValue.from ?? null;
const endTime = dateValue.to ?? null;
const isAdmin = all_admin_roles.includes(userRole);
const effectiveUserId = isAdmin ? null : userId;
const { data, loading, isFetchingMore } = usePaginatedDailyActivity({
fetchFn: userDailyActivityCall,
args: [accessToken, startTime, endTime, effectiveUserId],
enabled: !!accessToken && !!startTime && !!endTime,
});
const toolSpendEnabled = !!accessToken && !!startTime && !!endTime;
const rangeKey = startTime && endTime ? `${isoDay(startTime)}|${isoDay(endTime)}` : "";
const [toolSpendState, setToolSpendState] = useState<{ key: string; data: ToolSpendResponse } | null>(null);
const results = data.results as DailyData[];
useEffect(() => {
if (!accessToken || !startTime || !endTime) return;
let cancelled = false;
getToolSpend(accessToken, isoDay(startTime), isoDay(endTime))
.then((res) => {
if (!cancelled) setToolSpendState({ key: rangeKey, data: res });
})
.catch(() => {
if (!cancelled) setToolSpendState({ key: rangeKey, data: EMPTY_TOOL_SPEND });
});
return () => {
cancelled = true;
};
}, [accessToken, startTime, endTime, rangeKey]);
const toolSpend = toolSpendState?.key === rangeKey ? toolSpendState.data : null;
const toolSpendLoading = toolSpendEnabled && toolSpend === null;
const compressionTotal = useMemo(() => results.reduce((sum, d) => sum + compressionOf(d.metrics), 0), [results]);
const cachingTotal = useMemo(() => results.reduce((sum, d) => sum + cachingOf(d.metrics), 0), [results]);
@ -123,11 +133,27 @@ const UsageTab: React.FC<UsageTabProps> = ({ accessToken, userId, userRole }) =>
[compressionTotal, cachingTotal],
);
const topTools = useMemo(() => topToolsBySpend(toolSpend?.by_tool ?? []), [toolSpend]);
const topToolNames = useMemo(() => topTools.map((t) => t.tool_name), [topTools]);
const topToolsChart = useMemo<Record<string, string | number>[]>(
() => topTools.map((t) => ({ tool_name: t.tool_name, spend: t.spend })),
[topTools],
);
const dailyToolSeries = useMemo(
() =>
buildDailyToolSeries(toolSpend?.daily ?? [], topToolNames).map((point) => ({
...point,
date: shortDate(String(point.date)),
})),
[toolSpend, topToolNames],
);
const toolColors = useMemo(() => DEFAULT_COLOR_CYCLE.slice(0, Math.max(topToolNames.length, 1)), [topToolNames]);
return (
<div className="w-full space-y-6">
<div className="flex flex-wrap items-center justify-between gap-4">
<MethodologyNote />
<AdvancedDatePicker value={dateValue} onValueChange={(v) => setDateValue(v)} />
<AdvancedDatePicker value={dateValue} onValueChange={onDateChange} />
</div>
<div className="grid grid-cols-1 gap-6 sm:grid-cols-2 lg:grid-cols-3">
@ -177,6 +203,50 @@ const UsageTab: React.FC<UsageTabProps> = ({ accessToken, userId, userRole }) =>
</CardContent>
</Card>
</div>
<Card>
<CardHeader>
<CardTitle>Spend by tool</CardTitle>
<p className="text-sm text-muted-foreground">
Spend on requests that called each tool (MCP and client-side tools). A request that used multiple tools
counts its full spend toward each, so this attributes rather than partitions spend.
</p>
</CardHeader>
<CardContent>
{topTools.length === 0 ? (
<p className="py-8 text-center text-sm text-muted-foreground">
{toolSpendLoading ? "Loading..." : "No tool usage in this range."}
</p>
) : (
<div className="grid grid-cols-1 gap-6 lg:grid-cols-2">
<div>
<p className="mb-2 text-sm font-medium text-muted-foreground">Total by tool</p>
<BarChart
data={topToolsChart}
index="tool_name"
categories={["spend"]}
colors={["emerald"]}
layout="vertical"
yAxisWidth={140}
showLegend={false}
valueFormatter={usd}
/>
</div>
<div>
<p className="mb-2 text-sm font-medium text-muted-foreground">Daily spend by tool</p>
<BarChart
data={dailyToolSeries}
index="date"
categories={topToolNames}
colors={toolColors}
stack
valueFormatter={usd}
/>
</div>
</div>
)}
</CardContent>
</Card>
</div>
);
};

View file

@ -0,0 +1,223 @@
import { describe, expect, it } from "vitest";
import type { DailyData, SpendMetrics } from "@/components/UsagePage/types";
import type { ToolSpendDailyEntry, ToolSpendEntry } from "@/components/networking";
import { buildDailyToolSeries, computeCacheLeakage, isAnthropicModel, topToolsBySpend } from "./costOptimizationUtils";
const metrics = (overrides: Partial<SpendMetrics>): SpendMetrics => ({
spend: 0,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
api_requests: 0,
successful_requests: 0,
failed_requests: 0,
cache_read_input_tokens: 0,
cache_creation_input_tokens: 0,
...overrides,
});
const day = (
date: string,
keys: Record<string, { alias: string | null; metrics: Partial<SpendMetrics> }>,
): DailyData => ({
date,
metrics: metrics({}),
breakdown: {
models: {},
model_groups: {},
mcp_servers: {},
providers: {},
entities: {},
api_keys: Object.fromEntries(
Object.entries(keys).map(([hash, v]) => [
hash,
{ metrics: metrics(v.metrics), metadata: { key_alias: v.alias, team_id: null } },
]),
),
},
});
const modelDay = (date: string, models: Record<string, Partial<SpendMetrics>>): DailyData => ({
date,
metrics: metrics({}),
breakdown: {
models: Object.fromEntries(
Object.entries(models).map(([name, m]) => [name, { metrics: metrics(m), metadata: {}, api_key_breakdown: {} }]),
),
model_groups: {},
mcp_servers: {},
providers: {},
entities: {},
api_keys: {},
},
});
describe("computeCacheLeakage", () => {
it("aggregates a key's tokens and savings across multiple days", () => {
const results = [
day("2026-07-01", { h1: { alias: "svc-a", metrics: { prompt_tokens: 1000, cache_read_input_tokens: 0 } } }),
day("2026-07-02", { h1: { alias: "svc-a", metrics: { prompt_tokens: 500, cache_read_input_tokens: 0 } } }),
];
const { rows } = computeCacheLeakage(results);
expect(rows).toHaveLength(1);
expect(rows[0].uncachedPromptTokens).toBe(1500);
});
it("subtracts cache reads and writes from prompt tokens instead of double-counting them", () => {
const results = [
day("2026-07-01", {
h1: {
alias: "svc-a",
metrics: { prompt_tokens: 1000, cache_read_input_tokens: 400, cache_creation_input_tokens: 100 },
},
}),
];
const { rows } = computeCacheLeakage(results);
expect(rows).toHaveLength(1);
expect(rows[0].uncachedPromptTokens).toBe(500);
expect(rows[0].cacheHitRatio).toBeCloseTo(0.4, 6);
});
it("prices leakage at the portfolio's realized cache-read discount and drops fully cached keys", () => {
const results = [
day("2026-07-01", {
cacher: {
alias: "cacher",
metrics: { prompt_tokens: 1000, cache_read_input_tokens: 1000, prompt_caching_savings_spend: 2.0 },
},
leaker: { alias: "leaker", metrics: { prompt_tokens: 500 } },
}),
];
const { rows, discountPerToken } = computeCacheLeakage(results);
expect(discountPerToken).toBeCloseTo(0.002, 6);
expect(rows.map((r) => r.label)).toEqual(["leaker"]);
expect(rows[0].potentialSavings).toBeCloseTo(1.0, 6);
});
it("returns null estimate and ranks by uncached tokens when nobody used caching", () => {
const results = [
day("2026-07-01", {
big: { alias: "big", metrics: { prompt_tokens: 9000 } },
small: { alias: "small", metrics: { prompt_tokens: 100 } },
}),
];
const { rows, discountPerToken } = computeCacheLeakage(results);
expect(discountPerToken).toBeNull();
expect(rows.map((r) => r.label)).toEqual(["big", "small"]);
expect(rows.every((r) => r.potentialSavings === null)).toBe(true);
});
it("computes cache hit ratio against total prompt tokens and clamps inconsistent data at zero", () => {
const results = [
day("2026-07-01", {
onlycache: { alias: "onlycache", metrics: { cache_read_input_tokens: 100 } },
mixed: { alias: "mixed", metrics: { prompt_tokens: 1000, cache_read_input_tokens: 750 } },
}),
];
const { rows } = computeCacheLeakage(results);
expect(rows.map((r) => r.label)).toEqual(["mixed"]);
expect(rows[0].cacheHitRatio).toBeCloseTo(0.75, 6);
expect(rows[0].uncachedPromptTokens).toBe(250);
});
it("respects the row limit", () => {
const keys = Object.fromEntries(
Array.from({ length: 15 }, (_, i) => [`h${i}`, { alias: `k${i}`, metrics: { prompt_tokens: i + 1 } }]),
);
const { rows } = computeCacheLeakage([day("2026-07-01", keys)], "key", 5);
expect(rows).toHaveLength(5);
});
});
describe("computeCacheLeakage by model", () => {
it("aggregates only Anthropic models and ignores other providers", () => {
const models: Record<string, Partial<SpendMetrics>> = {
"claude-sonnet-5": { prompt_tokens: 10000, cache_read_input_tokens: 0 },
"anthropic/claude-haiku-4-5": { prompt_tokens: 4000, cache_read_input_tokens: 0 },
"bedrock/anthropic.claude-3-5-sonnet": { prompt_tokens: 2000, cache_read_input_tokens: 0 },
"gpt-4o": { prompt_tokens: 9000, cache_read_input_tokens: 0 },
"deepseek-chat": { prompt_tokens: 8000, cache_read_input_tokens: 0 },
};
const { rows } = computeCacheLeakage([modelDay("2026-07-01", models)], "model");
expect(rows.map((r) => r.id)).toEqual([
"claude-sonnet-5",
"anthropic/claude-haiku-4-5",
"bedrock/anthropic.claude-3-5-sonnet",
]);
});
it("labels model rows by model name with no sublabel", () => {
const results = [modelDay("2026-07-01", { "claude-sonnet-5": { prompt_tokens: 1000 } })];
const { rows } = computeCacheLeakage(results, "model");
expect(rows[0].label).toBe("claude-sonnet-5");
expect(rows[0].sublabel).toBeNull();
});
it("prices model leakage at the Anthropic realized cache-read discount", () => {
const results = [
modelDay("2026-07-01", {
"claude-sonnet-5": { prompt_tokens: 1000, cache_read_input_tokens: 1000, prompt_caching_savings_spend: 2.0 },
"claude-haiku-4-5": { prompt_tokens: 500 },
}),
];
const { rows, discountPerToken } = computeCacheLeakage(results, "model");
expect(discountPerToken).toBeCloseTo(0.002, 6);
expect(rows.map((r) => r.id)).toEqual(["claude-haiku-4-5"]);
expect(rows[0].potentialSavings).toBeCloseTo(1.0, 6);
});
});
describe("isAnthropicModel", () => {
it("matches Claude-family models across providers and rejects others", () => {
const anthropic = [
"claude-sonnet-5",
"anthropic/claude-haiku-4-5",
"bedrock/anthropic.claude-3-5-sonnet",
"vertex_ai/claude-opus-4-8",
];
const others = ["gpt-4o", "deepseek-chat", "gemini-2.5-pro", "mistral-large"];
expect(anthropic.every(isAnthropicModel)).toBe(true);
expect(others.some(isAnthropicModel)).toBe(false);
});
});
describe("buildDailyToolSeries", () => {
const daily: ToolSpendDailyEntry[] = [
{ date: "2026-07-01", tool_name: "search", spend: 1.0, call_count: 1 },
{ date: "2026-07-01", tool_name: "read", spend: 0.5, call_count: 1 },
{ date: "2026-07-02", tool_name: "search", spend: 2.0, call_count: 1 },
{ date: "2026-07-01", tool_name: "excluded", spend: 9.0, call_count: 1 },
];
it("pivots to per-date points keyed by the selected tools, dropping others", () => {
const series = buildDailyToolSeries(daily, ["search", "read"]);
expect(series).toEqual([
{ date: "2026-07-01", search: 1.0, read: 0.5 },
{ date: "2026-07-02", search: 2.0, read: 0 },
]);
});
it("sums repeated (date, tool) rows", () => {
const series = buildDailyToolSeries(
[
{ date: "2026-07-01", tool_name: "search", spend: 1.0, call_count: 1 },
{ date: "2026-07-01", tool_name: "search", spend: 2.5, call_count: 1 },
],
["search"],
);
expect(series[0].search).toBe(3.5);
});
});
describe("topToolsBySpend", () => {
const byTool: ToolSpendEntry[] = [
{ tool_name: "a", spend: 1, call_count: 1, total_tokens: 1 },
{ tool_name: "b", spend: 5, call_count: 1, total_tokens: 1 },
{ tool_name: "c", spend: 3, call_count: 1, total_tokens: 1 },
];
it("sorts by spend descending and truncates to the limit", () => {
expect(topToolsBySpend(byTool, 2).map((t) => t.tool_name)).toEqual(["b", "c"]);
});
});

View file

@ -0,0 +1,151 @@
import { DailyData, SpendMetrics } from "@/components/UsagePage/types";
import { ToolSpendDailyEntry, ToolSpendEntry } from "@/components/networking";
import { formatNumberWithCommas } from "@/utils/dataUtils";
export const usd = (value: number): string => {
const decimals = value > 0 && value < 1 ? 4 : 2;
return `$${formatNumberWithCommas(value, decimals)}`;
};
export const pct = (ratio: number): string => `${formatNumberWithCommas(ratio * 100, 1)}%`;
export type CacheLeakageDimension = "key" | "model";
export interface CacheLeakageRow {
id: string;
label: string;
sublabel: string | null;
uncachedPromptTokens: number;
cacheHitRatio: number;
potentialSavings: number | null;
}
export interface CacheLeakageResult {
rows: CacheLeakageRow[];
discountPerToken: number | null;
}
export const isAnthropicModel = (model: string): boolean => /claude|anthropic/i.test(model);
interface LeakageAccumulator {
alias: string | null;
teamId: string | null;
promptTokens: number;
cacheReadTokens: number;
cacheCreationTokens: number;
realizedCachingSavings: number;
}
const emptyAccumulator = (): LeakageAccumulator => ({
alias: null,
teamId: null,
promptTokens: 0,
cacheReadTokens: 0,
cacheCreationTokens: 0,
realizedCachingSavings: 0,
});
const addMetrics = (
acc: LeakageAccumulator,
m: SpendMetrics,
alias: string | null,
teamId: string | null,
): LeakageAccumulator => ({
alias: acc.alias ?? alias,
teamId: acc.teamId ?? teamId,
promptTokens: acc.promptTokens + (m.prompt_tokens ?? 0),
cacheReadTokens: acc.cacheReadTokens + (m.cache_read_input_tokens ?? 0),
cacheCreationTokens: acc.cacheCreationTokens + (m.cache_creation_input_tokens ?? 0),
realizedCachingSavings: acc.realizedCachingSavings + (m.prompt_caching_savings_spend ?? 0),
});
const aggregateByKey = (results: readonly DailyData[]): Map<string, LeakageAccumulator> => {
const byKey = new Map<string, LeakageAccumulator>();
for (const day of results) {
for (const [apiKey, entry] of Object.entries(day.breakdown?.api_keys ?? {})) {
const acc = byKey.get(apiKey) ?? emptyAccumulator();
byKey.set(
apiKey,
addMetrics(acc, entry.metrics, entry.metadata?.key_alias ?? null, entry.metadata?.team_id ?? null),
);
}
}
return byKey;
};
const aggregateByModel = (results: readonly DailyData[]): Map<string, LeakageAccumulator> => {
const byModel = new Map<string, LeakageAccumulator>();
for (const day of results) {
for (const [model, entry] of Object.entries(day.breakdown?.models ?? {})) {
if (!isAnthropicModel(model)) continue;
const acc = byModel.get(model) ?? emptyAccumulator();
byModel.set(model, addMetrics(acc, entry.metrics, null, null));
}
}
return byModel;
};
export const computeCacheLeakage = (
results: readonly DailyData[],
dimension: CacheLeakageDimension = "key",
limit = 10,
): CacheLeakageResult => {
const byEntity = dimension === "model" ? aggregateByModel(results) : aggregateByKey(results);
const totals = [...byEntity.values()].reduce(
(agg, a) => ({
cacheReadTokens: agg.cacheReadTokens + a.cacheReadTokens,
realizedCachingSavings: agg.realizedCachingSavings + a.realizedCachingSavings,
}),
{ cacheReadTokens: 0, realizedCachingSavings: 0 },
);
const discountPerToken = totals.cacheReadTokens > 0 ? totals.realizedCachingSavings / totals.cacheReadTokens : null;
const rows: CacheLeakageRow[] = [...byEntity.entries()]
.map(([id, a]) => {
const uncachedPromptTokens = Math.max(0, a.promptTokens - a.cacheReadTokens - a.cacheCreationTokens);
return {
id,
label: dimension === "model" ? id : a.alias ?? `${id.slice(0, 8)}...`,
sublabel: dimension === "model" ? null : a.teamId,
uncachedPromptTokens,
cacheHitRatio: a.promptTokens > 0 ? a.cacheReadTokens / a.promptTokens : 0,
potentialSavings: discountPerToken != null ? uncachedPromptTokens * discountPerToken : null,
};
})
.filter((row) => row.uncachedPromptTokens > 0);
const sorted = rows.sort((x, y) =>
discountPerToken != null
? (y.potentialSavings ?? 0) - (x.potentialSavings ?? 0)
: y.uncachedPromptTokens - x.uncachedPromptTokens,
);
return { rows: sorted.slice(0, limit), discountPerToken };
};
export interface DailyToolSpendPoint {
date: string;
[toolName: string]: string | number;
}
export const buildDailyToolSeries = (
daily: readonly ToolSpendDailyEntry[],
topToolNames: readonly string[],
): DailyToolSpendPoint[] => {
const top = new Set(topToolNames);
const byDate = new Map<string, DailyToolSpendPoint>();
for (const d of daily) {
if (!top.has(d.tool_name)) continue;
const point = byDate.get(d.date) ?? seedPoint(d.date, topToolNames);
point[d.tool_name] = (Number(point[d.tool_name]) || 0) + d.spend;
byDate.set(d.date, point);
}
return [...byDate.values()].sort((a, b) => a.date.localeCompare(b.date));
};
const seedPoint = (date: string, toolNames: readonly string[]): DailyToolSpendPoint =>
toolNames.reduce<DailyToolSpendPoint>((p, name) => ({ ...p, [name]: 0 }), { date });
export const topToolsBySpend = (byTool: readonly ToolSpendEntry[], limit = 8): ToolSpendEntry[] =>
[...byTool].sort((a, b) => b.spend - a.spend).slice(0, limit);

View file

@ -0,0 +1,39 @@
import { renderHook } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
const mockUsePaginatedDailyActivity = vi.fn();
vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({
usePaginatedDailyActivity: (args: unknown) => {
mockUsePaginatedDailyActivity(args);
return { data: { results: [] }, loading: false, isFetchingMore: false };
},
}));
vi.mock("@/components/networking", () => ({
userDailyActivityCall: vi.fn(),
}));
import { useDailyActivityRange } from "./useDailyActivityRange";
const argsOfLastCall = () => mockUsePaginatedDailyActivity.mock.calls.at(-1)?.[0].args as unknown[];
describe("useDailyActivityRange", () => {
it("queries every user's activity for an admin", () => {
renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin"));
expect(argsOfLastCall()).toEqual(["test-token", expect.any(Date), expect.any(Date), null]);
});
it("scopes the query to the caller for a non-admin", () => {
renderHook(() => useDailyActivityRange("test-token", "u1", "internal_user"));
expect(argsOfLastCall()).toEqual(["test-token", expect.any(Date), expect.any(Date), "u1"]);
});
it("stays disabled until an access token is available", () => {
renderHook(() => useDailyActivityRange(null, "u1", "proxy_admin"));
expect(mockUsePaginatedDailyActivity).toHaveBeenLastCalledWith(expect.objectContaining({ enabled: false }));
});
});

View file

@ -0,0 +1,49 @@
import { useMemo, useState } from "react";
import { userDailyActivityCall } from "@/components/networking";
import { DailyData } from "@/components/UsagePage/types";
import { all_admin_roles } from "@/utils/roles";
import { usePaginatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity";
const THIRTY_DAYS_MS = 30 * 24 * 60 * 60 * 1000;
export interface DateRange {
from?: Date;
to?: Date;
}
export interface DailyActivityRange {
dateValue: DateRange;
onDateChange: (value: DateRange) => void;
results: DailyData[];
loading: boolean;
isFetchingMore: boolean;
}
export const useDailyActivityRange = (
accessToken: string | null,
userId: string | null,
userRole: string,
): DailyActivityRange => {
const initialFrom = useMemo(() => new Date(new Date().getTime() - THIRTY_DAYS_MS), []);
const initialTo = useMemo(() => new Date(), []);
const [dateValue, setDateValue] = useState<DateRange>({ from: initialFrom, to: initialTo });
const startTime = dateValue.from ?? null;
const endTime = dateValue.to ?? null;
const effectiveUserId = all_admin_roles.includes(userRole) ? null : userId;
const { data, loading, isFetchingMore } = usePaginatedDailyActivity({
fetchFn: userDailyActivityCall,
args: [accessToken, startTime, endTime, effectiveUserId],
enabled: !!accessToken && !!startTime && !!endTime,
});
return {
dateValue,
onDateChange: setDateValue,
results: data.results as DailyData[],
loading,
isFetchingMore,
};
};

View file

@ -0,0 +1,82 @@
/* @vitest-environment jsdom */
import { renderHook } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
const { mockPush, navState } = vi.hoisted(() => ({
mockPush: vi.fn(),
navState: { pathname: "/logs" },
}));
vi.mock("next/navigation", () => ({
usePathname: () => navState.pathname,
useRouter: () => ({ push: mockPush }),
}));
vi.mock("@/components/networking", () => ({ serverRootPath: "" }));
import { createTabRoutes } from "@/utils/tabRoutes";
import { useTabRouting } from "./useTabRouting";
const routes = createTabRoutes("logs", ["audit", "deleted-keys", "deleted-teams"] as const);
const render = (ready = true) => {
const config = {
routes,
baseTabKey: "request-logs",
visibleKeys: ["audit", "deleted-keys", "deleted-teams"],
ready,
};
return renderHook(() => useTabRouting(config));
};
describe("useTabRouting", () => {
beforeEach(() => {
navState.pathname = "/logs";
mockPush.mockClear();
});
it("maps the base path to the base tab key", () => {
const { result } = render();
expect(result.current.activeSlug).toBe("");
expect(result.current.activeKey).toBe("request-logs");
});
it("uses the slug itself as the active key for a known nested tab", () => {
navState.pathname = "/ui/logs/audit";
const { result } = render();
expect(result.current.activeKey).toBe("audit");
});
it("falls back to the base tab key for an unknown slug", () => {
navState.pathname = "/ui/logs/bogus";
const { result } = render();
expect(result.current.activeKey).toBe("request-logs");
});
it("redirects an unknown slug to the base href once ready", () => {
const replaceMock = vi.fn();
const originalLocation = window.location;
Object.defineProperty(window, "location", { configurable: true, value: { replace: replaceMock } });
navState.pathname = "/ui/logs/bogus";
render(true);
expect(replaceMock).toHaveBeenCalledWith("/ui/logs/");
Object.defineProperty(window, "location", { configurable: true, value: originalLocation });
});
it("does not redirect while not ready (role/creds still loading)", () => {
const replaceMock = vi.fn();
const originalLocation = window.location;
Object.defineProperty(window, "location", { configurable: true, value: { replace: replaceMock } });
navState.pathname = "/ui/logs/bogus";
render(false);
expect(replaceMock).not.toHaveBeenCalled();
Object.defineProperty(window, "location", { configurable: true, value: originalLocation });
});
it("pushes the tab href on change, mapping the base key back to the empty slug", () => {
const { result } = render();
result.current.onTabChange("audit");
expect(mockPush).toHaveBeenCalledWith("/ui/logs/audit/");
result.current.onTabChange("request-logs");
expect(mockPush).toHaveBeenCalledWith("/ui/logs/");
});
});

View file

@ -0,0 +1,38 @@
import { useEffect } from "react";
import { usePathname, useRouter } from "next/navigation";
import type { TabRoutes } from "@/utils/tabRoutes";
interface UseTabRoutingArgs {
routes: Pick<TabRoutes<string>, "tabHref" | "slugFromPathname">;
baseTabKey: string;
visibleKeys: readonly string[];
ready?: boolean;
}
interface TabRoutingState {
activeSlug: string;
activeKey: string;
onTabChange: (key: string) => void;
}
export function useTabRouting({ routes, baseTabKey, visibleKeys, ready = true }: UseTabRoutingArgs): TabRoutingState {
const { tabHref, slugFromPathname } = routes;
const pathname = usePathname();
const router = useRouter();
const activeSlug = slugFromPathname(pathname);
const isKnownSlug = activeSlug === "" || visibleKeys.includes(activeSlug);
const activeKey = isKnownSlug ? activeSlug || baseTabKey : baseTabKey;
useEffect(() => {
if (ready && activeSlug !== "" && !isKnownSlug) {
window.location.replace(tabHref(""));
}
}, [ready, activeSlug, isKnownSlug, tabHref]);
const onTabChange = (key: string) => {
router.push(tabHref(key === baseTabKey ? "" : key));
};
return { activeSlug, activeKey, onTabChange };
}

View file

@ -1,662 +1,364 @@
import * as useAuthorizedModule from "@/app/(dashboard)/hooks/useAuthorized";
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import { renderWithProviders } from "../../../../../tests/test-utils";
import { render, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import AllModelsTab from "./AllModelsTab";
// Mock modelDeleteCall
import AllModelsTab from "./AllModelsTab";
import { STATUS_COLUMN_ID, toServerSortField } from "./ModelsTableColumns";
const mockModelDeleteCall = vi.fn().mockResolvedValue({});
const mockModelPatchUpdateCall = vi.fn().mockResolvedValue({});
vi.mock("@/components/networking", () => ({
modelDeleteCall: (...args: any[]) => mockModelDeleteCall(...args),
modelDeleteCall: (...args: unknown[]) => mockModelDeleteCall(...args),
modelPatchUpdateCall: (...args: unknown[]) => mockModelPatchUpdateCall(...args),
}));
// Mock NotificationsManager
vi.mock("@/components/molecules/notifications_manager", () => ({
default: {
success: vi.fn(),
fromBackend: vi.fn(),
default: { success: vi.fn(), fromBackend: vi.fn() },
}));
vi.mock("@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal", () => ({
default: function ModelSettingsModalMock({ isVisible }: { isVisible: boolean }) {
return isVisible ? <div data-testid="model-settings-modal" /> : null;
},
}));
// Mock react-query
const mockInvalidateQueries = vi.fn();
vi.mock("@tanstack/react-query", async (importOriginal) => {
const actual = (await importOriginal()) as any;
return {
...actual,
useQueryClient: () => ({
invalidateQueries: mockInvalidateQueries,
}),
};
const actual = await importOriginal<typeof import("@tanstack/react-query")>();
return { ...actual, useQueryClient: () => ({ invalidateQueries: mockInvalidateQueries }) };
});
// Mock the useModelsInfo hook
const mockUseModelsInfo = vi.fn(() => ({
data: { data: [], total_count: 0, current_page: 1, total_pages: 1, size: 50 },
isLoading: false,
error: null,
})) as any;
interface ModelsInfoArgs {
page?: number;
size?: number;
search?: string;
teamId?: string;
sortBy?: string;
sortOrder?: string;
}
const modelsInfoCalls: ModelsInfoArgs[] = [];
const mockRefetch = vi.fn();
let modelsInfoResult: Record<string, unknown> = {};
type UseModelsInfoArgs = [
page?: number,
size?: number,
search?: string,
modelId?: string,
teamId?: string,
sortBy?: string,
sortOrder?: string,
];
vi.mock("../../hooks/models/useModels", () => ({
useModelsInfo: (page?: number, size?: number, search?: string) => mockUseModelsInfo(page, size, search),
}));
// Mock the useModelCostMap hook
const mockUseModelCostMap = vi.fn(() => ({
data: {
"gpt-4": { litellm_provider: "openai" },
"gpt-3.5-turbo": { litellm_provider: "openai" },
"gpt-4-accessible": { litellm_provider: "openai" },
"gpt-3.5-turbo-blocked": { litellm_provider: "openai" },
"gpt-4-sales": { litellm_provider: "openai" },
"gpt-4-engineering": { litellm_provider: "openai" },
"gpt-4-personal": { litellm_provider: "openai" },
"gpt-4-team-only": { litellm_provider: "openai" },
"gpt-4-config": { litellm_provider: "openai" },
"gpt-4-db": { litellm_provider: "openai" },
useModelsInfo: (...args: UseModelsInfoArgs) => {
const [page, size, search, , teamId, sortBy, sortOrder] = args;
const call: ModelsInfoArgs = { page, size, search, teamId, sortBy, sortOrder };
modelsInfoCalls.push(call);
return { ...modelsInfoResult, refetch: mockRefetch };
},
isLoading: false,
error: null,
})) as any;
}));
vi.mock("../../hooks/models/useModelCostMap", () => ({
useModelCostMap: () => mockUseModelCostMap(),
useModelCostMap: () => ({ data: { "gpt-4": { litellm_provider: "openai" } }, isLoading: false, error: null }),
}));
// Mock the useTeams hook (react-query implementation)
const mockUseTeams = vi.fn(() => ({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
})) as any;
const mockTeams = [{ team_id: "team-1", team_alias: "Engineering" }];
vi.mock("../../hooks/teams/useTeams", () => ({
useTeams: () => mockUseTeams(),
useTeams: () => ({ data: mockTeams, isLoading: false, error: null, refetch: vi.fn() }),
}));
// Helper function to create model cost map mock return value
const createModelCostMapMock = (data: Record<string, any>) => ({
data,
isLoading: false,
error: null,
const BASE_MODEL_INFO = {
id: "model-1",
db_model: true,
created_by: "user-123",
created_at: "2024-01-01T00:00:00Z",
updated_at: "2024-01-02T00:00:00Z",
team_id: "team-1",
access_groups: [],
};
const makeRow = (overrides: Record<string, unknown> = {}) => ({
model_name: "gpt-4",
litellm_params: { model: "openai/gpt-4", custom_llm_provider: "openai" },
model_info: { ...BASE_MODEL_INFO, ...((overrides.model_info as Record<string, unknown>) ?? {}) },
});
// Helper function to create paginated model data mock
const createPaginatedModelData = (
models: any[],
totalCount: number = models.length,
currentPage: number = 1,
totalPages: number = 1,
size: number = 50,
) => ({
data: models,
total_count: totalCount,
current_page: currentPage,
total_pages: totalPages,
size: size,
});
const setModelsInfo = (rows: Record<string, unknown>[], totalCount = rows.length, isLoading = false) => {
modelsInfoResult = {
data: { data: rows, total_count: totalCount, current_page: 1, total_pages: 1, size: 50 },
isLoading,
isFetching: false,
error: null,
};
};
const lastModelsInfoCall = (): ModelsInfoArgs => modelsInfoCalls[modelsInfoCalls.length - 1];
const SEARCH_SETTLE_MS = 400;
const MOCK_AUTHORIZED = {
isLoading: false,
isAuthorized: true,
token: "mock-token",
accessToken: "mock-access-token",
userId: "user-123",
userEmail: "test@example.com",
userRole: "Admin",
premiumUser: true,
disabledPersonalKeyCreation: false,
showSSOBanner: false,
};
const mockSetSelectedModelGroup = vi.fn();
const mockSetSelectedModelId = vi.fn();
const mockSetSelectedTeamId = vi.fn();
const defaultProps = {
selectedModelGroup: "all",
setSelectedModelGroup: mockSetSelectedModelGroup,
availableModelGroups: ["gpt-4", "gpt-3.5-turbo"],
availableModelAccessGroups: ["sales-team"],
setSelectedModelId: mockSetSelectedModelId,
setSelectedTeamId: mockSetSelectedTeamId,
};
describe("AllModelsTab", () => {
const mockSetSelectedModelGroup = vi.fn();
const mockSetSelectedModelId = vi.fn();
const mockSetSelectedTeamId = vi.fn();
const defaultProps = {
selectedModelGroup: "all",
setSelectedModelGroup: mockSetSelectedModelGroup,
availableModelGroups: ["gpt-4", "gpt-3.5-turbo"],
availableModelAccessGroups: ["sales-team", "engineering-team"],
setSelectedModelId: mockSetSelectedModelId,
setSelectedTeamId: mockSetSelectedTeamId,
};
const mockUseAuthorized = {
token: "mock-token",
accessToken: "mock-access-token",
userId: "user-123",
userEmail: "test@example.com",
userRole: "Admin",
premiumUser: true,
disabledPersonalKeyCreation: false,
showSSOBanner: false,
};
beforeEach(() => {
vi.clearAllMocks();
vi.spyOn(useAuthorizedModule, "default").mockReturnValue(mockUseAuthorized);
modelsInfoCalls.length = 0;
setModelsInfo([makeRow()]);
vi.spyOn(useAuthorizedModule, "default").mockReturnValue(MOCK_AUTHORIZED);
});
it("should render with empty data", () => {
mockUseModelsInfo.mockReturnValueOnce({
data: createPaginatedModelData([], 0, 1, 1, 50),
isLoading: false,
error: null,
});
it("renders the fetched models and the server row count", async () => {
setModelsInfo([makeRow()], 137);
render(<AllModelsTab {...defaultProps} />);
mockUseTeams.mockReturnValueOnce({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
});
mockUseModelCostMap.mockReturnValueOnce(createModelCostMapMock({}));
renderWithProviders(<AllModelsTab {...defaultProps} />);
expect(screen.getByText("Current Team:")).toBeInTheDocument();
expect(await screen.findByText("gpt-4")).toBeInTheDocument();
expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 1-50 of 137");
});
it("should filter models by direct team access when current team is selected", async () => {
const mockTeams = [
{
team_id: "team-456",
team_alias: "Engineering Team",
models: ["gpt-4"],
max_budget: null,
budget_duration: null,
tpm_limit: null,
rpm_limit: null,
organization_id: "org-123",
created_at: "2024-01-01",
keys: [],
members_with_roles: [],
},
it("does not re-query after the mount-time debounced search settles unchanged", async () => {
render(<AllModelsTab {...defaultProps} />);
const callsAfterMount = modelsInfoCalls.length;
await new Promise((resolve) => setTimeout(resolve, SEARCH_SETTLE_MS));
expect(modelsInfoCalls.length).toBe(callsAfterMount);
});
it("shows the empty state when the proxy returns no models", () => {
setModelsInfo([], 0);
render(<AllModelsTab {...defaultProps} />);
expect(screen.getByText("No models found")).toBeInTheDocument();
});
it("shows the loading skeleton while the first page is in flight", () => {
setModelsInfo([], 0, true);
render(<AllModelsTab {...defaultProps} />);
expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0);
expect(screen.queryByText("No models found")).not.toBeInTheDocument();
});
describe("server sort contract", () => {
const sortHeader = (columnId: string): HTMLElement => screen.getByTestId(`sort-header-${columnId}`);
const expectIndicator = async (columnId: string, state: "asc" | "desc" | "none") => {
await waitFor(() => {
expect(sortHeader(columnId).querySelector(`[data-sort-indicator="${state}"]`)).not.toBeNull();
});
};
const cases: [string, string, string, "asc" | "desc"][] = [
["Model Information", "model_name", "model_name", "asc"],
["Created By", "model_info_created_by", "created_at", "asc"],
["Updated At", "model_info_updated_at", "updated_at", "asc"],
["Costs", "input_cost", "costs", "desc"],
];
mockUseTeams.mockReturnValueOnce({
data: mockTeams,
isLoading: false,
error: null,
refetch: vi.fn(),
it.each(cases)("sorts %s using the server field %s", async (_label, columnId, serverField, firstDirection) => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
await user.click(sortHeader(columnId));
await expectIndicator(columnId, firstDirection);
expect(lastModelsInfoCall().sortBy).toBe(serverField);
expect(lastModelsInfoCall().sortOrder).toBe(firstDirection);
});
mockUseModelCostMap.mockReturnValueOnce(
createModelCostMapMock({
"gpt-4-accessible": { litellm_provider: "openai" },
"gpt-3.5-turbo-blocked": { litellm_provider: "openai" },
}),
);
it("maps the hidden Status column to the server field status", () => {
expect(toServerSortField(STATUS_COLUMN_ID)).toBe("status");
});
const modelData = createPaginatedModelData(
[
{
model_name: "gpt-4-accessible",
model_info: {
id: "model-1",
access_via_team_ids: ["team-456"],
access_groups: [],
},
},
{
model_name: "gpt-3.5-turbo-blocked",
model_info: {
id: "model-2",
access_via_team_ids: ["team-789"],
access_groups: [],
},
},
],
2,
1,
1,
50,
);
it("cycles a sorted column back to unsorted", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null });
await user.click(sortHeader("model_info_updated_at"));
await expectIndicator("model_info_updated_at", "asc");
expect(lastModelsInfoCall().sortOrder).toBe("asc");
renderWithProviders(<AllModelsTab {...defaultProps} />);
await user.click(sortHeader("model_info_updated_at"));
await expectIndicator("model_info_updated_at", "desc");
expect(lastModelsInfoCall().sortOrder).toBe("desc");
// Component shows API total_count (2), not filtered count
// Since default is "personal" team and models don't have direct_access, they're filtered out
await waitFor(() => {
expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument();
await user.click(sortHeader("model_info_updated_at"));
await expectIndicator("model_info_updated_at", "none");
expect(lastModelsInfoCall().sortBy).toBeUndefined();
});
});
it("should filter models by access group matching when team models match model access groups", async () => {
const mockTeams = [
{
team_id: "team-sales",
team_alias: "Sales Team",
models: ["sales-model-group"],
max_budget: null,
budget_duration: null,
tpm_limit: null,
rpm_limit: null,
organization_id: "org-123",
created_at: "2024-01-01",
keys: [],
members_with_roles: [],
},
];
it("queries the selected team and resets to the first page", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
mockUseTeams.mockReturnValue({
data: mockTeams,
isLoading: false,
error: null,
refetch: vi.fn(),
});
expect(lastModelsInfoCall().teamId).toBeUndefined();
mockUseModelCostMap.mockReturnValueOnce(
createModelCostMapMock({
"gpt-4-sales": { litellm_provider: "openai" },
"gpt-4-engineering": { litellm_provider: "openai" },
}),
);
await user.click(screen.getByTestId("models-team-select"));
await user.click(await screen.findByRole("option", { name: "Engineering" }));
const modelData = createPaginatedModelData(
[
{
model_name: "gpt-4-sales",
model_info: {
id: "model-sales-1",
access_via_team_ids: [],
access_groups: ["sales-model-group"],
},
},
{
model_name: "gpt-4-engineering",
model_info: {
id: "model-eng-1",
access_via_team_ids: [],
access_groups: ["engineering-model-group"],
},
},
],
2,
1,
1,
50,
);
mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null });
renderWithProviders(<AllModelsTab {...defaultProps} />);
// Component shows API total_count (2), not filtered count
// Since default is "personal" team and models don't have direct_access, they're filtered out
await waitFor(() => {
expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument();
expect(lastModelsInfoCall().teamId).toBe("team-1");
});
expect(lastModelsInfoCall().page).toBe(1);
});
it("debounces the model name search into the server query", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
await user.type(screen.getByTestId("datatable-search"), "claude");
await waitFor(() => {
expect(lastModelsInfoCall().search).toBe("claude");
});
});
it("should filter models by direct_access for personal team", async () => {
mockUseTeams.mockReturnValue({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
});
it("applies a public model name filter through the drawer", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
mockUseModelCostMap.mockReturnValueOnce(
createModelCostMapMock({
"gpt-4-personal": { litellm_provider: "openai" },
"gpt-4-team-only": { litellm_provider: "openai" },
}),
);
await user.click(screen.getByTestId("datatable-filters-trigger"));
await user.click(await screen.findByPlaceholderText("Filter by Public Model Name"));
await user.click(await screen.findByRole("option", { name: "gpt-3.5-turbo" }));
await user.click(screen.getByTestId("filter-drawer-apply"));
const modelData = createPaginatedModelData(
[
{
model_name: "gpt-4-personal",
model_info: {
id: "model-personal-1",
direct_access: true,
access_via_team_ids: [],
access_groups: [],
},
},
{
model_name: "gpt-4-team-only",
model_info: {
id: "model-team-1",
direct_access: false,
access_via_team_ids: ["team-123"],
access_groups: [],
},
},
],
2,
1,
1,
50,
);
mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null });
renderWithProviders(<AllModelsTab {...defaultProps} />);
// Component shows API total_count (2), but only 1 model has direct_access
await waitFor(() => {
expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument();
expect(mockSetSelectedModelGroup).toHaveBeenCalledWith("gpt-3.5-turbo");
});
});
it("should show config model status for models defined in configs", async () => {
mockUseTeams.mockReturnValue({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
});
it("filters the fetched page down to the selected model group", () => {
setModelsInfo([makeRow(), { ...makeRow(), model_name: "claude-opus" }], 2);
render(<AllModelsTab {...defaultProps} selectedModelGroup="claude-opus" />);
mockUseModelCostMap.mockReturnValueOnce(
createModelCostMapMock({
"gpt-4-config": { litellm_provider: "openai" },
"gpt-4-db": { litellm_provider: "openai" },
}),
);
const table = screen.getByRole("table");
expect(within(table).getByText("claude-opus")).toBeInTheDocument();
expect(within(table).queryByText("gpt-4")).not.toBeInTheDocument();
});
const modelData = createPaginatedModelData(
[
{
model_name: "gpt-4-config",
litellm_model_name: "gpt-4-config",
provider: "openai",
model_info: {
id: "model-config-1",
db_model: false,
direct_access: true,
access_via_team_ids: [],
access_groups: [],
created_by: "user-123",
created_at: "2024-01-01",
updated_at: "2024-01-01",
},
},
{
model_name: "gpt-4-db",
litellm_model_name: "gpt-4-db",
provider: "openai",
model_info: {
id: "model-db-1",
db_model: true,
direct_access: true,
access_via_team_ids: [],
access_groups: [],
created_by: "user-123",
created_at: "2024-01-01",
updated_at: "2024-01-01",
},
},
],
2,
1,
1,
50,
);
it("resets search, filters, team and sorting from the drawer reset button", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} selectedModelGroup="gpt-4" />);
mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null });
await user.click(screen.getByTestId("models-team-select"));
await user.click(await screen.findByRole("option", { name: "Engineering" }));
await waitFor(() => expect(lastModelsInfoCall().teamId).toBe("team-1"));
renderWithProviders(<AllModelsTab {...defaultProps} />);
await user.click(screen.getByTestId("datatable-filters-trigger"));
await user.click(await screen.findByTestId("filter-drawer-reset"));
expect(mockSetSelectedModelGroup).toHaveBeenCalledWith("all");
await waitFor(() => {
expect(screen.getByText("Config Model")).toBeInTheDocument();
expect(screen.getByText("DB Model")).toBeInTheDocument();
expect(lastModelsInfoCall().teamId).toBeUndefined();
});
});
it("should show 'Defined in config' for models defined in configs", async () => {
mockUseTeams.mockReturnValue({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
});
it("opens the delete modal from the row and deletes the model", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
mockUseModelCostMap.mockReturnValueOnce(
createModelCostMapMock({
"gpt-4-config": { litellm_provider: "openai" },
}),
);
await user.click(await screen.findByTestId("model-delete-model-1"));
expect(await screen.findByText("Delete Model")).toBeInTheDocument();
const modelData = createPaginatedModelData(
[
{
model_name: "gpt-4-config",
litellm_model_name: "gpt-4-config",
provider: "openai",
model_info: {
id: "model-config-1",
db_model: false,
direct_access: true,
access_via_team_ids: [],
access_groups: [],
created_by: "user-123",
created_at: "2024-01-01",
updated_at: "2024-01-01",
},
},
],
1,
1,
1,
50,
);
mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null });
renderWithProviders(<AllModelsTab {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /^delete$/i }));
await waitFor(() => {
expect(screen.getByText("Defined in config")).toBeInTheDocument();
expect(mockModelDeleteCall).toHaveBeenCalledWith("mock-access-token", "model-1");
});
});
it("should handle pagination: Previous button is disabled on first page and Next button works", async () => {
mockUseTeams.mockReturnValue({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
});
it("pauses a model through the row toggle", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
mockUseModelCostMap.mockReturnValue(
createModelCostMapMock({
"gpt-4-page1": { litellm_provider: "openai" },
"gpt-4-page2": { litellm_provider: "openai" },
}),
);
// Mock first page response (page 1 of 2)
const page1Data = createPaginatedModelData(
[
{
model_name: "gpt-4-page1",
model_info: {
id: "model-page1-1",
direct_access: true,
access_via_team_ids: [],
access_groups: [],
},
},
],
2, // total_count
1, // current_page
2, // total_pages
50, // size
);
// Set up mock to return page1Data for page 1
mockUseModelsInfo.mockImplementation((page: number = 1, size?: number, search?: string) => {
return { data: page1Data, isLoading: false, error: null };
});
renderWithProviders(<AllModelsTab {...defaultProps} />);
await user.click(await screen.findByTestId("model-pause-toggle-model-1"));
await waitFor(() => {
// Component calculates: ((1-1)*50)+1 = 1, Math.min(1*50, 2) = 2
expect(screen.getByText("Showing 1 - 2 of 2 results")).toBeInTheDocument();
expect(mockModelPatchUpdateCall).toHaveBeenCalledWith("mock-access-token", { blocked: true }, "model-1");
});
// Check that Previous button is disabled on first page
const previousButton = screen.getByRole("button", { name: /previous/i });
expect(previousButton).toBeDisabled();
// Check that Next button is enabled (since we're on page 1 of 2)
const nextButton = screen.getByRole("button", { name: /next/i });
expect(nextButton).not.toBeDisabled();
});
it("should handle pagination: Next button is disabled on last page", async () => {
mockUseTeams.mockReturnValue({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
});
it("opens the model settings modal from the toolbar", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
mockUseModelCostMap.mockReturnValue(
createModelCostMapMock({
"gpt-4-page2": { litellm_provider: "openai" },
}),
);
// Mock single page response (page 1 of 1 - last page)
const singlePageData = createPaginatedModelData(
[
{
model_name: "gpt-4-page2",
model_info: {
id: "model-page2-1",
direct_access: true,
access_via_team_ids: [],
access_groups: [],
},
},
],
1, // total_count
1, // current_page
1, // total_pages (only 1 page, so this is the last page)
50, // size
);
mockUseModelsInfo.mockImplementation((page?: number, size?: number, search?: string) => {
return { data: singlePageData, isLoading: false, error: null };
});
renderWithProviders(<AllModelsTab {...defaultProps} />);
await waitFor(() => {
expect(screen.getByText("Showing 1 - 1 of 1 results")).toBeInTheDocument();
});
// When there's only 1 page (last page), Next should be disabled
const nextButton = screen.getByRole("button", { name: /next/i });
expect(nextButton).toBeDisabled();
// Previous should also be disabled on the first (and only) page
const previousButton = screen.getByRole("button", { name: /previous/i });
expect(previousButton).toBeDisabled();
expect(screen.queryByTestId("model-settings-modal")).not.toBeInTheDocument();
await user.click(screen.getByTestId("models-settings-trigger"));
expect(screen.getByTestId("model-settings-modal")).toBeInTheDocument();
});
it("should pass setDeleteModalModelId to columns for delete functionality", async () => {
// This test verifies that the delete modal setter is passed to columns
// The actual modal rendering is handled by DeleteResourceModal component
mockUseTeams.mockReturnValue({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
});
it("opens the model detail view from the model ID cell", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
mockUseModelCostMap.mockReturnValue(
createModelCostMapMock({
"gpt-4-delete-test": { litellm_provider: "openai" },
}),
);
await user.click(await screen.findByTestId("model-id-model-1"));
const modelData = createPaginatedModelData(
[
{
model_name: "gpt-4-delete-test",
litellm_model_name: "gpt-4-delete-test",
provider: "openai",
model_info: {
id: "model-to-delete",
db_model: true,
direct_access: true,
access_via_team_ids: [],
access_groups: [],
created_by: "user-123",
created_at: "2024-01-01",
updated_at: "2024-01-01",
},
},
],
1,
1,
1,
50,
);
mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null, refetch: vi.fn() });
renderWithProviders(<AllModelsTab {...defaultProps} />);
await waitFor(() => {
expect(screen.getByText("gpt-4-delete-test")).toBeInTheDocument();
});
// Verify the DB Model badge is shown (indicating it can be deleted)
expect(screen.getByText("DB Model")).toBeInTheDocument();
expect(mockSetSelectedModelId).toHaveBeenCalledWith("model-1");
});
it("should render clickable model ID that calls setSelectedModelId", async () => {
mockUseTeams.mockReturnValue({
data: [],
isLoading: false,
error: null,
refetch: vi.fn(),
it("opens the team detail view from the team ID cell", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
await user.click(await screen.findByTestId("model-team-id-model-1"));
expect(mockSetSelectedTeamId).toHaveBeenCalledWith("team-1");
});
describe("virtual key hint", () => {
it("explains personal key creation while viewing current team models", () => {
render(<AllModelsTab {...defaultProps} />);
expect(screen.getByText(/create a Virtual Key without selecting a team/i)).toBeInTheDocument();
});
mockUseModelCostMap.mockReturnValue(
createModelCostMapMock({
"gpt-4-clickable": { litellm_provider: "openai" },
}),
);
it("names the selected team in the hint", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
const modelData = createPaginatedModelData(
[
{
model_name: "gpt-4-clickable",
litellm_model_name: "gpt-4-clickable",
provider: "openai",
model_info: {
id: "clickable-model-id",
db_model: true,
direct_access: true,
access_via_team_ids: [],
access_groups: [],
created_by: "user-123",
created_at: "2024-01-01",
updated_at: "2024-01-01",
},
},
],
1,
1,
1,
50,
);
await user.click(screen.getByTestId("models-team-select"));
await user.click(await screen.findByRole("option", { name: "Engineering" }));
mockUseModelsInfo.mockReturnValue({ data: modelData, isLoading: false, error: null, refetch: vi.fn() });
renderWithProviders(<AllModelsTab {...defaultProps} />);
await waitFor(() => {
expect(screen.getByText("gpt-4-clickable")).toBeInTheDocument();
expect(await screen.findByText(/select Team as "Engineering"/i)).toBeInTheDocument();
});
// Click on the Model ID cell which should call setSelectedModelId
const modelIdCell = screen.getByText("clickable-model-id");
expect(modelIdCell).toBeInTheDocument();
it("hides the hint when viewing all available models", async () => {
const user = userEvent.setup();
render(<AllModelsTab {...defaultProps} />);
fireEvent.click(modelIdCell);
await user.click(screen.getByTestId("models-view-select"));
await user.click(await screen.findByRole("option", { name: "All Available Models" }));
await waitFor(() => {
expect(mockSetSelectedModelId).toHaveBeenCalledWith("clickable-model-id");
await waitFor(() => {
expect(screen.queryByText(/create a Virtual Key/i)).not.toBeInTheDocument();
});
});
});
});

View file

@ -1,27 +1,33 @@
"use client";
import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap";
import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { Team } from "@/components/key_team_helpers/key_list";
import { AllModelsDataTable } from "@/components/model_dashboard/all_models_table";
import { columns } from "@/components/molecules/models/columns";
import { getDisplayModelName } from "@/components/view_model/model_name_display";
import DeleteResourceModal from "@/components/common_components/DeleteResourceModal";
import ModelSettingsModal from "@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal";
import { ModelData } from "@/components/model_dashboard/types";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { modelDeleteCall, modelPatchUpdateCall } from "@/components/networking";
import { InfoCircleOutlined, SettingOutlined } from "@ant-design/icons";
import { PaginationState, SortingState } from "@tanstack/react-table";
import { useQueryClient } from "@tanstack/react-query";
import { Grid } from "@tremor/react";
import { Badge, Button, Select, Skeleton, Space, Typography } from "antd";
import ModelSettingsModal from "@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal";
import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer";
import { useEffect, useMemo, useState } from "react";
import { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table";
import { Info } from "lucide-react";
import { useCallback, useEffect, useMemo, useState } from "react";
import { useModelsInfo } from "../../hooks/models/useModels";
import { transformModelData } from "../utils/modelDataTransformer";
type ModelViewMode = "all" | "current_team";
import {
ALL_MODEL_GROUPS_VALUE,
AllModelsTable,
ModelViewMode,
PERSONAL_TEAM_VALUE,
WILDCARD_MODEL_GROUP_VALUE,
} from "./AllModelsTable";
import { ACCESS_GROUPS_COLUMN_ID, MODEL_NAME_COLUMN_ID, toServerSortField } from "./ModelsTableColumns";
const SEARCH_DEBOUNCE_WAIT_MS = 200;
const { Text } = Typography;
const DEFAULT_PAGE_SIZE = 50;
const DEFAULT_PAGINATION: PaginationState = { pageIndex: 0, pageSize: DEFAULT_PAGE_SIZE };
interface AllModelsTabProps {
selectedModelGroup: string | null;
@ -41,31 +47,30 @@ const AllModelsTab = ({
setSelectedTeamId,
}: AllModelsTabProps) => {
const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap();
const { accessToken, userId, userRole, premiumUser } = useAuthorized();
const { accessToken, userId, userRole } = useAuthorized();
const { data: teams, isLoading: isLoadingTeams } = useTeams();
const queryClient = useQueryClient();
const [modelNameSearch, setModelNameSearch] = useState<string>("");
const [debouncedSearch, setDebouncedSearch] = useState<string>("");
const [modelViewMode, setModelViewMode] = useState<ModelViewMode>("current_team");
const [currentTeam, setCurrentTeam] = useState<Team | "personal">("personal");
const [showFilters, setShowFilters] = useState<boolean>(false);
const [selectedTeamValue, setSelectedTeamValue] = useState<string>(PERSONAL_TEAM_VALUE);
const [selectedModelAccessGroupFilter, setSelectedModelAccessGroupFilter] = useState<string | null>(null);
const [expandedRows, setExpandedRows] = useState<Set<string>>(new Set());
const [currentPage, setCurrentPage] = useState<number>(1);
const [pageSize] = useState<number>(50);
const [pagination, setPagination] = useState<PaginationState>({
pageIndex: 0,
pageSize: 50,
});
const [pagination, setPagination] = useState<PaginationState>(DEFAULT_PAGINATION);
const [sorting, setSorting] = useState<SortingState>([]);
const [isModelSettingsModalVisible, setIsModelSettingsModalVisible] = useState(false);
const [deleteModalModelId, setDeleteModalModelId] = useState<string | null>(null);
const [deleteLoading, setDeleteLoading] = useState(false);
const [pausingModelId, setPausingModelId] = useState<string | null>(null);
const resetToFirstPage = useCallback(() => {
setPagination((previous) => (previous.pageIndex === 0 ? previous : { ...previous, pageIndex: 0 }));
}, []);
const debouncedUpdateSearch = useDebouncedCallback(
(value: string) => {
setDebouncedSearch(value);
setCurrentPage(1);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
resetToFirstPage();
},
{ wait: SEARCH_DEBOUNCE_WAIT_MS },
);
@ -74,125 +79,130 @@ const AllModelsTab = ({
debouncedUpdateSearch(modelNameSearch);
}, [modelNameSearch, debouncedUpdateSearch]);
// Determine teamId to pass to the query - only pass if not "personal"
const teamIdForQuery = currentTeam === "personal" ? undefined : currentTeam.team_id;
const teamIdForQuery = selectedTeamValue === PERSONAL_TEAM_VALUE ? undefined : selectedTeamValue;
// Convert sorting state to sortBy and sortOrder for API
const sortBy = useMemo(() => {
if (sorting.length === 0) return undefined;
const sort = sorting[0];
const columnIdToServerField: Record<string, string> = {
input_cost: "costs", // Map input_cost column to "costs" for server-side sorting
model_info_db_model: "status", // Map model_info.db_model column to "status" for server-side sorting
model_info_created_by: "created_at", // Map model_info.created_by column to "created_at" for server-side sorting
model_info_updated_at: "updated_at", // Map model_info.updated_at column to "updated_at" for server-side sorting
};
return columnIdToServerField[sort.id] || sort.id;
return toServerSortField(sorting[0].id);
}, [sorting]);
const sortOrder = useMemo(() => {
if (sorting.length === 0) return undefined;
const sort = sorting[0];
return sort.desc ? "desc" : "asc";
return sorting[0].desc ? "desc" : "asc";
}, [sorting]);
const {
data: rawModelData,
isLoading: isLoadingModelsInfo,
isFetching: isFetchingModelsInfo,
refetch: refetchModels,
} = useModelsInfo(currentPage, pageSize, debouncedSearch || undefined, undefined, teamIdForQuery, sortBy, sortOrder);
} = useModelsInfo(
pagination.pageIndex + 1,
pagination.pageSize,
debouncedSearch || undefined,
undefined,
teamIdForQuery,
sortBy,
sortOrder,
);
const isLoading = isLoadingModelsInfo || isLoadingModelCostMap;
const getProviderFromModel = (model: string) => {
if (modelCostMapData !== null && modelCostMapData !== undefined) {
if (typeof modelCostMapData == "object" && model in modelCostMapData) {
return modelCostMapData[model]["litellm_provider"];
const getProviderFromModel = useCallback(
(model: string) => {
if (modelCostMapData !== null && modelCostMapData !== undefined) {
if (typeof modelCostMapData == "object" && model in modelCostMapData) {
return modelCostMapData[model]["litellm_provider"];
}
}
}
return "openai";
};
return "openai";
},
[modelCostMapData],
);
const modelData = useMemo(() => {
if (!rawModelData) return { data: [] };
return transformModelData(rawModelData, getProviderFromModel);
}, [rawModelData, modelCostMapData]);
}, [rawModelData, getProviderFromModel]);
const [deleteModalModelId, setDeleteModalModelId] = useState<string | null>(null);
const [deleteLoading, setDeleteLoading] = useState(false);
// Get pagination metadata from the response
const paginationMeta = useMemo(() => {
if (!rawModelData) {
return {
total_count: 0,
current_page: 1,
total_pages: 1,
size: pageSize,
};
}
return {
total_count: rawModelData.total_count ?? 0,
current_page: rawModelData.current_page ?? 1,
total_pages: rawModelData.total_pages ?? 1,
size: rawModelData.size ?? pageSize,
};
}, [rawModelData, pageSize]);
const filteredData = useMemo(() => {
const filteredData = useMemo<ModelData[]>(() => {
if (!modelData || !modelData.data || modelData.data.length === 0) {
return [];
}
// Server-side search is now handled by the API, so we only filter by other criteria
return modelData.data.filter((model: any) => {
return modelData.data.filter((model: ModelData) => {
const modelNameMatch =
selectedModelGroup === "all" ||
selectedModelGroup === ALL_MODEL_GROUPS_VALUE ||
model.model_name === selectedModelGroup ||
!selectedModelGroup ||
(selectedModelGroup === "wildcard" && model.model_name?.includes("*"));
(selectedModelGroup === WILDCARD_MODEL_GROUP_VALUE && model.model_name?.includes("*"));
const accessGroupMatch =
selectedModelAccessGroupFilter === "all" ||
model.model_info["access_groups"]?.includes(selectedModelAccessGroupFilter) ||
selectedModelAccessGroupFilter === ALL_MODEL_GROUPS_VALUE ||
model.model_info["access_groups"]?.includes(selectedModelAccessGroupFilter ?? "") ||
!selectedModelAccessGroupFilter;
// Team filtering is now handled server-side via teamId query parameter
// Only apply client-side filtering for model groups and access groups
return modelNameMatch && accessGroupMatch;
});
}, [modelData, selectedModelGroup, selectedModelAccessGroupFilter]);
useEffect(() => {
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
setCurrentPage(1);
}, [selectedModelGroup, selectedModelAccessGroupFilter]);
const columnFilters = useMemo<ColumnFiltersState>(
() =>
[
selectedModelGroup && selectedModelGroup !== ALL_MODEL_GROUPS_VALUE
? { id: MODEL_NAME_COLUMN_ID, value: selectedModelGroup }
: null,
selectedModelAccessGroupFilter ? { id: ACCESS_GROUPS_COLUMN_ID, value: selectedModelAccessGroupFilter } : null,
].filter((entry) => entry !== null),
[selectedModelGroup, selectedModelAccessGroupFilter],
);
// Reset pagination when team changes
useEffect(() => {
setCurrentPage(1);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}, [teamIdForQuery]);
const handleColumnFiltersChange: OnChangeFn<ColumnFiltersState> = (updater) => {
const next = typeof updater === "function" ? updater(columnFilters) : updater;
const modelGroup = next.find((entry) => entry.id === MODEL_NAME_COLUMN_ID)?.value;
const accessGroup = next.find((entry) => entry.id === ACCESS_GROUPS_COLUMN_ID)?.value;
setSelectedModelGroup(typeof modelGroup === "string" ? modelGroup : ALL_MODEL_GROUPS_VALUE);
setSelectedModelAccessGroupFilter(typeof accessGroup === "string" ? accessGroup : null);
resetToFirstPage();
};
// Reset pagination when sorting changes
useEffect(() => {
setCurrentPage(1);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}, [sorting]);
const handleSortingChange: OnChangeFn<SortingState> = (updater) => {
setSorting(typeof updater === "function" ? updater(sorting) : updater);
resetToFirstPage();
};
const handleTeamChange = (value: string) => {
setSelectedTeamValue(value);
resetToFirstPage();
};
const resetFilters = () => {
setModelNameSearch("");
setSelectedModelGroup("all");
setSelectedModelGroup(ALL_MODEL_GROUPS_VALUE);
setSelectedModelAccessGroupFilter(null);
setCurrentTeam("personal");
setSelectedTeamValue(PERSONAL_TEAM_VALUE);
setModelViewMode("current_team");
setCurrentPage(1);
setPagination({ pageIndex: 0, pageSize: 50 });
setPagination(DEFAULT_PAGINATION);
setSorting([]);
};
const teamOptions = useMemo(
() => [
{ value: PERSONAL_TEAM_VALUE, label: "Personal" },
...(teams ?? [])
.filter((team) => team.team_id)
.map((team) => ({ value: team.team_id, label: team.team_alias ? team.team_alias : team.team_id })),
],
[teams],
);
const selectedTeam = useMemo(
() => (teams ?? []).find((team) => team.team_id === selectedTeamValue) ?? null,
[teams, selectedTeamValue],
);
const modelToDelete = useMemo(() => {
if (!deleteModalModelId || !modelData?.data) return null;
return modelData.data.find((model: any) => model.model_info.id === deleteModalModelId);
return modelData.data.find((model: ModelData) => model.model_info.id === deleteModalModelId);
}, [deleteModalModelId, modelData]);
const handleDeleteModel = async () => {
@ -212,356 +222,99 @@ const AllModelsTab = ({
}
};
const [pausingModelId, setPausingModelId] = useState<string | null>(null);
const handleTogglePause = useCallback(
async (modelId: string, blocked: boolean) => {
if (!accessToken) return;
try {
setPausingModelId(modelId);
await modelPatchUpdateCall(accessToken, { blocked }, modelId);
NotificationsManager.success(blocked ? "Model paused" : "Model resumed");
// invalidateQueries already schedules a refetch for active observers
// on this key — no need to also call refetchModels() (would double-fetch).
queryClient.invalidateQueries({ queryKey: ["models", "list"] });
} catch (error) {
console.error("Error toggling model pause state:", error);
NotificationsManager.fromBackend(error);
} finally {
setPausingModelId(null);
}
},
[accessToken, queryClient],
);
const handleTogglePause = async (modelId: string, blocked: boolean) => {
if (!accessToken) return;
try {
setPausingModelId(modelId);
await modelPatchUpdateCall(accessToken, { blocked }, modelId);
NotificationsManager.success(blocked ? "Model paused" : "Model resumed");
// invalidateQueries already schedules a refetch for active observers
// on this key — no need to also call refetchModels() (would double-fetch).
queryClient.invalidateQueries({ queryKey: ["models", "list"] });
} catch (error) {
console.error("Error toggling model pause state:", error);
NotificationsManager.fromBackend(error);
} finally {
setPausingModelId(null);
}
};
const handleRefresh = useCallback(() => {
void refetchModels();
}, [refetchModels]);
const handleDeleteClick = useCallback((modelId: string) => {
setDeleteModalModelId(modelId);
}, []);
const handleOpenModelSettings = useCallback(() => {
setIsModelSettingsModalVisible(true);
}, []);
const teamAccessLabel = selectedTeam?.team_alias || selectedTeam?.team_id || "";
return (
<div className="w-full">
<Grid>
<div className="flex flex-col space-y-4">
<div className="bg-white rounded-lg shadow-sm">
{/* Current Team and View Mode Selector - Prominent Section */}
<div className="border-b px-6 py-4 bg-gray-50">
<div className="flex items-center justify-between">
<div className="flex items-center gap-4">
<Text className="text-lg font-semibold text-gray-900">Current Team:</Text>
<div className="w-80">
{isLoading ? (
<Skeleton.Input active block size="large" />
) : (
<Select
style={{ width: "100%" }}
size="large"
defaultValue="personal"
value={currentTeam === "personal" ? "personal" : currentTeam.team_id}
onChange={(value) => {
if (value === "personal") {
setCurrentTeam("personal");
// Reset to page 1 when team changes
setCurrentPage(1);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
} else {
const team = teams?.find((t) => t.team_id === value);
if (team) {
setCurrentTeam(team);
// Reset to page 1 when team changes
setCurrentPage(1);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}
}
}}
loading={isLoadingTeams}
options={[
{
value: "personal",
label: (
<Space direction="horizontal" align="center">
<Badge color="blue" size="small" />
<Text style={{ fontSize: 16 }}>Personal</Text>
</Space>
),
},
...(teams
?.filter((team) => team.team_id)
.map((team) => ({
value: team.team_id,
label: (
<Space direction="horizontal" align="center">
<Badge color="green" size="small" />
<Text ellipsis style={{ fontSize: 16 }}>
{team.team_alias ? team.team_alias : team.team_id}
</Text>
</Space>
),
})) ?? []),
]}
/>
)}
</div>
</div>
<div className="flex items-center gap-4">
<Text className="text-lg font-semibold text-gray-900">View:</Text>
<div className="w-64">
{isLoading ? (
<Skeleton.Input active block size="large" />
) : (
<Select
style={{ width: "100%" }}
size="large"
defaultValue="current_team"
value={modelViewMode}
onChange={(value) => setModelViewMode(value as "current_team" | "all")}
options={[
{
value: "current_team",
label: (
<Space direction="horizontal" align="center">
<Badge color="purple" size="small" />
<Text style={{ fontSize: 16 }}>Current Team Models</Text>
</Space>
),
},
{
value: "all",
label: (
<Space direction="horizontal" align="center">
<Badge color="gray" size="small" />
<Text style={{ fontSize: 16 }}>All Available Models</Text>
</Space>
),
},
]}
/>
)}
</div>
</div>
</div>
<div className="flex flex-col gap-3">
<AllModelsTable
data={filteredData}
rowCount={rawModelData?.total_count ?? 0}
isLoading={isLoading}
isRefreshing={isFetchingModelsInfo}
onRefresh={handleRefresh}
sorting={sorting}
onSortingChange={handleSortingChange}
pagination={pagination}
onPaginationChange={setPagination}
columnFilters={columnFilters}
onColumnFiltersChange={handleColumnFiltersChange}
onResetFilters={resetFilters}
searchValue={modelNameSearch}
onSearchChange={setModelNameSearch}
teamOptions={teamOptions}
selectedTeamValue={selectedTeamValue}
onTeamChange={handleTeamChange}
isLoadingTeams={isLoadingTeams}
viewMode={modelViewMode}
onViewModeChange={setModelViewMode}
onOpenModelSettings={handleOpenModelSettings}
availableModelGroups={availableModelGroups}
availableModelAccessGroups={availableModelAccessGroups}
userRole={userRole}
userID={userId}
onModelIdClick={setSelectedModelId}
onTeamIdClick={setSelectedTeamId}
onDeleteClick={handleDeleteClick}
onTogglePauseClick={handleTogglePause}
pausingModelId={pausingModelId}
/>
{modelViewMode === "current_team" && (
<div className="flex items-start gap-2 mt-3">
<InfoCircleOutlined className="text-gray-400 mt-0.5 shrink-0 text-xs" />
<div className="text-xs text-gray-500">
{currentTeam === "personal" ? (
<span>
To access these models: Create a Virtual Key without selecting a team on the{" "}
<a
href="/public?login=success&page=api-keys"
className="text-gray-600 hover:text-gray-800 underline"
>
Virtual Keys page
</a>
</span>
) : (
<span>
To access these models: Create a Virtual Key and select Team as &quot;
{typeof currentTeam !== "string" ? currentTeam.team_alias || currentTeam.team_id : ""}&quot; on
the{" "}
<a
href="/public?login=success&page=api-keys"
className="text-gray-600 hover:text-gray-800 underline"
>
Virtual Keys page
</a>
</span>
)}
</div>
</div>
)}
</div>
{/* Search and Filter Controls */}
<div className="border-b px-6 py-4">
<div className="flex flex-col space-y-4">
{/* Search and Filter Controls */}
<div className="flex items-center justify-between gap-3">
<div className="flex flex-wrap items-center gap-3">
{/* Model Name Search */}
<div className="relative w-64">
<input
type="text"
placeholder="Search model names..."
data-testid="model-search-input"
className="w-full px-3 py-2 pl-8 border rounded-md text-sm focus:outline-hidden focus:ring-2 focus:ring-blue-500 focus:border-blue-500"
value={modelNameSearch}
onChange={(e) => setModelNameSearch(e.target.value)}
/>
<svg
className="absolute left-2.5 top-2.5 h-4 w-4 text-gray-500"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M21 21l-6-6m2-5a7 7 0 11-14 0 7 7 0 0114 0z"
/>
</svg>
</div>
{/* Filter Button */}
<button
className={`px-3 py-2 text-sm border rounded-md hover:bg-gray-50 flex items-center gap-2 ${showFilters ? "bg-gray-100" : ""}`}
onClick={() => setShowFilters(!showFilters)}
>
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M3 4a1 1 0 011-1h16a1 1 0 011 1v2.586a1 1 0 01-.293.707l-6.414 6.414a1 1 0 00-.293.707V17l-4 4v-6.586a1 1 0 00-.293-.707L3.293 7.293A1 1 0 013 6.586V4z"
/>
</svg>
Filters
</button>
{/* Reset Filters Button */}
<button
className="px-3 py-2 text-sm border rounded-md hover:bg-gray-50 flex items-center gap-2"
onClick={resetFilters}
>
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24">
<path
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15"
/>
</svg>
Reset Filters
</button>
</div>
{/* Model Settings Button */}
<Button
icon={<SettingOutlined />}
onClick={() => setIsModelSettingsModalVisible(true)}
title="Model Settings"
/>
</div>
{/* Additional Filters */}
{showFilters && (
<div className="flex flex-wrap items-center gap-3 mt-3">
{/* Model Name Filter */}
<div className="w-64">
<Select
className="w-full"
value={selectedModelGroup ?? "all"}
onChange={(value) => setSelectedModelGroup(value === "all" ? "all" : value)}
placeholder="Filter by Public Model Name"
showSearch
options={[
{ value: "all", label: "All Models" },
{ value: "wildcard", label: "Wildcard Models (*)" },
...availableModelGroups.map((group, idx) => ({
value: group,
label: group,
})),
]}
/>
</div>
{/* Model Access Group Filter */}
<div className="w-64">
<Select
className="w-full"
value={selectedModelAccessGroupFilter ?? "all"}
onChange={(value) => setSelectedModelAccessGroupFilter(value === "all" ? null : value)}
placeholder="Filter by Model Access Group"
showSearch
options={[
{ value: "all", label: "All Model Access Groups" },
...availableModelAccessGroups.map((accessGroup, idx) => ({
value: accessGroup,
label: accessGroup,
})),
]}
/>
</div>
</div>
)}
{/* Results Count and Pagination Controls */}
<div className="flex justify-between items-center">
{isLoading ? (
<Skeleton.Input active style={{ width: 184, height: 20 }} />
) : (
<span data-testid="models-results-count" className="text-sm text-gray-700">
{paginationMeta.total_count > 0
? `Showing ${(currentPage - 1) * pageSize + 1} - ${Math.min(currentPage * pageSize, paginationMeta.total_count)} of ${paginationMeta.total_count} results`
: "Showing 0 results"}
</span>
)}
<div className="flex items-center space-x-2">
{isLoading ? (
<Skeleton.Button active style={{ width: 84, height: 30 }} />
) : (
<button
onClick={() => {
const newPage = currentPage - 1;
setCurrentPage(newPage);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}}
disabled={currentPage === 1}
className={`px-3 py-1 text-sm border rounded-md ${
currentPage === 1 ? "bg-gray-100 text-gray-400 cursor-not-allowed" : "hover:bg-gray-50"
}`}
>
Previous
</button>
)}
{isLoading ? (
<Skeleton.Button active style={{ width: 56, height: 30 }} />
) : (
<button
onClick={() => {
const newPage = currentPage + 1;
setCurrentPage(newPage);
setPagination((prev: PaginationState) => ({ ...prev, pageIndex: 0 }));
}}
disabled={currentPage >= paginationMeta.total_pages}
className={`px-3 py-1 text-sm border rounded-md ${
currentPage >= paginationMeta.total_pages
? "bg-gray-100 text-gray-400 cursor-not-allowed"
: "hover:bg-gray-50"
}`}
>
Next
</button>
)}
</div>
</div>
</div>
</div>
<AllModelsDataTable
columns={columns(
userRole,
userId,
premiumUser,
setSelectedModelId,
setSelectedTeamId,
getDisplayModelName,
() => {},
() => {},
expandedRows,
setExpandedRows,
setDeleteModalModelId,
handleTogglePause,
pausingModelId,
)}
data={filteredData}
isLoading={isLoadingModelsInfo}
sorting={sorting}
onSortingChange={setSorting}
pagination={pagination}
onPaginationChange={setPagination}
enablePagination={true}
onRowClick={(model: any) => setSelectedModelId(model.model_info.id)}
/>
{modelViewMode === "current_team" && (
<div className="flex items-start gap-2 px-1 text-xs text-muted-foreground">
<Info className="mt-0.5 size-3.5 shrink-0" />
{selectedTeamValue === PERSONAL_TEAM_VALUE ? (
<span>
To access these models, create a Virtual Key without selecting a team on the{" "}
<a href="/public?login=success&page=api-keys" className="font-medium text-blue-600 hover:underline">
Virtual Keys page
</a>
.
</span>
) : (
<span>
To access these models, create a Virtual Key and select Team as &quot;{teamAccessLabel}&quot; on the{" "}
<a href="/public?login=success&page=api-keys" className="font-medium text-blue-600 hover:underline">
Virtual Keys page
</a>
.
</span>
)}
</div>
</div>
</Grid>
)}
</div>
<DeleteResourceModal
isOpen={!!deleteModalModelId}

Some files were not shown because too many files have changed in this diff Show more