mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
commit
961e52964b
159 changed files with 14351 additions and 8881 deletions
14
.github/workflows/codspeed.yml
vendored
14
.github/workflows/codspeed.yml
vendored
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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("_")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -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(
|
||||
|
|
|
|||
637
litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py
Normal file
637
litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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) ##
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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("/"):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 [
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()}"
|
||||
|
|
|
|||
|
|
@ -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()}"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
58
tests/e2e/mcp/linear_session_capture.py
Normal file
58
tests/e2e/mcp/linear_session_capture.py
Normal 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))))
|
||||
271
tests/e2e/mcp/oauth_chat_client.py
Normal file
271
tests/e2e/mcp/oauth_chat_client.py
Normal 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)
|
||||
197
tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py
Normal file
197
tests/e2e/mcp/test_mcp_chat_completion_oauth_e2e.py
Normal 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}"
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?"}],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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];
|
||||
|
|
|
|||
|
|
@ -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'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's base URL. We'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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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} />,
|
||||
},
|
||||
];
|
||||
|
||||
|
|
|
|||
|
|
@ -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 })));
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
|
|
@ -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 }));
|
||||
});
|
||||
});
|
||||
|
|
@ -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,
|
||||
};
|
||||
};
|
||||
|
|
@ -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/");
|
||||
});
|
||||
});
|
||||
|
|
@ -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 };
|
||||
}
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
{typeof currentTeam !== "string" ? currentTeam.team_alias || currentTeam.team_id : ""}" 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 "{teamAccessLabel}" 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
Loading…
Add table
Reference in a new issue