Merge branch 'BerriAI:litellm_internal_staging' into feature/improve-gigachat-provider

This commit is contained in:
Yuriy 2026-07-08 10:13:29 +03:00 • committed by GitHub
commit 7a73138877
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
126 changed files with 6597 additions and 1031 deletions

View file

@ -1029,6 +1029,8 @@ jobs:
- *python312_image
working_directory: ~/project
resource_class: large
environment:
REQUEST_TIMEOUT: "180"
steps:
- checkout
@ -1058,7 +1060,8 @@ jobs:
-v -x \
--junitxml=test-results/junit.xml \
--durations=5 \
-n 8"
-n 8 \
--reruns 1 --only-rerun Timeout"
no_output_timeout: 15m
# Store test results

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -5,6 +5,10 @@ on:
branches:
- main
- litellm_internal_staging
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
workflow_dispatch:
permissions:

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -7,6 +7,10 @@ on:
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths-ignore:
- "ui/**"
- "**.md"
- "**.mdx"
permissions:
contents: read

View file

@ -16,6 +16,7 @@ jobs:
timeout-minutes: 30
strategy:
fail-fast: false
matrix:
root_path: ["/api/v1", "/llmproxy"]
@ -108,8 +109,26 @@ jobs:
- name: Install UI deps and Chromium
working-directory: ui/litellm-dashboard
run: |
npm ci
npx playwright install --with-deps chromium
retry() {
local attempt=1
local max_attempts=4
until "$@"; do
if [ "$attempt" -ge "$max_attempts" ]; then
echo "Command failed after $attempt attempts: $*"
return 1
fi
echo "Attempt $attempt failed: $*. Retrying in $((attempt * 15))s..."
sleep $((attempt * 15))
attempt=$((attempt + 1))
done
}
npm config set fetch-retries 5
npm config set fetch-retry-mintimeout 20000
npm config set fetch-retry-maxtimeout 120000
retry npm ci
retry npx playwright install --with-deps chromium
- name: Run SERVER_ROOT_PATH redirect e2e
working-directory: ui/litellm-dashboard

View file

@ -21,7 +21,7 @@ End-to-end tests belong in `tests/e2e/` and must follow the harness conventions
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout
When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank

View file

@ -0,0 +1,8 @@
-- Timestamp sorts before some already-applied migrations; this is safe: the
-- runner is `prisma migrate deploy`, which applies every pending migration
-- regardless of name order (utils.py has an informational check for exactly
-- this), and IF NOT EXISTS keeps a re-apply idempotent.
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "token_exchange_endpoint" TEXT;
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "audience" TEXT;
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "subject_token_type" TEXT;

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "token_exchange_profile" TEXT;

View file

@ -329,6 +329,12 @@ model LiteLLM_MCPServerTable {
token_url String?
registration_url String?
oauth2_flow String?
token_exchange_endpoint String?
// Named for the RFC 8693 "audience" token-exchange request parameter (that flow only).
// RFC 8707 resource indicators are a separate concept, named "resource" in the v2 egress types.
audience String?
subject_token_type String?
token_exchange_profile String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)

View file

@ -3619,6 +3619,7 @@ class PrometheusLogger(CustomLogger):
hashed_token=user_api_key_dict.token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_cache_only=True,
)
if key_object:
user_api_key_dict.budget_reset_at = key_object.budget_reset_at

View file

@ -50,6 +50,8 @@ _STS_REGION_FROM_ENDPOINT_PATTERN = re.compile(
r"(?:^|\.)sts(?:-fips)?\.([a-z0-9-]+)\.(?:amazonaws\.com(?:\.cn)?|vpce\.amazonaws\.com)"
)
SIGV4_COMPUTED_HEADERS = frozenset({"authorization", "x-amz-date", "x-amz-security-token", "date"})
class Boto3CredentialsInfo(BaseModel):
credentials: Credentials
@ -1400,11 +1402,13 @@ class BaseAWSLLM:
# Add back all original headers (including forwarded ones) after signature calculation
for header_name, header_value in headers.items():
if header_value is not None:
if header_value is not None and header_name.lower() not in SIGV4_COMPUTED_HEADERS:
request.headers[header_name] = header_value
if (
extra_headers is not None and "Authorization" in extra_headers
extra_headers is not None
and "Authorization" in extra_headers
and not extra_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
): # prevent sigv4 from overwriting the auth header
request.headers["Authorization"] = extra_headers["Authorization"]
prepped = request.prepare()
@ -1527,9 +1531,15 @@ class BaseAWSLLM:
# Add back original headers after signing. Only headers in SignedHeaders
# are integrity-protected; forwarded headers (x-forwarded-*) must remain unsigned.
for header_name, header_value in headers.items():
if header_value is not None:
if header_value is not None and header_name.lower() not in SIGV4_COMPUTED_HEADERS:
request_headers_dict[header_name] = header_value
if headers is not None and "Authorization" in headers: # prevent sigv4 from overwriting the auth header
request_headers_dict["Authorization"] = headers["Authorization"]
incoming_authorization = next(
(value for name, value in headers.items() if name.lower() == "authorization" and value is not None),
None,
)
if incoming_authorization is not None and not incoming_authorization.startswith(
"AWS4-HMAC-SHA256"
): # prevent sigv4 from overwriting the auth header
request_headers_dict["Authorization"] = incoming_authorization
return request_headers_dict, request.body

View file

@ -9,6 +9,7 @@ import json
import os
import threading
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
from urllib.parse import urlparse
import litellm
from litellm._logging import verbose_logger
@ -315,6 +316,9 @@ class VertexBase:
api_base=api_base,
)
if partner == VertexPartnerProvider.llama:
return default_api_base
if len(default_api_base.split(":")) > 1:
endpoint = default_api_base.split(":")[-1]
else:
@ -615,7 +619,8 @@ class VertexBase:
Handles custom api_base for:
1. Gemini (Google AI Studio) - constructs /models/{model}:{endpoint}
2. Vertex AI with standard proxies - constructs {api_base}:{endpoint}
2. Vertex AI with standard proxies - constructs {api_base}:{endpoint};
if api_base has no path (bare host), grafts the default vertex URL path onto it
3. Vertex AI with PSC endpoints - constructs full path structure
{api_base}/v1/projects/{project}/locations/{location}/endpoints/{model}:{endpoint}
(only when use_psc_endpoint_format=True)
@ -660,8 +665,9 @@ class VertexBase:
model_for_url,
endpoint,
)
elif urlparse(api_base).path in ("", "/"):
url = api_base.rstrip("/") + urlparse(url).path
else:
# Fallback to simple format if we don't have all parameters
url = "{}:{}".format(api_base, endpoint)
if stream is True:
url = url + "?alt=sse"

View file

@ -23503,6 +23503,76 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-2.1": {
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image": 5e-06,
"input_cost_per_token": 4e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"regional_processing_uplift_multiplier_eu": 1.1,
"regional_processing_uplift_multiplier_us": 1.1,
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"image",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-2.1-mini": {
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,
"cache_read_input_token_cost": 6e-08,
"input_cost_per_audio_token": 1e-05,
"input_cost_per_image": 8e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"regional_processing_uplift_multiplier_eu": 1.1,
"regional_processing_uplift_multiplier_us": 1.1,
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"image",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-mini": {
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,

View file

@ -83,6 +83,15 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
token_url: Optional[str] = None
registration_url: Optional[str] = None
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
# Token Exchange (OBO) fields — RFC 8693. ``audience`` is named for the RFC's
# request parameter (token-exchange only); RFC 8707 resource indicators are a
# separate concept named ``resource`` in the v2 egress types. A null
# ``subject_token_type`` means DEFAULT_SUBJECT_TOKEN_TYPE (litellm.types.mcp),
# applied at the egress build sites.
token_exchange_endpoint: Optional[str] = None
audience: Optional[str] = None
subject_token_type: Optional[str] = None
token_exchange_profile: Optional[str] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False

View file

@ -28,6 +28,7 @@ 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
@ -35,8 +36,6 @@ if TYPE_CHECKING:
# RFC 8693 grant type constant
TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange"
DEFAULT_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token"
class TokenExchangeHandler:
"""Handles OAuth 2.0 Token Exchange (RFC 8693) for MCP servers.

View file

@ -46,6 +46,35 @@ from litellm.types.mcp import MCPCredentials
if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
_AUTH_FLOW_SCOPED_FIELDS: frozenset = frozenset(
{
"authorization_url",
"token_url",
"registration_url",
"oauth2_flow",
"token_exchange_endpoint",
"audience",
"subject_token_type",
"token_exchange_profile",
}
)
# Token-exchange settings with dedicated columns that also exist on
# ``MCPCredentials`` as a legacy shape (rows and REST callers that predate the
# columns). Every write lifts blob values into the columns and strips them from
# the stored blob, so the read-time ``column or blob`` fallback only serves rows
# the current code has never written — a cleared column can then never be
# silently resurrected by a stale blob copy. These keys are stored plaintext
# (endpoints/identifiers, not secrets), so values lift as-is.
_TOKEN_EXCHANGE_COLUMN_FIELDS: frozenset = frozenset(
{
"token_exchange_endpoint",
"audience",
"subject_token_type",
"token_exchange_profile",
}
)
def _is_global_env_var_scope(scope: Any) -> bool:
"""``scope="user"`` entries are placeholders the user fills in; everything
@ -241,6 +270,14 @@ def _prepare_mcp_server_data(
# Handle credentials serialization
credentials = data_dict.get("credentials")
if credentials is not None:
# Lift legacy blob-shaped token-exchange settings into their dedicated
# columns (an explicit top-level value wins, including an explicit
# null) and strip them from the blob so it never seeds the read-time
# fallback for rows written by current code.
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
blob_value = credentials.pop(te_field, None)
if blob_value is not None and te_field not in data_dict:
data_dict[te_field] = blob_value
data_dict["credentials"] = encrypt_credentials(credentials=credentials, encryption_key=_get_salt_key())
data_dict["credentials"] = safe_dumps(data_dict["credentials"])
@ -603,19 +640,41 @@ async def update_mcp_server(
# Pre-fetch existing record once if we need it for auth_type or credential logic
existing = None
has_credentials = "credentials" in data_dict and data_dict["credentials"] is not None
if data.auth_type or has_credentials:
# An explicit token-exchange column write (set or clear) also migrates the
# legacy blob copies below, so the existing row is needed for those updates.
explicit_te_write = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys())
if data.auth_type or has_credentials or explicit_te_write:
existing = await MCPServerRepository(prisma_client).table.find_unique(where={"server_id": data.server_id})
auth_type_changed = bool(
data.auth_type and existing and existing.auth_type is not None and existing.auth_type != data.auth_type
)
# Clear stale credentials when auth_type changes but no new credentials provided
if (
data.auth_type
and "credentials" not in data_dict
and existing
and existing.auth_type is not None
and existing.auth_type != data.auth_type
):
if auth_type_changed and "credentials" not in data_dict:
data_dict["credentials"] = None
if auth_type_changed:
data_dict.update({field: None for field in _AUTH_FLOW_SCOPED_FIELDS if field not in data_dict})
# An explicit column write that does not touch credentials must still migrate
# the row's legacy blob copies: lift values for columns the caller left
# untouched, strip every copy from the blob. Without this, clearing a column
# (e.g. to re-enable RFC 9728/8414 discovery) would leave the blob copy in
# place, and the next credentials update's migrate-on-write would silently
# repopulate the column the admin just cleared. (When credentials ARE in the
# update, the merge below performs the same migration.)
if explicit_te_write and "credentials" not in data_dict and existing is not None and existing.credentials:
existing_creds = (
json.loads(existing.credentials) if isinstance(existing.credentials, str) else dict(existing.credentials)
)
if _TOKEN_EXCHANGE_COLUMN_FIELDS & existing_creds.keys():
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
legacy_value = existing_creds.pop(te_field, None)
if legacy_value is not None and te_field not in data_dict and getattr(existing, te_field, None) is None:
data_dict[te_field] = legacy_value
data_dict["credentials"] = safe_dumps(existing_creds)
# Merge credentials: preserve existing fields not present in the update.
# Without this, a partial credential update (e.g. changing only region)
# would wipe encrypted secrets that the UI cannot display back.
@ -638,6 +697,19 @@ async def update_mcp_server(
)
# New values override existing; existing keys not in update are preserved
merged = {**existing_creds, **new_creds}
# Migrate-on-write for legacy rows: token-exchange settings the
# old blob shape carried move to their dedicated columns (unless
# the caller set the column this update, or the row already has
# one) and are never re-persisted in the blob. Stored plaintext,
# so the merged value lifts as-is.
for te_field in _TOKEN_EXCHANGE_COLUMN_FIELDS:
legacy_value = merged.pop(te_field, None)
if (
legacy_value is not None
and te_field not in data_dict
and getattr(existing, te_field, None) is None
):
data_dict[te_field] = legacy_value
data_dict["credentials"] = safe_dumps(merged)
# Add audit fields

View file

@ -448,6 +448,12 @@ async def _store_per_user_token_server_side(
)
return # Don't warm Redis if DB write failed
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id)
# Warm the Redis cache so the first subsequent MCP call is a cache hit
ttl = _compute_per_user_token_ttl(server, expires_in)
await mcp_per_user_token_cache.set(

View file

@ -70,6 +70,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import
to_server_spec,
to_subject,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InvalidatableOAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
)
@ -115,7 +118,7 @@ from litellm.proxy.common_utils.user_api_key_cache import get_management_object_
from litellm.proxy.utils import ProxyLogging, get_server_root_path
from litellm.repositories.table_repositories import MCPServerRepository
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPAuth, MCPStdioConfig
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, MCPStdioConfig
from litellm.types.mcp_server.mcp_server_manager import (
MCPInfo,
MCPOAuthMetadata,
@ -689,9 +692,16 @@ class MCPServerManager:
"""
return auth_type == MCPAuth.oauth2_token_exchange and not (token_exchange_endpoint or token_url)
def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None):
def __init__(
self,
cred_provider: Optional[UpstreamCredentialProvider] = None,
per_user_oauth_token_store: Optional[InvalidatableOAuthTokenStore] = None,
):
self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore(
self.get_mcp_server_by_id
)
self._cred_provider = cred_provider or UpstreamCredentialProvider(
oauth_token_store=LazyPerUserOAuthTokenStore(self.get_mcp_server_by_id),
oauth_token_store=self._per_user_oauth_token_store,
token_exchanger=build_token_exchanger(),
)
self.registry: dict[str, MCPServer] = {}
@ -715,8 +725,10 @@ class MCPServerManager:
# Per-server outbound tool-call concurrency limiters, lazily created from
# each server's max_concurrent_requests. Keyed by server_id so the cap
# survives the registry atomic-swap on config reload; a missing key means
# the server has no configured limit.
self._server_call_semaphores: dict[str, asyncio.Semaphore] = {}
# the server has no configured limit. The limit is cached alongside the
# semaphore so an edited limit rebuilds it instead of keeping the old cap
# until restart.
self._server_call_semaphores: dict[str, tuple[int, asyncio.Semaphore]] = {}
self.tool_name_to_mcp_server_name_mapping: dict[str, str] = {}
"""
{
@ -972,7 +984,7 @@ class MCPServerManager:
audience=server_config.get("audience", None),
subject_token_type=server_config.get(
"subject_token_type",
"urn:ietf:params:oauth:token-type:access_token",
DEFAULT_SUBJECT_TOKEN_TYPE,
),
token_exchange_profile=server_config.get("token_exchange_profile", "rfc8693"),
allow_sampling=bool(server_config.get("allow_sampling", False)),
@ -1283,7 +1295,8 @@ class MCPServerManager:
(auth_type == MCPAuth.oauth2 and not mcp_server.authorization_url)
or self._obo_needs_endpoint_discovery(
auth_type,
credentials_dict.get("token_exchange_endpoint") if credentials_dict else None,
mcp_server.token_exchange_endpoint
or (credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
mcp_server.token_url,
)
)
@ -1349,12 +1362,16 @@ class MCPServerManager:
aws_role_name=aws_creds.get("aws_role_name"),
aws_session_name=aws_creds.get("aws_session_name"),
instructions=mcp_server.instructions,
# Token Exchange (OBO) fields — read from credentials JSON blob
token_exchange_endpoint=(credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
audience=(credentials_dict.get("audience") if credentials_dict else None),
subject_token_type=(credentials_dict.get("subject_token_type") if credentials_dict else None)
or "urn:ietf:params:oauth:token-type:access_token",
token_exchange_profile=(credentials_dict.get("token_exchange_profile") if credentials_dict else None)
# Token exchange (OBO) fields: dedicated columns, with the credentials blob as a
# back-compat fallback for servers persisted before the columns existed.
token_exchange_endpoint=mcp_server.token_exchange_endpoint
or (credentials_dict.get("token_exchange_endpoint") if credentials_dict else None),
audience=mcp_server.audience or (credentials_dict.get("audience") if credentials_dict else None),
subject_token_type=mcp_server.subject_token_type
or (credentials_dict.get("subject_token_type") if credentials_dict else None)
or DEFAULT_SUBJECT_TOKEN_TYPE,
token_exchange_profile=mcp_server.token_exchange_profile
or (credentials_dict.get("token_exchange_profile") if credentials_dict else None)
or "rfc8693",
timeout=getattr(mcp_server, "timeout", None),
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
@ -3589,10 +3606,11 @@ class MCPServerManager:
limit = mcp_server.max_concurrent_requests
if limit is None or limit <= 0:
return None
semaphore = self._server_call_semaphores.get(mcp_server.server_id)
if semaphore is None:
semaphore = asyncio.Semaphore(limit)
self._server_call_semaphores[mcp_server.server_id] = semaphore
cached = self._server_call_semaphores.get(mcp_server.server_id)
if cached is not None and cached[0] == limit:
return cached[1]
semaphore = asyncio.Semaphore(limit)
self._server_call_semaphores[mcp_server.server_id] = (limit, semaphore)
return semaphore
@asynccontextmanager
@ -3806,17 +3824,21 @@ class MCPServerManager:
# OBO: the exchanged token may have been revoked/rotated upstream since it was cached, so
# an upstream 401 gets one re-mint + retry. Gated to this mode; all others keep the plain
# single call below.
tool_call_coro = self._obo_call_tool_with_retry(
client=client,
call_tool_params=call_tool_params,
host_progress_callback=host_progress_callback,
mcp_server=mcp_server,
server_auth_header=server_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
async def _obo_call_tool_limited():
async with self._limit_outbound_concurrency(mcp_server):
return await self._obo_call_tool_with_retry(
client=client,
call_tool_params=call_tool_params,
host_progress_callback=host_progress_callback,
mcp_server=mcp_server,
server_auth_header=server_auth_header,
extra_headers=extra_headers,
stdio_env=stdio_env,
subject_token=subject_token,
user_api_key_auth=user_api_key_auth,
)
tool_call_coro = _obo_call_tool_limited()
else:
async def _call_tool_via_client(client, params):
@ -3910,6 +3932,19 @@ class MCPServerManager:
return False
return await self._cred_provider.has_user_token(to_subject(user_api_key_auth, None), spec)
async def invalidate_user_oauth_token_cache(self, user_id: str, server_id: str) -> None:
"""Drop the v2 chain's cached token for ``(user_id, server_id)`` after the credential row
changes (re-auth, revoke), so the next resolve reads the new row instead of serving the
replaced token until its cache TTL. Best-effort: a cache-drop failure is logged, never
raised, because the DB write already succeeded and the TTL remains the backstop.
"""
try:
await self._per_user_oauth_token_store.invalidate(user_id, server_id)
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
verbose_logger.warning(
"Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
)
async def _resolve_oauth2_headers_for_tool_call(
self,
mcp_server: MCPServer,
@ -4630,6 +4665,10 @@ class MCPServerManager:
token_url=server.token_url,
registration_url=server.registration_url,
oauth2_flow=server.oauth2_flow,
token_exchange_endpoint=server.token_exchange_endpoint,
audience=server.audience,
subject_token_type=server.subject_token_type,
token_exchange_profile=server.token_exchange_profile,
allow_all_keys=server.allow_all_keys,
instructions=server.instructions,
timeout=server.timeout,
@ -4734,6 +4773,10 @@ class MCPServerManager:
token_url=server.token_url,
registration_url=server.registration_url,
oauth2_flow=server.oauth2_flow,
token_exchange_endpoint=server.token_exchange_endpoint,
audience=server.audience,
subject_token_type=server.subject_token_type,
token_exchange_profile=server.token_exchange_profile,
allow_all_keys=server.allow_all_keys,
available_on_public_internet=server.available_on_public_internet,
delegate_auth_to_upstream=server.delegate_auth_to_upstream,

View file

@ -28,7 +28,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
Subject,
TokenExchangeConfig,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth
if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
@ -124,7 +124,7 @@ def _token_exchange_spec(server: MCPServer, resource: str) -> Optional[ServerSpe
resource=resource,
config=TokenExchangeConfig(
profile=profile,
subject_token_type=server.subject_token_type or "urn:ietf:params:oauth:token-type:access_token",
subject_token_type=server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
token_exchange_endpoint=endpoint,
audience=server.audience,
client_id=server.client_id,

View file

@ -69,6 +69,17 @@ class OAuthTokenStore(Protocol):
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: ...
class InvalidatableOAuthTokenStore(OAuthTokenStore, Protocol):
"""An ``OAuthTokenStore`` whose cached entry for a ``(user, server)`` pair can be dropped.
The write side calls ``invalidate`` after a (re)authorization or revocation changes the
credential row, so reads stop serving the replaced token immediately instead of until its
cache TTL. ``CachedOAuthTokenStore`` (the top of the per-user chain) satisfies this.
"""
async def invalidate(self, user_id: str, server_id: str) -> None: ...
class TokenRefresher(Protocol):
"""Mints a fresh token from an expired one and persists it, returning the new token.

View file

@ -24,8 +24,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_toke
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
CachedOAuthTokenStore,
InvalidatableOAuthTokenStore,
OAuthToken,
OAuthTokenStore,
RefreshCoordinator,
RefreshingTokenStore,
TokenCacheBackend,
@ -51,7 +51,7 @@ if TYPE_CHECKING:
_DEFAULT_TTL_SECONDS = 300.0
ServerLookup = Callable[[str], "MCPServer | None"]
StoreBuilder = Callable[[ServerLookup], tuple[OAuthTokenStore, bool]]
StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]]
async def _read_credential(user_id: str, server_id: str) -> dict[str, object] | None:
@ -185,7 +185,7 @@ class LazyPerUserOAuthTokenStore:
self._server_lookup = server_lookup
self._store_builder = store_builder
self._redis_available = redis_available
self._store: OAuthTokenStore | None = None
self._store: InvalidatableOAuthTokenStore | None = None
self._uses_redis = False
self._fetch_lock = asyncio.Condition()
self._local_fetches = 0
@ -203,7 +203,26 @@ class LazyPerUserOAuthTokenStore:
if not uses_redis:
await self._finish_local_fetch()
async def _store_for_fetch(self) -> tuple[OAuthTokenStore, bool]:
async def invalidate(self, user_id: str, server_id: str) -> None:
"""Drop the chain's cached entry for ``(user_id, server_id)`` after the credential row
changes (re-auth, revoke). Builds the chain if no fetch has run yet, so a shared (Redis)
cache entry written by another worker is dropped too; the in-process case is then a no-op
on an empty cache.
"""
if self._uses_redis:
store = self._store
if store is not None:
await store.invalidate(user_id, server_id)
return
store, uses_redis = await self._store_for_fetch()
try:
await store.invalidate(user_id, server_id)
finally:
if not uses_redis:
await self._finish_local_fetch()
async def _store_for_fetch(self) -> tuple[InvalidatableOAuthTokenStore, bool]:
async with self._fetch_lock:
while (
self._store is not None and not self._uses_redis and self._redis_available() and self._local_fetches > 0

View file

@ -39,6 +39,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Ok,
Result,
)
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
class AuthSpecKind(str, Enum):
@ -215,7 +216,7 @@ class TokenExchangeConfig(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal[AuthSpecKind.token_exchange] = AuthSpecKind.token_exchange
profile: Literal["rfc8693", "entra_obo"] = "rfc8693"
subject_token_type: str = "urn:ietf:params:oauth:token-type:access_token"
subject_token_type: str = DEFAULT_SUBJECT_TOKEN_TYPE
token_exchange_endpoint: str | None = None
audience: str | None = None
client_id: str | None = None

View file

@ -716,7 +716,7 @@ if MCP_AVAILABLE:
if not (host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta):
return None
host_token = getattr(host_ctx.meta, "progressToken", None)
if not (host_token and hasattr(host_ctx, "session") and host_ctx.session):
if host_token is None or not (hasattr(host_ctx, "session") and host_ctx.session):
return None
host_session = host_ctx.session
@ -732,7 +732,7 @@ if MCP_AVAILABLE:
except Exception as e:
verbose_logger.error(f"Failed to forward progress to Host: {e}")
verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...")
verbose_logger.debug(f"Host progressToken captured: {str(host_token)[:8]}...")
return forward_progress
async def _build_virtual_call_logging_obj(

View file

@ -1046,6 +1046,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
model_rpm_limit: Optional[dict] = None
model_tpm_limit: Optional[dict] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
tag_rpm_limit: Optional[dict[str, int]] = None
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
prompts: Optional[List[str]] = None
@ -1256,6 +1257,14 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
token_url: Optional[str] = None
registration_url: Optional[str] = None
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
# Token Exchange (OBO) fields — RFC 8693. These top-level fields are the
# canonical shape; the same keys inside ``credentials`` are the legacy
# pre-column REST shape and are lifted into these columns on write (an
# explicit top-level value wins) and stripped from the stored blob.
token_exchange_endpoint: Optional[str] = None
audience: Optional[str] = None
subject_token_type: Optional[str] = None
token_exchange_profile: Optional[str] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
@ -1342,6 +1351,14 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
token_url: Optional[str] = None
registration_url: Optional[str] = None
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
# Token Exchange (OBO) fields — RFC 8693. These top-level fields are the
# canonical shape; the same keys inside ``credentials`` are the legacy
# pre-column REST shape and are lifted into these columns on write (an
# explicit top-level value wins) and stripped from the stored blob.
token_exchange_endpoint: Optional[str] = None
audience: Optional[str] = None
subject_token_type: Optional[str] = None
token_exchange_profile: Optional[str] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
@ -3854,6 +3871,7 @@ LiteLLM_ManagementEndpoint_MetadataFields = [
"model_rpm_limit",
"model_tpm_limit",
"mcp_rpm_limit",
"tag_rpm_limit",
"rpm_limit_type",
"tpm_limit_type",
"enforced_params",

View file

@ -1474,7 +1474,7 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c
elif last_db_access_time[key][0] is not None: # check db for non-null values (for refresh operations)
return True
elif last_db_access_time[key][0] is None:
if current_time - last_db_access_time[key] >= db_cache_expiry:
if current_time - last_db_access_time[key][1] >= db_cache_expiry:
return True
return False
@ -1649,6 +1649,12 @@ async def get_user_object(
include={"organization_memberships": True},
)
else:
if should_check_db:
_update_last_db_access_time(
key=db_access_time_key,
value=None,
last_db_access_time=last_db_access_time,
)
raise Exception
if response.organization_memberships is not None and len(response.organization_memberships) > 0:

View file

@ -975,6 +975,20 @@ def get_team_mcp_rpm_limit(
return None
def get_key_tag_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[dict[str, int]]:
"""
Get the per-request-tag rpm limit configured on a given api key.
The returned dict is keyed by request tag, so each tag/group tracked on
the key gets its own independent RPM counter.
"""
if user_api_key_dict.metadata:
return user_api_key_dict.metadata.get("tag_rpm_limit")
return None
def get_project_model_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:

View file

@ -31,8 +31,12 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_str_from_messages,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata
from litellm.proxy.auth.auth_utils import (
get_key_tag_rpm_limit,
get_model_rate_limit_from_metadata,
)
from litellm.proxy.auth.budget_throttle import throttled_limit
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
from litellm.proxy.common_utils.proxy_rate_limit_error import (
ProxyRateLimitError,
map_v3_rate_limit_type,
@ -1300,6 +1304,43 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
)
def _add_tag_per_key_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
data: dict,
descriptors: list[RateLimitDescriptor],
) -> None:
"""
Add per-request-tag rpm limit descriptors for the API key.
Each tag carried on the request that has a configured limit gets its own
``{api_key}:{tag}`` counter, so a burst on one tag/group never consumes
another's budget. Tags without a configured limit fall through to the
key-level descriptor.
"""
if not user_api_key_dict.api_key:
return
tag_rpm_limit = get_key_tag_rpm_limit(user_api_key_dict) or {}
if not tag_rpm_limit:
return
for tag in dict.fromkeys(get_tags_from_request_body(data)):
rpm_limit = tag_rpm_limit.get(tag)
if rpm_limit is None:
continue
descriptors.append(
RateLimitDescriptor(
key="tag_per_key",
value=f"{user_api_key_dict.api_key}:{tag}",
rate_limit={
"requests_per_unit": rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
def _add_mcp_per_key_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -1645,6 +1686,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
descriptors=descriptors,
)
# Per-request-tag rate limits scoped to this key
self._add_tag_per_key_rate_limit_descriptor(
user_api_key_dict=user_api_key_dict,
data=data,
descriptors=descriptors,
)
# REST MCP calls pass the raw body through this hook before server
# resolution; only the later synthetic hook payload may carry this key.
if call_type == CallTypes.call_mcp_tool.value and "server_id" not in data:
@ -1961,6 +2009,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# Org Level Rate Limits
descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model))
# Only check rate limits if we have descriptors with actual limits
if descriptors:
# First pass: RPM and max_parallel_requests sliding-window check.

View file

@ -377,6 +377,7 @@ async def new_user(
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user.
- tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user.
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- agent_id: Optional[str] - The agent id associated with the user.
@ -1379,6 +1380,7 @@ async def user_update(
- budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}.
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user.
- tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user.
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- agent_id: Optional[str] - The agent id associated with the user.

View file

@ -1496,6 +1496,7 @@ async def generate_key_fn(
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tag_rpm_limit: Optional[dict] - key-specific per-request-tag rpm limit, keyed by request tag. Example - {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; requests whose tag is absent fall back to the key-level rpm limit.
- tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request
@ -2514,6 +2515,7 @@ async def update_key_fn(
- rpm_limit: Optional[int] - Requests per minute limit
- model_rpm_limit: Optional[dict] - Model-specific RPM limits {"gpt-4": 100, "claude-v1": 200}
- mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200}
- tag_rpm_limit: Optional[dict] - Per-request-tag RPM limits, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; absent tags fall back to the key-level rpm limit.
- model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000}
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
@ -3551,6 +3553,7 @@ async def generate_key_helper_fn(
model_rpm_limit: Optional[dict] = None,
model_tpm_limit: Optional[dict] = None,
mcp_rpm_limit: Optional[dict] = None,
tag_rpm_limit: Optional[dict] = None,
guardrails: Optional[list] = None,
policies: Optional[list] = None,
prompts: Optional[list] = None,
@ -3624,6 +3627,9 @@ async def generate_key_helper_fn(
if mcp_rpm_limit is not None:
metadata = metadata or {}
metadata["mcp_rpm_limit"] = mcp_rpm_limit
if tag_rpm_limit is not None:
metadata = metadata or {}
metadata["tag_rpm_limit"] = tag_rpm_limit
if guardrails is not None:
metadata = metadata or {}
metadata["guardrails"] = guardrails

View file

@ -537,6 +537,10 @@ if MCP_AVAILABLE:
sanitized.authorization_url = None
sanitized.token_url = None
sanitized.registration_url = None
sanitized.token_exchange_endpoint = None
sanitized.audience = None
sanitized.subject_token_type = None
sanitized.token_exchange_profile = None
# Drop env vars entirely rather than only blanking global values: the
# names alone (DB_PASSWORD, GITHUB_API_KEY, ...) leak what secrets the
# admin configured. Non-admins get the per-user vars they must fill in
@ -578,6 +582,10 @@ if MCP_AVAILABLE:
sanitized.authorization_url = None
sanitized.token_url = None
sanitized.registration_url = None
sanitized.token_exchange_endpoint = None
sanitized.audience = None
sanitized.subject_token_type = None
sanitized.token_exchange_profile = None
sanitized.health_check_error = None
sanitized.last_health_check = None
@ -1905,6 +1913,11 @@ if MCP_AVAILABLE:
expires_in=payload.expires_in,
scopes=payload.scopes,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id)
# Read back the persisted record so the response reflects the stored
# expires_at rather than recomputing it here (which could diverge by
# milliseconds or if the storage logic ever adds a grace period).
@ -1945,6 +1958,11 @@ if MCP_AVAILABLE:
await delete_user_credential(prisma_client, user_id, server_id)
except RecordNotFoundError:
pass # Already gone — treat as a successful delete
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id)
return MCPOAuthUserCredentialStatus(
server_id=server_id,
has_credential=False,

View file

@ -2134,12 +2134,16 @@ async def cli_poll_key(
models=session_data.get("models", []),
)
user_db_obj = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
try:
user_db_obj = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
)
except ValueError as e:
verbose_proxy_logger.debug(f"CLI poll: user lookup failed, proceeding without user budget: {e}")
user_db_obj = None
user_budget = user_db_obj.max_budget if user_db_obj is not None else None
team_budget: Optional[float] = None

View file

@ -7,7 +7,7 @@ import traceback
from base64 import b64encode
from datetime import datetime
from itertools import groupby
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast
from typing import Any, AsyncGenerator, Dict, List, Mapping, Optional, Tuple, Union, cast
from urllib.parse import urlencode, urlparse
import httpx
@ -389,18 +389,24 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
forward_multipart: bool = False,
) -> httpx.Response:
"""
Handle non-streaming HTTP requests
Handle non-SSE HTTP requests
Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests
Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests.
GET and generic requests are sent with httpx stream semantics so the caller can
decide from the response headers whether to buffer the body (JSON, inspected for
logging/guardrails) or relay it to the client without materializing it in memory
(LIT-4009: large batch results files must not be buffered in proxy RSS).
"""
if request.method == "GET":
response = await async_client.request(
method=request.method,
url=url,
get_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
)
elif HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart:
return await async_client.send(get_request, stream=True)
if HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart:
# Forward multipart via make_multipart_http_request even when _parsed_body is
# non-empty (pass_through_request always injects litellm_logging_obj, etc.).
# forward_multipart is False when custom_body was supplied (JSON body despite
@ -412,16 +418,14 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
headers=headers,
requested_query_params=requested_query_params,
)
else:
# Generic httpx method
response = await async_client.request(
method=request.method,
url=url,
headers=headers,
params=requested_query_params,
json=_parsed_body,
)
return response
generic_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
json=_parsed_body,
)
return await async_client.send(generic_request, stream=True)
@staticmethod
def is_multipart(request: Request) -> bool:
@ -820,21 +824,7 @@ async def pass_through_request(
forward_headers=forward_headers,
)
# Apply default query parameters if provided, regardless of merge_query_params setting
if default_query_params or merge_query_params:
# Determine what to merge based on settings
request_params = dict(request.query_params) if merge_query_params else {}
# Create a new URL with the merged query params
url = url.copy_with(
query=urlencode(
HttpPassThroughEndpointHelpers.get_merged_query_parameters(
existing_url=url,
request_query_params=request_params,
default_query_params=default_query_params,
)
).encode("ascii")
)
requested_query_params: Optional[dict] = query_params or dict(request.query_params)
endpoint_type: EndpointType = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url))
@ -952,9 +942,6 @@ async def pass_through_request(
)
logging_obj.model_call_details["litellm_call_id"] = litellm_call_id
# combine url with query params for logging
requested_query_params: Optional[dict] = query_params or dict(request.query_params)
## PASSTHROUGH MANAGED ID RESOLUTION (INPUT) ##
# Resolve managed IDs in path, query params, and body back to raw
# provider IDs before forwarding upstream. Gated by feature flag and
@ -1024,6 +1011,20 @@ async def pass_through_request(
request.method,
)
# Apply default query parameters if provided, regardless of merge_query_params setting
if default_query_params or merge_query_params:
# Create a new URL with the merged query params
url = url.copy_with(
query=urlencode(
HttpPassThroughEndpointHelpers.get_merged_query_parameters(
existing_url=url,
request_query_params=requested_query_params,
default_query_params=default_query_params,
)
).encode("ascii")
)
requested_query_params = None
## PASSTHROUGH MANAGED LIST (DB-only response) ##
# For GET /v1/files and GET /v1/batches passthrough routes, serve the
# listing entirely from our DB so each caller only sees their own IDs.
@ -1164,13 +1165,14 @@ async def pass_through_request(
if state_raw_body is not None:
# SigV4-signed callers (Bedrock) require the exact pre-signed bytes
# to be forwarded so the signature/Content-Length stay valid.
response = await async_client.request(
method=request.method,
url=url,
raw_body_request = async_client.build_request(
request.method,
url,
headers=headers,
params=requested_query_params,
content=state_raw_body,
)
response = await async_client.send(raw_body_request, stream=True)
else:
response = await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
request=request,
@ -1226,6 +1228,40 @@ async def pass_through_request(
status_code=response.status_code,
)
if not _should_buffer_passthrough_response(response):
relay_custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=litellm_call_id,
model_id=None,
cache_key=None,
api_base=str(url._uri_reference),
)
relay_callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=_parsed_body or {},
user_api_key_dict=user_api_key_dict,
response=response,
request_headers=dict(request.headers),
)
if relay_callback_headers:
relay_custom_headers.update(relay_callback_headers)
return StreamingResponse(
_relay_passthrough_response_bytes(
response=response,
request_body=_parsed_body or {},
url_route=str(url),
start_time=start_time,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
success_handler_kwargs=kwargs,
),
status_code=response.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=response.headers,
custom_headers=relay_custom_headers,
),
)
content = await response.aread()
## POST-CALL GUARDRAILS ##
@ -2214,6 +2250,70 @@ def _is_streaming_response(response: httpx.Response) -> bool:
return False
def _should_buffer_passthrough_response(response: httpx.Response) -> bool:
"""
Decide from the response headers whether the body must be read into memory.
JSON bodies (and upstream errors) stay buffered: spend logging, guardrails and
managed-id rewriting inspect them, and they are small in practice. Everything
else (jsonl batch results, octet-stream files, ...) is relayed to the client
chunk by chunk so a large body is never resident in full (LIT-4009). A missing
content-type is buffered because the body cannot be classified.
"""
if response.status_code >= 400:
return True
media_type = response.headers.get("content-type", "").split(";")[0].strip().lower()
return media_type in ("", "application/json") or media_type.endswith("+json")
async def _relay_passthrough_response_bytes(
response: httpx.Response,
request_body: dict,
url_route: str,
start_time: datetime,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str],
success_handler_kwargs: dict,
) -> AsyncGenerator[bytes, None]:
"""
Yield upstream bytes to the client without accumulating them, then fire the
passthrough success handler with response_body=None (uninspected body). The
finally block also runs on client disconnect (GeneratorExit) so partial
downloads still produce a spend-log row, mirroring chunk_processor; a
disconnect additionally logs a warning with the number of bytes relayed so
partial deliveries are distinguishable from complete ones in proxy logs.
"""
bytes_relayed = 0
upstream_fully_relayed = False
try:
async for chunk in response.aiter_bytes():
bytes_relayed += len(chunk)
yield chunk
upstream_fully_relayed = True
finally:
if not upstream_fully_relayed:
verbose_proxy_logger.warning(
f"Passthrough stream for {url_route} ended before upstream body was fully relayed; "
f"{bytes_relayed} bytes were sent to the client"
)
await response.aclose()
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler(
httpx_response=response,
response_body=None,
url_route=url_route,
result="",
start_time=start_time,
end_time=datetime.now(),
logging_obj=logging_obj,
cache_hit=False,
request_body=request_body,
custom_llm_provider=custom_llm_provider,
**success_handler_kwargs,
)
)
def _extract_model_from_vertex_ai_setup(setup_response: dict) -> Optional[str]:
"""
Extract the model name from Vertex AI Live setup response.

View file

@ -34,6 +34,18 @@ from .llm_provider_handlers.vertex_passthrough_logging_handler import (
cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler()
def _safe_response_text(httpx_response: httpx.Response) -> str:
"""
Streamed passthrough responses are relayed to the client without being read
into memory, so accessing .text on them raises ResponseNotRead. Their body is
intentionally uninspected; log an empty string instead of failing the row.
"""
try:
return httpx_response.text
except httpx.ResponseNotRead:
return ""
class PassThroughEndpointLogging:
def __init__(self):
self.TRACKED_VERTEX_ROUTES = [
@ -306,7 +318,9 @@ class PassThroughEndpointLogging:
]
kwargs = normalized_llm_passthrough_logging_payload["kwargs"]
if standard_logging_response_object is None:
standard_logging_response_object = StandardPassThroughResponseObject(response=httpx_response.text)
standard_logging_response_object = StandardPassThroughResponseObject(
response=_safe_response_text(httpx_response)
)
kwargs = self._set_cost_per_request(
logging_obj=logging_obj,

View file

@ -1117,45 +1117,6 @@ _OPENAPI_HTTP_METHODS = {
# `_SSO_SENSITIVE_FIELDS` / `_CACHE_SENSITIVE_FIELDS` constants in the SSO
# and cache endpoint files.
_ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"}
_DB_LITELLM_PARAM_ENV_REF_KEYS = frozenset(
{
"api_key",
"client_secret",
"vertex_credentials",
"vertex_ai_credentials",
"aws_access_key_id",
"aws_secret_access_key",
"aws_session_token",
"aws_region_name",
"aws_session_name",
"aws_profile_name",
"aws_role_name",
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_external_id",
"aws_bedrock_runtime_endpoint",
"aws_bedrock_project_id",
"aws_batch_role_arn",
"aws_workspace_id",
}
)
def _db_model_is_team_scoped(model: object) -> bool:
model_info = getattr(model, "model_info", None)
if isinstance(model_info, BaseModel):
return getattr(model_info, "team_id", None) is not None
if isinstance(model_info, str):
try:
model_info = json.loads(model_info)
except (TypeError, ValueError):
model_info = None
if isinstance(model_info, dict) and model_info.get("team_id") is not None:
return True
if getattr(model_info, "team_id", None) is not None:
return True
model_name = getattr(model, "model_name", None)
return isinstance(model_name, str) and model_name.startswith("model_name_")
def _strip_operation_id_method_suffix(operation_id: str) -> str:
@ -5009,17 +4970,12 @@ class ProxyConfig:
deleted_deployments += 1
return deleted_deployments
def _resolve_db_litellm_param(self, key: str, value: object, resolve_env_refs: bool = True) -> object:
def _resolve_db_litellm_param(self, key: str, value: object) -> object:
if not isinstance(value, str):
return value
decrypted_value = decrypt_value_helper(value=value, key=key, return_original_value=True)
if (
resolve_env_refs
and key in _DB_LITELLM_PARAM_ENV_REF_KEYS
and isinstance(decrypted_value, str)
and decrypted_value.startswith("os.environ/")
):
if isinstance(decrypted_value, str) and decrypted_value.startswith("os.environ/"):
return get_secret(decrypted_value)
return decrypted_value
@ -5040,13 +4996,10 @@ class ProxyConfig:
## ADD MODEL LOGIC
for m in db_models:
_litellm_params = m.litellm_params
resolve_env_refs = not _db_model_is_team_scoped(m)
if isinstance(_litellm_params, dict):
# decrypt values
for k, v in _litellm_params.items():
_litellm_params[k] = self._resolve_db_litellm_param(
key=k, value=v, resolve_env_refs=resolve_env_refs
)
_litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v)
_litellm_params = LiteLLM_Params(**_litellm_params)
else:
@ -5072,15 +5025,12 @@ class ProxyConfig:
_model_list: list = []
for m in new_models:
_litellm_params = m.litellm_params
resolve_env_refs = not _db_model_is_team_scoped(m)
if isinstance(_litellm_params, BaseModel):
_litellm_params = _litellm_params.model_dump()
if isinstance(_litellm_params, dict):
# decrypt values
for k, v in _litellm_params.items():
_litellm_params[k] = self._resolve_db_litellm_param(
key=k, value=v, resolve_env_refs=resolve_env_refs
)
_litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v)
_litellm_params = LiteLLM_Params(**_litellm_params)
else:
verbose_proxy_logger.error(

View file

@ -329,6 +329,12 @@ model LiteLLM_MCPServerTable {
token_url String?
registration_url String?
oauth2_flow String?
token_exchange_endpoint String?
// Named for the RFC 8693 "audience" token-exchange request parameter (that flow only).
// RFC 8707 resource indicators are a separate concept, named "resource" in the v2 egress types.
audience String?
subject_token_type String?
token_exchange_profile String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)

View file

@ -204,6 +204,9 @@ class ResponsesAPIRequestUtils:
if response_id is None:
return responses_api_response
if ResponsesAPIRequestUtils._is_litellm_encoded_response_id(response_id):
return responses_api_response
updated_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
model_id=model_id,
custom_llm_provider=custom_llm_provider,
@ -470,6 +473,14 @@ class ResponsesAPIRequestUtils:
response_id=response_id,
)
@staticmethod
def _is_litellm_encoded_response_id(response_id: str) -> bool:
decoded_response_id = ResponsesAPIRequestUtils._decode_responses_api_response_id(response_id)
return (
decoded_response_id.get("model_id") is not None
or decoded_response_id.get("custom_llm_provider") is not None
)
@staticmethod
def get_model_id_from_response_id(response_id: Optional[str]) -> Optional[str]:
"""Get the model_id from the response_id"""

View file

@ -40,6 +40,12 @@ class MCPAuth(str, enum.Enum):
oauth2_token_exchange = "oauth2_token_exchange"
# RFC 8693 default subject_token_type. A NULL column / omitted config key means
# "use this default"; it is applied at every egress build site via this single
# constant rather than a DB-level DEFAULT (Prisma writes explicit values on
# insert, so a column default would rarely apply anyway).
DEFAULT_SUBJECT_TOKEN_TYPE = "urn:ietf:params:oauth:token-type:access_token"
# MCP Literals
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025]
@ -122,18 +128,31 @@ class MCPCredentials(TypedDict, total=False):
audience: Optional[str]
"""
Target audience for OAuth 2.0 Token Exchange (RFC 8693)
Target audience for OAuth 2.0 Token Exchange (RFC 8693).
Legacy input shape: this setting has a dedicated ``audience`` column, which is
authoritative. A value sent here is accepted for back-compat (the pre-column
REST shape, released since 2026-05), lifted into the column on write, and
stripped from the stored blob. Prefer the top-level request field.
"""
token_exchange_endpoint: Optional[str]
"""
IDP token endpoint for OAuth 2.0 Token Exchange (RFC 8693)
IDP token endpoint for OAuth 2.0 Token Exchange (RFC 8693).
Legacy input shape: lifted into the dedicated ``token_exchange_endpoint``
column on write and stripped from the stored blob; the column is
authoritative. Prefer the top-level request field.
"""
subject_token_type: Optional[str]
"""
Subject token type for OAuth 2.0 Token Exchange (RFC 8693).
Default: urn:ietf:params:oauth:token-type:access_token
Default: DEFAULT_SUBJECT_TOKEN_TYPE (urn:ietf:params:oauth:token-type:access_token).
Legacy input shape: lifted into the dedicated ``subject_token_type`` column on
write and stripped from the stored blob; the column is authoritative. Prefer
the top-level request field.
"""
token_endpoint_auth_method: Optional[MCPTokenEndpointAuthMethod]
@ -147,6 +166,10 @@ class MCPCredentials(TypedDict, total=False):
Token exchange wire dialect: "rfc8693" (default, the standard token-exchange grant) or
"entra_obo" (Microsoft Entra On-Behalf-Of, the RFC 7523 jwt-bearer grant + requested_token_use
extension). Not a secret; stored unencrypted.
Legacy input shape: lifted into the dedicated ``token_exchange_profile`` column on
write and stripped from the stored blob; the column is authoritative. Prefer the
top-level request field.
"""

View file

@ -4,6 +4,7 @@ from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel, ConfigDict
from litellm.types.mcp import (
DEFAULT_SUBJECT_TOKEN_TYPE,
MCPAuth,
MCPAuthType,
MCPTokenEndpointAuthMethod,
@ -68,7 +69,7 @@ class MCPServer(BaseModel):
# Token Exchange (OBO) fields
token_exchange_endpoint: Optional[str] = None
audience: Optional[str] = None
subject_token_type: str = "urn:ietf:params:oauth:token-type:access_token"
subject_token_type: str = DEFAULT_SUBJECT_TOKEN_TYPE
# Wire dialect: "rfc8693" (standard token-exchange grant) or "entra_obo" (Microsoft Entra
# On-Behalf-Of, the RFC 7523 jwt-bearer grant + requested_token_use extension)
token_exchange_profile: str = "rfc8693"

View file

@ -2619,8 +2619,9 @@ _CACHE_PRICING_FIELDS = (
def _resolve_builtin_model_cost_entry(key: str, provider: str) -> Optional[Dict[str, Any]]:
"""Best-effort lookup of a built-in ``model_cost`` entry for a custom key
whose shape ``get_model_info`` cannot resolve (double provider prefixes
like ``bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region aliases).
whose shape ``get_model_info`` cannot resolve (repeated provider prefixes
like ``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region
aliases).
Returns a copy of the matching entry so the caller can inherit its defaults
(most importantly cache pricing) without mutating the shared built-in.
@ -5052,9 +5053,9 @@ def _get_model_info_from_generalization(
candidates = [
potential_model_names["combined_model_name"],
model,
potential_model_names["split_model"],
potential_model_names["combined_stripped_model_name"],
potential_model_names["stripped_model_name"],
potential_model_names["split_model"],
]
for candidate in candidates:
generalized_info = match_fallback_generalization(candidate)
@ -5094,6 +5095,11 @@ def _get_potential_model_names(
stripped_model_name,
)
if custom_llm_provider in ("bedrock", "bedrock_converse"):
from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix
split_model = strip_bedrock_routing_prefix(split_model)
return PotentialModelNamesAndCustomLLMProvider(
split_model=split_model,
combined_model_name=combined_model_name,
@ -5261,9 +5267,9 @@ def _get_model_info_helper(
Check if: (in order of specificity)
1. 'custom_llm_provider/model' in litellm.model_cost. Checks "groq/llama3-8b-8192" if model="llama3-8b-8192" and custom_llm_provider="groq"
2. 'model' in litellm.model_cost. Checks "gemini-1.5-pro-002" in litellm.model_cost if model="gemini-1.5-pro-002" and custom_llm_provider=None
3. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given.
4. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given.
5. 'split_model' in litellm.model_cost. Checks "llama3-8b-8192" in litellm.model_cost if model="groq/llama3-8b-8192"
3. 'split_model' in litellm.model_cost. Checks "au.anthropic.claude-opus-4-8" in litellm.model_cost if model="bedrock/au.anthropic.claude-opus-4-8"
4. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given.
5. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given.
"""
_model_info: Optional[Dict[str, Any]] = None
@ -5289,6 +5295,16 @@ def _get_model_info_helper(
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(split_model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(combined_stripped_model_name)
if _matched_key is not None:
@ -5309,16 +5325,6 @@ def _get_model_info_helper(
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(split_model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info,
custom_llm_provider=model_cost_custom_llm_provider,
):
_model_info = None
if _model_info is None:
generalization = _get_model_info_from_generalization(

View file

@ -23670,6 +23670,76 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-2.1": {
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image": 5e-06,
"input_cost_per_token": 4e-06,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"regional_processing_uplift_multiplier_eu": 1.1,
"regional_processing_uplift_multiplier_us": 1.1,
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"image",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-2.1-mini": {
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,
"cache_read_input_token_cost": 6e-08,
"input_cost_per_audio_token": 1e-05,
"input_cost_per_image": 8e-07,
"input_cost_per_token": 6e-07,
"litellm_provider": "openai",
"max_input_tokens": 128000,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_audio_token": 2e-05,
"output_cost_per_token": 2.4e-06,
"regional_processing_uplift_multiplier_eu": 1.1,
"regional_processing_uplift_multiplier_us": 1.1,
"supported_endpoints": [
"/v1/realtime"
],
"supported_modalities": [
"text",
"image",
"audio"
],
"supported_output_modalities": [
"text",
"audio"
],
"supports_audio_input": true,
"supports_audio_output": true,
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-realtime-mini": {
"cache_creation_input_audio_token_cost": 3e-07,
"cache_read_input_audio_token_cost": 3e-07,

View file

@ -329,6 +329,12 @@ model LiteLLM_MCPServerTable {
token_url String?
registration_url String?
oauth2_flow String?
token_exchange_endpoint String?
// Named for the RFC 8693 "audience" token-exchange request parameter (that flow only).
// RFC 8707 resource indicators are a separate concept, named "resource" in the v2 egress types.
audience String?
subject_token_type String?
token_exchange_profile String?
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)

View file

@ -13,7 +13,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `realtime/` - realtime websocket sessions, including the pipecat audio path
- `budgets/` - budget definition, enforcement, and reset windows (key, team, tag, soft, multi-window)
- `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
- `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials; also the dashboard UI behavior on top of them, driven through the proxy-served UI at /ui with playwright (optional dep behind importorskip)
- `logging/` - logging-integration delivery (datadog and friends)
- `security/` - secret handling and log-leak protection
- `router/` - routing and reliability behavior (rate limits, fallbacks, cooldowns)
@ -63,23 +63,26 @@ The harness is fully typed and new code must not add `Any` or widen the basedpyr
The set of tests we want is a registry checked into this repo, one row per behavior; that file is the definition of done and the denominator. Each e2e test declares what it covers with `@pytest.mark.covers("...")`, and a small collector diffs the registry against the tests and ships coverage to the existing Grafana. No Allure, no new dependencies
Coverage is organized as module > feature > test. There are six modules: LLMs, MCPs, Management/UI, Reliability & Performance, Logging & Guardrails, and Other. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works`
Coverage is organized as module > feature > test. Dashboard modules are Core LLMs, Non-Core LLMs, MCPs, Management/UI, Reliability & Performance, Logging & Guardrails, and Other. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works`
The metric is coverage: the share of registry rows that have a passing covering test, reported to Grafana per module so a gap surfaces as an uncovered row rather than a silent absence
Tests do not declare a dashboard module directly. They only declare the registry cell id with `@pytest.mark.covers("...")`; the registry row decides the module, tier, endpoint, and dashboard rollup. Run `python -m coverage_registry.collector --strict` when you want CI to reject unknown marker ids. Add `--fail-on-collection-errors` when the job should also fail on pytest collection errors.
### Naming grammar per module
LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix
LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix. `chat_completions`, `messages`, and `responses` are Core LLMs. Other LLM endpoints, including `batches` and `realtime`, roll up as Non-Core LLMs.
```
llm.<endpoint>.<route>.<capability>.<streaming>.<assertion>
endpoint : chat_completions | messages | responses | embeddings | batches | files
| rerank | images_generations | audio_speech | audio_transcriptions | moderations
route : openai | azure_openai | anthropic | bedrock_invoke | bedrock_converse | vertex | azure_foundry
| realtime
route : openai | azure_openai | anthropic | bedrock_converse | vertex | azure_foundry
| cohere | together_ai
(vocab varies per endpoint; messages is anthropic-format only)
capability : basic | tool_use | prompt_cache_5m | prompt_cache_1h | vision | thinking
| thinking_tool_use | pdf_input | web_search | structured_output | count_tokens
| tool_search | long_context_1m
capability : basic | tool_use | prompt_cache_5m | vision | thinking | structured_output
| service_tier
streaming : stream | nonstream (omit where n/a)
assertion : works | cost_logged
label (not in id): model = haiku-4.5 | sonnet-4.6 | opus-4.7 | gpt-*

View file

@ -34,6 +34,20 @@ The suites run against a live proxy, so bring one up first. `docker-compose.yml`
uv run pytest tests/e2e/llm_translation/ -v
```
The browser tests in the `management/` suite drive the dashboard the proxy serves at `/ui` through playwright, an optional dependency behind `importorskip` (the suite's API tests run without it). Install it once into your environment along with its browser:
```bash
uv pip install playwright
uv run playwright install chromium
```
They also need a proxy whose bundled UI contains the change under test. The published `main-latest` image ships the UI from the last release; to test local UI changes, build the image from your branch and point the compose stack at it:
```bash
docker build -t litellm-local .
LITELLM_E2E_IMAGE=litellm-local docker compose up -d
```
4. Tear it down when you're done:
```bash

View file

@ -33,6 +33,10 @@ def pytest_configure(config: pytest.Config) -> None:
"markers",
"e2e: live test that requires a running proxy and real provider keys",
)
config.addinivalue_line(
"markers",
"covers(cell_id, *, exercised_on=()): coverage-registry cell(s) this test covers",
)
def _liveness_reason(label: str, base_url: str) -> str | None:

View file

@ -0,0 +1,72 @@
# e2e coverage registry
This directory is the **denominator** for e2e test coverage: the set of behaviors we
want covered, one row per behavior, checked into the repo so coverage is a number we
can track instead of a guess. It implements the plan in the "E2E Coverage Tracking"
note; the naming grammar lives in `tests/e2e/CLAUDE.md`.
## The model
A **cell** is one customer-noticeable behavior a single e2e test can assert pass/fail
on, for example `llm.chat_completions.bedrock_converse.tool_use.stream.works`. Cells are
grouped `module > feature > test`, with LLM cells split into Core LLMs and Non-Core
LLMs for dashboarding. Each cell carries a tier (P0/P1/P2), a source, and a
`fail_before_fix` flag.
The rows live in per-prefix YAML files (`llm_*.yaml`, `mgmt.yaml`, `mcp.yaml`,
`reliability.yaml`, `logging.yaml`, `guardrail.yaml`, `other.yaml`) and validate against
the discriminated union in `schema.py`, so an LLM row cannot carry a guardrail field and
vice versa. `llm` rows with `subject_endpoint` of `chat_completions`, `messages`, or
`responses` roll up to "Core LLMs"; all other LLM endpoints roll up to "Non-Core
LLMs". LLM endpoint, route, and capability values are typed in `schema.py`, so new
taxonomy values require an explicit schema change. `logging` and `guardrail` are two
id-prefixes that roll up into the single "Logging & Guardrails" dashboard module.
A test declares what it covers with a marker:
```python
@pytest.mark.covers("llm.chat_completions.openai.tool_use.stream.works")
def test_openai_streaming_tool_calls(self) -> None:
...
```
## The number
`collector.py` diffs the registry against those markers and reports coverage per module.
It is static: a collect-only pass reads the markers, so it runs no test and needs no live
proxy. Whether a covered cell currently passes or fails is a separate, live concern.
```
cd tests/e2e && PYTHONPATH=. python -m coverage_registry.collector
```
Use `--format prometheus` or `--format json` for CI jobs that publish coverage to
Grafana.
The headline is overall coverage. The collector also lists markers that point at ids
not in the registry, so a typo or an unenumerated behavior surfaces instead of being
silently dropped.
Use strict mode in CI once existing draft markers are reconciled:
```
cd tests/e2e && PYTHONPATH=. python -m coverage_registry.collector --strict
```
Strict mode exits non-zero on `@pytest.mark.covers(...)` ids that are not checked into
the registry. Add `--fail-on-collection-errors` when the job should also fail on pytest
collection errors.
## Status: this is a draft for review
The cells were enumerated from the codebase and the tiers are a first proposal. Known
things to settle before treating the set as final:
- tiers are proposed, not signed off; 125 P0 is a lot to prove fail-before-fix, so P0 may
want tightening
- a few cells need a support check or a prune (for example `llm.embeddings.anthropic.*`
and `reliability.perf.throughput.under_slo`)
- auth is covered in two places (`other.auth.*` and the mgmt authz assertions); the
boundary needs a decision, and the auth cluster may deserve promotion to its own module
- the P2 "niche" cells each stand in for a large tail of integrations/providers by design,
so the denominator is deliberately P0-weighted rather than a full inventory

View file

@ -0,0 +1,8 @@
"""The e2e coverage registry: the denominator for e2e test coverage.
`schema.py` defines one validated row per customer-noticeable behavior (a "cell").
The `*.yaml` files hold the rows, one file per id-prefix. `registry.py` loads and
validates them; `collector.py` diffs the registry against the `@pytest.mark.covers`
markers on the live tests and reports coverage per module. See tests/e2e/CLAUDE.md
for the naming grammar.
"""

View file

@ -0,0 +1,280 @@
"""Diff the registry (denominator) against the @pytest.mark.covers markers on the
live tests (numerator) and report coverage per module.
Coverage here is static: it reads the markers via a collect-only pass, so it runs
no test and needs no live proxy. Whether a covered cell currently passes or fails
(covered_pass vs covered_fail) is a separate, live concern layered on top later.
cd tests/e2e && PYTHONPATH=. python -m coverage_registry.collector
"""
from __future__ import annotations
import contextlib
import io
import json
import sys
from argparse import ArgumentParser
from dataclasses import dataclass
from pathlib import Path
import pytest
from .registry import load_registry
from .schema import MODULE_ORDER, Cell, Tier, dashboard_module
E2E_DIR = Path(__file__).resolve().parent.parent
class _CoversSink:
"""Pytest plugin: after collection, capture every cell id declared via
@pytest.mark.covers(...), plus any nodes that failed to import."""
def __init__(self) -> None:
self.covered_ids: frozenset[str] = frozenset()
self.collection_errors: tuple[str, ...] = ()
def pytest_collection_finish(self, session: pytest.Session) -> None:
self.covered_ids = frozenset(
arg
for item in session.items
for marker in item.iter_markers(name="covers")
for arg in marker.args
if isinstance(arg, str)
)
def pytest_collectreport(self, report: pytest.CollectReport) -> None:
if report.failed:
self.collection_errors = (*self.collection_errors, report.nodeid)
def collect_covered_ids(
e2e_dir: Path = E2E_DIR,
) -> tuple[frozenset[str], tuple[str, ...]]:
"""Return (covered cell ids, nodeids that failed to import)."""
sink = _CoversSink()
with contextlib.redirect_stdout(io.StringIO()):
pytest.main(
[
"--collect-only",
"-qq",
"--continue-on-collection-errors",
"-p",
"no:cacheprovider",
str(e2e_dir),
],
plugins=[sink],
)
return sink.covered_ids, sink.collection_errors
@dataclass(frozen=True, slots=True)
class ModuleCoverage:
module: str
total: int
covered: int
p0_total: int
p0_covered: int
@property
def coverage_percent(self) -> float:
return _percent(self.covered, self.total)
@dataclass(frozen=True, slots=True)
class CoverageReport:
modules: tuple[ModuleCoverage, ...]
total: int
covered: int
p0_total: int
p0_covered: int
p0_gaps: tuple[str, ...]
orphan_markers: tuple[str, ...]
collection_errors: tuple[str, ...]
@property
def coverage_percent(self) -> float:
return _percent(self.covered, self.total)
def _percent(covered: int, total: int) -> float:
return (100.0 * covered / total) if total else 0.0
def _module_coverage(
module: str, cells: tuple[Cell, ...], covered: frozenset[str]
) -> ModuleCoverage:
in_module = tuple(c for c in cells if dashboard_module(c) == module)
p0 = tuple(c for c in in_module if c.tier is Tier.P0)
return ModuleCoverage(
module=module,
total=len(in_module),
covered=sum(1 for c in in_module if c.id in covered),
p0_total=len(p0),
p0_covered=sum(1 for c in p0 if c.id in covered),
)
def compute_coverage(
cells: tuple[Cell, ...],
covered: frozenset[str],
collection_errors: tuple[str, ...] = (),
) -> CoverageReport:
p0_cells = tuple(c for c in cells if c.tier is Tier.P0)
registry_ids = frozenset(c.id for c in cells)
return CoverageReport(
modules=tuple(_module_coverage(m, cells, covered) for m in MODULE_ORDER),
total=len(cells),
covered=sum(1 for c in cells if c.id in covered),
p0_total=len(p0_cells),
p0_covered=sum(1 for c in p0_cells if c.id in covered),
p0_gaps=tuple(sorted(c.id for c in p0_cells if c.id not in covered)),
orphan_markers=tuple(sorted(covered - registry_ids)),
collection_errors=collection_errors,
)
def _row(label: str, covered: int, total: int) -> str:
frac = f"{covered}/{total}"
return f"{label:30}{frac:>12}{_percent(covered, total):>11.1f}%"
def render(report: CoverageReport) -> str:
rows = tuple(_row(m.module, m.covered, m.total) for m in report.modules)
lines = (
f"{'MODULE':30}{'COVERED':>12}{'COVERAGE':>12}",
*rows,
"-" * 54,
_row("ALL", report.covered, report.total),
"",
f"Headline coverage: {report.covered}/{report.total} ({report.coverage_percent:.1f}%)",
)
orphans = (
(
f"\n{len(report.orphan_markers)} marker(s) point at ids not in the registry "
f"(reconcile: fix the marker or add the cell):\n "
+ "\n ".join(report.orphan_markers),
)
if report.orphan_markers
else ()
)
warning = (
(
f"\nWARNING: {len(report.collection_errors)} node(s) failed to import during "
f"collection, so coverage may undercount:\n "
+ "\n ".join(report.collection_errors),
)
if report.collection_errors
else ()
)
return "\n".join((*lines, *orphans, *warning))
def _report_dict(report: CoverageReport) -> dict[str, object]:
return {
"covered": report.covered,
"total": report.total,
"coverage_percent": report.coverage_percent,
"modules": [
{
"module": m.module,
"covered": m.covered,
"total": m.total,
"coverage_percent": m.coverage_percent,
"p0_covered": m.p0_covered,
"p0_total": m.p0_total,
}
for m in report.modules
],
"orphan_markers": list(report.orphan_markers),
"collection_errors": list(report.collection_errors),
}
def render_json(report: CoverageReport) -> str:
return json.dumps(_report_dict(report), indent=2, sort_keys=True)
def _label_value(value: str) -> str:
return value.replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n")
def render_prometheus(report: CoverageReport) -> str:
lines = [
"# HELP litellm_e2e_coverage_cells E2E coverage registry cells by module and state.",
"# TYPE litellm_e2e_coverage_cells gauge",
]
for module in report.modules:
label = _label_value(module.module)
lines.append(
f'litellm_e2e_coverage_cells{{module="{label}",state="covered"}} {module.covered}'
)
lines.append(
f'litellm_e2e_coverage_cells{{module="{label}",state="total"}} {module.total}'
)
lines.extend(
[
f'litellm_e2e_coverage_cells{{module="ALL",state="covered"}} {report.covered}',
f'litellm_e2e_coverage_cells{{module="ALL",state="total"}} {report.total}',
"# HELP litellm_e2e_coverage_percent E2E coverage percent by module.",
"# TYPE litellm_e2e_coverage_percent gauge",
]
)
for module in report.modules:
label = _label_value(module.module)
lines.append(
f'litellm_e2e_coverage_percent{{module="{label}"}} {module.coverage_percent:.6f}'
)
lines.extend(
[
f'litellm_e2e_coverage_percent{{module="ALL"}} {report.coverage_percent:.6f}',
"# HELP litellm_e2e_coverage_orphan_markers Coverage markers not found in the registry.",
"# TYPE litellm_e2e_coverage_orphan_markers gauge",
f"litellm_e2e_coverage_orphan_markers {len(report.orphan_markers)}",
"# HELP litellm_e2e_coverage_collection_errors Pytest nodes that failed during collection.",
"# TYPE litellm_e2e_coverage_collection_errors gauge",
f"litellm_e2e_coverage_collection_errors {len(report.collection_errors)}",
]
)
return "\n".join(lines)
def main() -> int:
parser = ArgumentParser()
parser.add_argument(
"--format",
choices=("text", "json", "prometheus"),
default="text",
help="Output format. Use prometheus or json for Grafana ingestion jobs.",
)
parser.add_argument(
"--strict",
action="store_true",
help="Exit non-zero if markers outside the registry are found.",
)
parser.add_argument(
"--fail-on-collection-errors",
action="store_true",
help="Exit non-zero if pytest collection errors are found.",
)
args = parser.parse_args()
cells = load_registry()
covered, errors = collect_covered_ids()
report = compute_coverage(cells, covered, errors)
output = {
"text": render,
"json": render_json,
"prometheus": render_prometheus,
}[
args.format
](report)
print(output) # noqa: T201 # CLI entrypoint output
if args.strict and report.orphan_markers:
return 1
if args.fail_on_collection_errors and report.collection_errors:
return 1
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -0,0 +1,29 @@
# Guardrail enforcement (behavior features). Grounded in litellm/proxy/guardrails/guardrail_hooks/.
# Rolls up into the "Logging & Guardrails" dashboard module together with logging.*
- {id: guardrail.presidio.pre_call.masks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "PII masking pre-call; data-leak blast radius"}
- {id: guardrail.presidio.post_call.masks, module: guardrail, tier: P0, hook_point: post_call, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Mask PII in model output"}
- {id: guardrail.presidio.logging_only.masks, module: guardrail, tier: P0, hook_point: logging_only, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Redact in logs without blocking"}
- {id: guardrail.bedrock.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "AWS content guardrail blocks harmful input"}
- {id: guardrail.bedrock.during.blocks, module: guardrail, tier: P0, hook_point: during, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "During-call moderation for streaming"}
- {id: guardrail.bedrock.post_call.blocks, module: guardrail, tier: P0, hook_point: post_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "Block harmful output"}
- {id: guardrail.lakera.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/lakera_ai_v2.py", rationale: "Prompt-injection block pre-execution"}
- {id: guardrail.lakera.post_call.blocks, module: guardrail, tier: P0, hook_point: post_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/lakera_ai_v2.py", rationale: "Post-call injection on multi-turn chains"}
- {id: guardrail.openai_moderations.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/openai/moderations.py", rationale: "Content policy for regulated industries"}
- {id: guardrail.aim.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/aim/aim.py", rationale: "Security guardrail malicious-input"}
- {id: guardrail.aim.post_call.blocks, module: guardrail, tier: P1, hook_point: post_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/aim/aim.py", rationale: "Output security check"}
- {id: guardrail.ibm_guardrails.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/ibm_guardrails/ibm_detector.py", rationale: "Enterprise multi-policy"}
- {id: guardrail.ibm_guardrails.post_call.blocks, module: guardrail, tier: P1, hook_point: post_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/ibm_guardrails/ibm_detector.py", rationale: "Output policy validation"}
- {id: guardrail.semantic_guard.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/semantic_guard", rationale: "Semantic policy compliance"}
- {id: guardrail.block_code_execution.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/block_code_execution", rationale: "Code-injection prevention"}
- {id: guardrail.tool_permission.pre_call.allows, module: guardrail, tier: P1, hook_point: pre_call, assertions: [allows], exercised_on: [chat_completions], source: "guardrail_hooks/tool_permission.py", rationale: "Grant allowed tools"}
- {id: guardrail.tool_permission.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/tool_permission.py", rationale: "Block unauthorized tools"}
- {id: guardrail.microsoft_purview.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/microsoft_purview/purview_dlp.py", rationale: "DLP sensitive-data disclosure"}
- {id: guardrail.headroom.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/headroom/headroom.py", rationale: "Anomaly detection threshold"}
- {id: guardrail.generic_guardrail_api.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py", rationale: "Vendor-agnostic custom API"}
- {id: guardrail.pangea.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/pangea/pangea.py", rationale: "API security + DLP"}
- {id: guardrail.niche_providers.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: grammar, rationale: "SMOKE cohort: lasso/hiddenlayer/model_armor/qualifire/guardrails_ai/cato/cisco/akto/prompt_security/promptguard/zscaler/vigil/etc"}
- {id: guardrail.niche_providers.post_call.blocks, module: guardrail, tier: P2, hook_point: post_call, assertions: [blocks], exercised_on: [chat_completions], source: grammar, rationale: "SMOKE niche output filtering"}
- {id: guardrail.niche_providers.pre_call.allows, module: guardrail, tier: P2, hook_point: pre_call, assertions: [allows], exercised_on: [chat_completions], source: grammar, rationale: "SMOKE niche allow-path passthrough"}
- {id: guardrail.tool_policy.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/tool_policy/tool_policy_guardrail.py", rationale: "Tool-use policy enforcement"}
- {id: guardrail.mcp_security.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/mcp_security", rationale: "MCP protocol security"}
- {id: guardrail.llm_as_a_judge.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/llm_as_a_judge", rationale: "LLM-based judgment guardrail"}

View file

@ -0,0 +1,54 @@
# LLM conversational endpoints (chat_completions, messages, responses). Grounded in proxy handlers + model_prices json.
- {id: llm.chat_completions.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "Core endpoint/route/capability"}
- {id: llm.chat_completions.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Core streaming"}
- {id: llm.chat_completions.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "proxy_server.py:8455", rationale: "Cost logging regression catch"}
- {id: llm.chat_completions.openai.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "OpenAI function_calling; high usage"}
- {id: llm.chat_completions.openai.tool_use.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: tool_use, streaming: stream, assertions: [works], source: "model_prices json", rationale: "Tool calls over streaming"}
- {id: llm.chat_completions.openai.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "gpt-4o vision; high usage"}
- {id: llm.chat_completions.openai.prompt_cache_5m.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Prompt caching cost optimization"}
- {id: llm.chat_completions.openai.service_tier.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: service_tier, streaming: nonstream, assertions: [works], source: "OpenAI service_tier param", rationale: "OpenAI scale-tier request option is forwarded and echoed"}
- {id: llm.chat_completions.openai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "o-series reasoning; emerging"}
- {id: llm.chat_completions.openai.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: structured_output, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "response_schema extraction"}
- {id: llm.chat_completions.anthropic.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route translated to Anthropic"}
- {id: llm.chat_completions.anthropic.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming translation"}
- {id: llm.chat_completions.anthropic.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude tool_use; high usage"}
- {id: llm.chat_completions.anthropic.tool_use.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: tool_use, streaming: stream, assertions: [works], source: "model_prices json", rationale: "Streaming tool calls"}
- {id: llm.chat_completions.anthropic.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude vision; high usage"}
- {id: llm.chat_completions.anthropic.prompt_cache_5m.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude prompt caching"}
- {id: llm.chat_completions.anthropic.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: anthropic, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude extended thinking"}
- {id: llm.chat_completions.anthropic.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: anthropic, capability: structured_output, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Claude response_schema"}
- {id: llm.chat_completions.bedrock_converse.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Bedrock Converse unified"}
- {id: llm.chat_completions.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming over Converse"}
- {id: llm.chat_completions.bedrock_converse.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Converse function_calling; AWS adoption"}
- {id: llm.chat_completions.bedrock_converse.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Bedrock vision (Anthropic/Nova)"}
- {id: llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic-on-Bedrock caching"}
- {id: llm.chat_completions.bedrock_converse.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic thinking on Bedrock"}
- {id: llm.chat_completions.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Vertex AI"}
- {id: llm.chat_completions.vertex.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming over Vertex"}
- {id: llm.chat_completions.vertex.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Vertex Gemini function_calling"}
- {id: llm.chat_completions.vertex.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Gemini vision"}
- {id: llm.chat_completions.vertex.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: vertex, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Vertex Gemini prompt caching"}
- {id: llm.chat_completions.azure_openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Azure OpenAI deployments"}
- {id: llm.chat_completions.azure_openai.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Azure OpenAI function_calling"}
- {id: llm.chat_completions.azure_foundry.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: azure_foundry, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "Azure Foundry (azure_ai); newer, smoke"}
- {id: llm.messages.anthropic.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: basic, streaming: nonstream, assertions: [works], source: "anthropic_endpoints/endpoints.py:64", rationale: "Core endpoint; Anthropic Messages native"}
- {id: llm.messages.anthropic.basic.stream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: basic, streaming: stream, assertions: [works], source: "anthropic_endpoints/endpoints.py:64", rationale: "Streaming Messages API"}
- {id: llm.messages.anthropic.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "anthropic_endpoints/endpoints.py:64", rationale: "Cost logged on passthrough"}
- {id: llm.messages.anthropic.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Tool calls via Messages API"}
- {id: llm.messages.anthropic.tool_use.stream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: tool_use, streaming: stream, assertions: [works], source: "model_prices json", rationale: "Streaming tool calls"}
- {id: llm.messages.anthropic.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Vision via Messages API"}
- {id: llm.messages.anthropic.prompt_cache_5m.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Prompt caching via Messages API"}
- {id: llm.messages.anthropic.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Extended thinking via Messages API"}
- {id: llm.responses.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Core endpoint; OpenAI Responses native"}
- {id: llm.responses.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: stream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Streaming via /v1/responses"}
- {id: llm.responses.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "response_api_endpoints/endpoints.py:26", rationale: "Cost logged on responses"}
- {id: llm.responses.openai.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Tool calls via Responses API"}
- {id: llm.responses.openai.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Vision via Responses API"}
- {id: llm.responses.anthropic.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: anthropic, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Anthropic translation (smoke)"}
- {id: llm.responses.anthropic.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: anthropic, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Anthropic"}
- {id: llm.responses.bedrock_converse.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Bedrock Converse (smoke)"}
- {id: llm.responses.bedrock_converse.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: bedrock_converse, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Converse"}
- {id: llm.responses.vertex.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Vertex (smoke)"}
- {id: llm.responses.vertex.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: vertex, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Vertex"}
- {id: llm.responses.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Azure OpenAI (smoke)"}
- {id: llm.responses.azure_openai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Azure OpenAI"}

View file

@ -0,0 +1,45 @@
# LLM non-conversational endpoints. Grounded in litellm/proxy endpoints + llms/ handlers.
- {id: llm.embeddings.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: embeddings, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_embeddings_endpoint_e2e.py:23", rationale: "Core endpoint, live vector response"}
- {id: llm.embeddings.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: embeddings, route: openai, capability: basic, streaming: nonstream, assertions: [cost_logged], source: "SPEND_TRACKING_COVERAGE_MATRIX.md:34", rationale: "Cost tracking on embeddings"}
- {id: llm.embeddings.azure_openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: embeddings, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/azure/azure.py", rationale: "Azure embeddings via translation"}
- {id: llm.embeddings.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: embeddings, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "llms/bedrock/embed/embedding.py", rationale: "Bedrock Titan embeddings"}
- {id: llm.embeddings.vertex.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: embeddings, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "vertex_embeddings/embedding_handler.py", rationale: "Vertex embeddings"}
- {id: llm.embeddings.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: embeddings, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "llms/cohere/embed/handler.py", rationale: "Cohere embeddings"}
- {id: llm.embeddings.anthropic.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: embeddings, route: anthropic, capability: basic, streaming: nonstream, assertions: [works], source: "llms/anthropic/chat/handler.py", rationale: "Anthropic vector API (verify support)"}
- {id: llm.batches.openai.create.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Core batch create"}
- {id: llm.batches.openai.retrieve.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Batch retrieve, id round-trip + status"}
- {id: llm.batches.openai.cancel.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Batch cancel"}
- {id: llm.batches.openai.list.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Batch list envelope"}
- {id: llm.batches.openai.file_lifecycle.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "File upload/retrieve/delete for batch flow"}
- {id: llm.batches.openai_encoded.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py", rationale: "Encoded scenario lifecycle"}
- {id: llm.batches.openai_unified.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py", rationale: "Unified/managed-id scenario"}
- {id: llm.batches.openai_model_param.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py", rationale: "Model-param scenario"}
- {id: llm.batches.openai_provider_fallback.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py", rationale: "Provider-fallback raw-id scenario"}
- {id: llm.batches.azure_openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Azure batches all scenarios"}
- {id: llm.batches.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Vertex batches"}
- {id: llm.batches.bedrock.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Bedrock batches (encoded/unified only)"}
- {id: llm.batches.openai.key_model_access_denied.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Key model restriction 403 on upload/create"}
- {id: llm.files.openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "openai_files_endpoints/files_endpoints.py:46", rationale: "File upload returns OpenAIFileObject"}
- {id: llm.files.openai.retrieve.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "files_endpoints.py", rationale: "File retrieve by id"}
- {id: llm.files.openai.delete.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "files_endpoints.py", rationale: "File delete returns deleted=true"}
- {id: llm.files.openai.list.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "files_endpoints.py", rationale: "File list paginated"}
- {id: llm.files.azure_openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:45", rationale: "Azure file upload managed backend"}
- {id: llm.files.vertex.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:52", rationale: "Vertex file upload to GCS"}
- {id: llm.files.bedrock.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:59", rationale: "Bedrock file upload to S3"}
- {id: llm.rerank.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "test_rerank_e2e.py:29", rationale: "Cohere rerank, top_n + relevance_score"}
- {id: llm.rerank.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "llms/bedrock/rerank/handler.py", rationale: "Bedrock rerank"}
- {id: llm.rerank.together_ai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: together_ai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/together_ai/rerank/handler.py", rationale: "Together rerank"}
- {id: llm.images_generations.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_image_generation_e2e.py:22", rationale: "OpenAI image gen, b64/url"}
- {id: llm.images_generations.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/azure/azure.py", rationale: "Azure DALL-E"}
- {id: llm.images_generations.vertex.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "vertex_ai/image_generation/image_generation_handler.py", rationale: "Vertex Imagen"}
- {id: llm.images_generations.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "bedrock/image_generation/image_handler.py", rationale: "Bedrock Titan Image"}
- {id: llm.images_generations.black_forest_labs.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "black_forest_labs/image_generation/handler.py", rationale: "BFL Flux via OpenAI-compat"}
- {id: llm.audio_speech.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_speech, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_audio_speech_e2e.py:22", rationale: "OpenAI TTS binary audio"}
- {id: llm.audio_speech.openai.basic.stream.works, module: llm, tier: P1, subject_endpoint: audio_speech, route: openai, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:9043", rationale: "TTS streaming chunk generator"}
- {id: llm.audio_speech.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_speech, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/azure/azure.py", rationale: "Azure TTS"}
- {id: llm.audio_speech.vertex.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_speech, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "vertex_ai/text_to_speech/text_to_speech_handler.py", rationale: "Vertex TTS"}
- {id: llm.audio_transcriptions.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_transcriptions, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "openai/transcriptions/handler.py", rationale: "OpenAI Whisper"}
- {id: llm.audio_transcriptions.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_transcriptions, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "azure/audio_transcriptions.py", rationale: "Azure STT"}
- {id: llm.audio_transcriptions.soniox.basic.nonstream.works, module: llm, tier: P2, subject_endpoint: audio_transcriptions, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "soniox/audio_transcription/handler.py", rationale: "Soniox via OpenAI-compat (smoke)"}
- {id: llm.audio_transcriptions.nvidia_riva.basic.nonstream.works, module: llm, tier: P2, subject_endpoint: audio_transcriptions, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "nvidia_riva/audio_transcription/handler.py", rationale: "NVIDIA Riva (smoke)"}
- {id: llm.moderations.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: moderations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py", rationale: "OpenAI moderations (only provider)"}

View file

@ -0,0 +1,25 @@
# Logging integration delivery (behavior features). Grounded in litellm/integrations/.
- {id: logging.langfuse.success.logs_spend, module: logging, tier: P0, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages, embeddings], source: "integrations/langfuse/langfuse.py", rationale: "Primary tracing backend; cost accuracy"}
- {id: logging.langfuse.failure.logs_spend, module: logging, tier: P0, event: failure, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/langfuse/langfuse.py", rationale: "Failure path must still track spend"}
- {id: logging.langfuse.stream.logs_spend, module: logging, tier: P0, event: stream, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/langfuse/langfuse.py", rationale: "Streaming token counts aggregate"}
- {id: logging.s3.success.writes_object, module: logging, tier: P0, event: success, assertions: [writes_object], exercised_on: [chat_completions, messages, embeddings], source: "integrations/s3_v2.py", rationale: "Primary audit trail; batch flush no-drop"}
- {id: logging.s3.failure.writes_object, module: logging, tier: P0, event: failure, assertions: [writes_object], exercised_on: [chat_completions, messages], source: "integrations/s3_v2.py", rationale: "Failed calls persisted for compliance"}
- {id: logging.gcs_bucket.success.writes_object, module: logging, tier: P0, event: success, assertions: [writes_object], exercised_on: [chat_completions, messages, embeddings], source: "integrations/gcs_bucket/gcs_bucket.py", rationale: "GCS parallel to S3"}
- {id: logging.datadog.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/datadog/datadog.py", rationale: "Powers dashboards/alerts; cardinality regressions common"}
- {id: logging.datadog.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions], source: "integrations/datadog/datadog.py", rationale: "Failure metrics for alerting/SLO"}
- {id: logging.prometheus.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/prometheus.py", rationale: "Standard OSS metrics; per-key cardinality (existing e2e)"}
- {id: logging.otel.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/otel/logger.py", rationale: "OTEL spans on every call path"}
- {id: logging.otel.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions, messages], source: "integrations/otel/logger.py", rationale: "Error spans for observability continuity"}
- {id: logging.braintrust.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/braintrust_logging.py", rationale: "Evals platform spend"}
- {id: logging.langsmith.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/langsmith.py", rationale: "LangChain ecosystem"}
- {id: logging.arize.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, embeddings], source: "integrations/arize/arize.py", rationale: "ML-ops observability"}
- {id: logging.mlflow.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions], source: "integrations/mlflow.py", rationale: "Experiment tracking cost/run"}
- {id: logging.opik.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions], source: "integrations/opik/opik.py", rationale: "Eval platform spend/case"}
- {id: logging.openmeter.success.exports_metric, module: logging, tier: P1, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/openmeter.py", rationale: "Usage metering for billing"}
- {id: logging.literal_ai.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions, messages], source: "integrations/literal_ai.py", rationale: "Tracing platform spend"}
- {id: logging.posthog.success.exports_metric, module: logging, tier: P1, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages], source: "integrations/posthog.py", rationale: "Product analytics batching"}
- {id: logging.azure_storage.success.writes_object, module: logging, tier: P1, event: success, assertions: [writes_object], exercised_on: [chat_completions, messages], source: "integrations/azure_storage/azure_storage.py", rationale: "Azure blob for enterprise"}
- {id: logging.cloudzero.success.logs_spend, module: logging, tier: P1, event: success, assertions: [logs_spend], exercised_on: [chat_completions], source: "integrations/cloudzero/cloudzero.py", rationale: "Cost ops correlation"}
- {id: logging.focus.success.writes_object, module: logging, tier: P1, event: success, assertions: [writes_object], exercised_on: [chat_completions, messages], source: "integrations/focus/focus_logger.py", rationale: "Cost mgmt multi-destination export"}
- {id: logging.niche_integrations.success.logs_spend, module: logging, tier: P2, event: success, assertions: [logs_spend], exercised_on: [chat_completions], source: grammar, rationale: "SMOKE cohort: athina/galileo/deepeval/langtrace/weave/lunary/humanloop/traceloop/helicone/argilla/newrelic/sqs/supabase/dynamodb/agentops/lago/etc"}
- {id: logging.niche_integrations.failure.logs_spend, module: logging, tier: P2, event: failure, assertions: [logs_spend], exercised_on: [chat_completions], source: grammar, rationale: "SMOKE niche failure path"}

View file

@ -0,0 +1,113 @@
# MCP module. Grounded in litellm/proxy/_experimental/mcp_server/. See tests/e2e/CLAUDE.md for the grammar.
- id: mcp.list_tools.api_key.succeeds
module: mcp
tier: P0
operation: list_tools
auth_family: api_key
assertions: [succeeds]
source: "server.py:637"
rationale: Core operation; most common auth path; high usage
- id: mcp.list_tools.api_key.denied_without_permission
module: mcp
tier: P0
operation: list_tools
auth_family: api_key
assertions: [denied_without_permission]
source: "mcp_server_manager.py:1409"
rationale: Permission guard is high blast-radius; multi-tenant safety
- id: mcp.call_tool.api_key.succeeds
module: mcp
tier: P0
operation: call_tool
auth_family: api_key
assertions: [succeeds]
source: "server.py:849"
rationale: Primary operation; customer-critical; high usage
- id: mcp.call_tool.api_key.denied_without_permission
module: mcp
tier: P0
operation: call_tool
auth_family: api_key
assertions: [denied_without_permission]
source: "rest_endpoints.py:305-386"
rationale: Tool-level permission guard; multi-tenant safety
- id: mcp.list_tools.bearer.succeeds
module: mcp
tier: P1
operation: list_tools
auth_family: bearer
assertions: [succeeds]
source: "server.py:662"
rationale: OAuth/bearer token flow; upstream delegation
- id: mcp.call_tool.bearer.succeeds
module: mcp
tier: P1
operation: call_tool
auth_family: bearer
assertions: [succeeds]
source: "server.py:886"
rationale: Bearer token forwarding for tool invocation
- id: mcp.list_tools.oauth.succeeds
module: mcp
tier: P1
operation: list_tools
auth_family: oauth
assertions: [succeeds]
source: "rest_endpoints.py:138-188"
rationale: Interactive OAuth2 flow; live token management
- id: mcp.call_tool.oauth.succeeds
module: mcp
tier: P1
operation: call_tool
auth_family: oauth
assertions: [succeeds]
source: "db.py user_oauth_credential lookup"
rationale: OAuth2 token passthrough; per-user credential storage
- id: mcp.list_tools.none.succeeds
module: mcp
tier: P1
operation: list_tools
auth_family: none
assertions: [succeeds]
source: "mcp_server_manager.py:1485-1492"
rationale: Public/anonymous servers; delegate_auth_to_upstream
- id: mcp.call_tool.none.succeeds
module: mcp
tier: P1
operation: call_tool
auth_family: none
assertions: [succeeds]
source: "rest_endpoints.py:305-334"
rationale: No upstream auth required; demo servers
- id: mcp.get_prompt.api_key.succeeds
module: mcp
tier: P1
operation: get_prompt
auth_family: api_key
assertions: [succeeds]
source: "server.py:1042"
rationale: Prompt op; same auth stack as tools
- id: mcp.read_resource.api_key.succeeds
module: mcp
tier: P1
operation: read_resource
auth_family: api_key
assertions: [succeeds]
source: "server.py:1177"
rationale: Resource op; same permission model as tools
- id: mcp.list_prompts.api_key.succeeds
module: mcp
tier: P2
operation: list_prompts
auth_family: api_key
assertions: [succeeds]
source: "server.py:993"
rationale: Smoke-level; same auth stack as list_tools
- id: mcp.list_resources.api_key.succeeds
module: mcp
tier: P2
operation: list_resources
auth_family: api_key
assertions: [succeeds]
source: "server.py:1089"
rationale: Smoke; rarely used; same auth model as tools

View file

@ -0,0 +1,68 @@
# Management/UI endpoint features. Grounded in litellm/proxy/management_endpoints/.
- {id: mgmt.key.generate.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:1444", rationale: "API key survives DB roundtrip"}
- {id: mgmt.key.generate.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "key_management_endpoints.py:1444", rationale: "Only master/team-admin creates keys"}
- {id: mgmt.key.generate.happy_path, module: mgmt, tier: P0, surface: ui, assertions: [happy_path], source: "ui_sso.py:420", rationale: "SSO-driven key gen (UI path)"}
- {id: mgmt.key.update.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:2462", rationale: "Budget/model changes persist"}
- {id: mgmt.key.update.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "key_management_endpoints.py:2462", rationale: "Non-admin cannot escalate perms"}
- {id: mgmt.key.update.happy_path, module: mgmt, tier: P1, surface: ui, assertions: [happy_path], source: "key_management_endpoints.py:2462", rationale: "Key edit through the dashboard"}
- {id: mgmt.key.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:3122", rationale: "Deletion revokes future calls"}
- {id: mgmt.key.delete.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "key_management_endpoints.py:3122", rationale: "Non-owner cannot delete"}
- {id: mgmt.key.info.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:3380", rationale: "Info reflects all writes"}
- {id: mgmt.team.new.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "team_endpoints.py:897", rationale: "team_id/alias/budgets stored"}
- {id: mgmt.team.new.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "team_endpoints.py:897", rationale: "Only org-admin/master creates teams"}
- {id: mgmt.team.member_add.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "team_endpoints.py:2424", rationale: "Membership + per-member budget persist"}
- {id: mgmt.team.member_add.member_forbidden, module: mgmt, tier: P0, surface: api, assertions: [member_forbidden], source: "team_endpoints.py:2424", rationale: "Non-admin forbidden to add"}
- {id: mgmt.team.member_delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "team_endpoints.py:2800", rationale: "Removal revokes team key access"}
- {id: mgmt.team.member_delete.member_forbidden, module: mgmt, tier: P0, surface: api, assertions: [member_forbidden], source: "team_endpoints.py:2800", rationale: "Non-admin forbidden to remove"}
- {id: mgmt.budget.new.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "budget_management_endpoints.py:40", rationale: "max/soft/reset windows persist"}
- {id: mgmt.budget.new.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "budget_management_endpoints.py:40", rationale: "Requires master/admin"}
- {id: mgmt.model.add.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "model_management_endpoints.py:1201", rationale: "Registration persists for routing"}
- {id: mgmt.model.add.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "model_management_endpoints.py:1201", rationale: "Non-admin cannot inject model config"}
- {id: mgmt.user.new.happy_path, module: mgmt, tier: P0, surface: api, assertions: [happy_path], source: "internal_user_endpoints.py:360", rationale: "User creation full cycle"}
- {id: mgmt.key.list.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:5119", rationale: "Key inventory pagination"}
- {id: mgmt.key.block.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "key_management_endpoints.py:5849", rationale: "Blocked stays blocked on restart"}
- {id: mgmt.key.unblock.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "key_management_endpoints.py:5960", rationale: "Unblock restores access"}
- {id: mgmt.key.regenerate.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:6071", rationale: "Rotation: new works, old invalid"}
- {id: mgmt.key.health.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:4292", rationale: "Key health endpoint"}
- {id: mgmt.key.bulk_update.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:2677", rationale: "Batch key updates"}
- {id: mgmt.team.update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py:1582", rationale: "Metadata/budget updates persist"}
- {id: mgmt.team.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py:1750", rationale: "Deletion prevents key access"}
- {id: mgmt.team.block.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py", rationale: "Block suspends all members"}
- {id: mgmt.team.info.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "team_endpoints.py:2244", rationale: "Metadata+members+budgets"}
- {id: mgmt.team.list.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "team_endpoints.py:3645", rationale: "Pagination/filtering"}
- {id: mgmt.team.member_update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "team_endpoints.py:2768", rationale: "Member budget/role updates persist"}
- {id: mgmt.user.update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "internal_user_endpoints.py:555", rationale: "Metadata/perm updates persist"}
- {id: mgmt.user.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "internal_user_endpoints.py:640", rationale: "Deletion revokes keys+teams"}
- {id: mgmt.user.list.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "internal_user_endpoints.py:475", rationale: "Admin view all users"}
- {id: mgmt.user.info.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "internal_user_endpoints.py:440", rationale: "Roles/perms/team membership"}
- {id: mgmt.organization.new.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "organization_endpoints.py:403", rationale: "Org for multi-tenant isolation"}
- {id: mgmt.organization.update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "organization_endpoints.py:545", rationale: "Org metadata updates persist"}
- {id: mgmt.organization.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "organization_endpoints.py:710", rationale: "Cascades to teams/keys"}
- {id: mgmt.organization.member_add.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "organization_endpoints.py:835", rationale: "Org member onboarding"}
- {id: mgmt.customer.new.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "customer_endpoints.py:372", rationale: "End-user for spend tracking"}
- {id: mgmt.customer.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "customer_endpoints.py:480", rationale: "Removes from spend tracking"}
- {id: mgmt.end_user.new.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "customer_endpoints.py:730", rationale: "End-user create (synonym)"}
- {id: mgmt.tag.new.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "tag_management_endpoints.py:160", rationale: "Tag for spend categorization"}
- {id: mgmt.tag.list.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "tag_management_endpoints.py:315", rationale: "Tag enumeration"}
- {id: mgmt.tag.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "tag_management_endpoints.py:390", rationale: "Stops future tagging"}
- {id: mgmt.model.update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "model_management_endpoints.py:1358", rationale: "Pricing/concurrency persist"}
- {id: mgmt.model.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "model_management_endpoints.py:1045", rationale: "Removes from registry"}
- {id: mgmt.model.block.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "model_management_endpoints.py", rationale: "Blocked model stays blocked"}
- {id: mgmt.access_group.new.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "model_access_group_management_endpoints.py:450", rationale: "Model permissioning group"}
- {id: mgmt.access_group.info.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "model_access_group_management_endpoints.py:600", rationale: "Access group membership query"}
- {id: mgmt.mcp_server.register.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "mcp_management_endpoints.py:880", rationale: "MCP server registration"}
- {id: mgmt.mcp_server.approve.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:1200", rationale: "Admin approval persists"}
- {id: mgmt.budget.update.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "budget_management_endpoints.py:155", rationale: "Limit changes apply"}
- {id: mgmt.budget.delete.persists, module: mgmt, tier: P1, surface: api, assertions: [persists], source: "budget_management_endpoints.py:280", rationale: "Clears limits"}
- {id: mgmt.budget.list.happy_path, module: mgmt, tier: P1, surface: api, assertions: [happy_path], source: "budget_management_endpoints.py:215", rationale: "Budget enumeration"}
- {id: mgmt.callback.list.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "callback_management_endpoints.py", rationale: "Callback config (smoke)"}
- {id: mgmt.cache_settings.update.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "cache_settings_endpoints.py", rationale: "Cache config (smoke)"}
- {id: mgmt.cost_tracking.estimate.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "cost_tracking_settings.py", rationale: "Cost estimate (smoke)"}
- {id: mgmt.router_settings.update.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "router_settings_endpoints.py", rationale: "Router config (smoke)"}
- {id: mgmt.jwt_key_mapping.new.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "jwt_key_mapping_endpoints.py", rationale: "JWT->key mapping (smoke)"}
- {id: mgmt.compliance.gdpr.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "compliance_endpoints.py", rationale: "GDPR ops (smoke)"}
- {id: mgmt.tool_management.list.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "tool_management_endpoints.py", rationale: "Tool inventory (smoke)"}
- {id: mgmt.fallback_management.update.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "fallback_management_endpoints.py", rationale: "Fallback config (smoke)"}
- {id: mgmt.config_override.hashicorp_vault.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "config_override_endpoints.py", rationale: "Vault integration (smoke)"}
- {id: mgmt.workflow.list.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "workflow_management_endpoints.py", rationale: "Workflow tracking (smoke)"}
- {id: mgmt.credential_migration.check.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:4252", rationale: "Encryption migration (smoke)"}

View file

@ -0,0 +1,28 @@
# Other (holding pen). Grounded in litellm/proxy/auth/ + health_endpoints/ + proxy_server.py.
# PROMOTION NOTE: the auth cluster (~14 cells) is a candidate to promote to its own module once stable.
- {id: other.auth.master_key.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "user_api_key_auth.py:1569-1588", rationale: "Master key authenticates; timing-safe compare"}
- {id: other.auth.master_key.invalid_denied, module: other, tier: P0, area: auth, assertions: [invalid_denied], source: "user_api_key_auth.py:1580", rationale: "Invalid master key rejected"}
- {id: other.auth.jwt.valid_token_allows, module: other, tier: P0, area: auth, assertions: [valid_token_allows], source: "handle_jwt.py:77-150", rationale: "Valid JWT with correct issuer + claims grants access"}
- {id: other.auth.jwt.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "handle_jwt.py:125-135", rationale: "Expired JWT rejected even with valid signature"}
- {id: other.auth.jwt.invalid_signature_denied, module: other, tier: P0, area: auth, assertions: [invalid_signature_denied], source: "handle_jwt.py:145-150", rationale: "Bad/missing signature fails verification"}
- {id: other.auth.virtual_key.route_permission_enforced, module: other, tier: P0, area: auth, assertions: [route_permission_enforced], source: "route_checks.py:89-151", rationale: "allowed_routes whitelist denies disallowed routes"}
- {id: other.auth.virtual_key.route_group_allowed, module: other, tier: P1, area: auth, assertions: [route_group_allowed], source: "route_checks.py:106-128", rationale: "allowed_routes=[llm_api_routes] grants all LLM endpoints"}
- {id: other.auth.passthrough.model_allowlist_enforced, module: other, tier: P1, area: auth, assertions: [model_allowlist_enforced], source: "route_checks.py:135-151", rationale: "Passthrough enforces per-key model allow-lists"}
- {id: other.auth.oauth2.token_valid_allows, module: other, tier: P1, area: auth, assertions: [token_valid_allows], source: "oauth2_check.py:15-73", rationale: "OAuth2 introspection grants active token"}
- {id: other.auth.oauth2.token_invalid_denied, module: other, tier: P1, area: auth, assertions: [token_invalid_denied], source: "oauth2_check.py:37-73", rationale: "Expired/inactive OAuth2 token denied"}
- {id: other.auth.ip_allowlist.internal_ip_allows, module: other, tier: P1, area: auth, assertions: [internal_ip_allows], source: "ip_address_utils.py:54-76", rationale: "Internal CIDR bypasses public-API restriction"}
- {id: other.auth.ip_allowlist.external_ip_denied_to_private, module: other, tier: P1, area: auth, assertions: [external_ip_denied_to_private], source: "ip_address_utils.py:54-76", rationale: "External IP cannot reach internal-only resources"}
- {id: other.lifecycle.readiness.public_probe, module: other, tier: P0, area: lifecycle, assertions: [public_probe], source: "_health_endpoints.py:1551-1570", rationale: "Unauthenticated /health/readiness safe for LBs"}
- {id: other.lifecycle.readiness.reports_db_status, module: other, tier: P0, area: lifecycle, assertions: [reports_db_status], source: "_health_endpoints.py:1551-1570", rationale: "readiness distinguishes healthy vs DB-unreachable"}
- {id: other.lifecycle.readiness.shutting_down_returns_503, module: other, tier: P0, area: lifecycle, assertions: [shutting_down_returns_503], source: "_health_endpoints.py:1554-1556", rationale: "Graceful shutdown drains LB via 503"}
- {id: other.lifecycle.readiness_details.authenticated_diagnostics, module: other, tier: P1, area: lifecycle, assertions: [authenticated_diagnostics], source: "_health_endpoints.py:1574-1584", rationale: "Auth'd details expose cache/callback status"}
- {id: other.lifecycle.liveness.ping, module: other, tier: P1, area: lifecycle, assertions: [ping], source: "_health_endpoints.py:134-155", rationale: "Liveness confirms server responding"}
- {id: other.lifecycle.startup.config_loads, module: other, tier: P0, area: lifecycle, assertions: [config_loads], source: "proxy_server.py:4020-4100", rationale: "Startup loads YAML, resolves env, persists to DB"}
- {id: other.lifecycle.startup.env_vars_resolved, module: other, tier: P1, area: lifecycle, assertions: [env_vars_resolved], source: "proxy_server.py:3984-4010", rationale: "os.environ/ refs resolved at startup"}
- {id: other.lifecycle.background_health_check.interval_configurable, module: other, tier: P1, area: lifecycle, assertions: [interval_configurable], source: "proxy_server.py:3245-3310", rationale: "Background checks run at configurable interval"}
- {id: other.config.runtime_update.applies_at_runtime, module: other, tier: P0, area: config, assertions: [applies_at_runtime], source: "proxy_server.py:14014-14060", rationale: "/config/update persists to DB + invalidates cache"}
- {id: other.config.general_settings.alert_webhook_side_effect, module: other, tier: P1, area: config, assertions: [alert_webhook_side_effect], source: "proxy_server.py:14215", rationale: "alert_to_webhook_url auto-enables slack alerting"}
- {id: other.config.secret_resolution.kms_integration, module: other, tier: P1, area: config, assertions: [kms_integration], source: "proxy_server.py:3984-4010", rationale: "Resolves secrets from Vault/KMS at startup"}
- {id: other.config.overrides.audit_logged, module: other, tier: P1, area: config, assertions: [audit_logged], source: "config_override_endpoints.py:67-100", rationale: "Config override mutations audit-logged, values redacted"}
- {id: other.key_mgmt.regenerate.grace_period_honored, module: other, tier: P1, area: auth, assertions: [grace_period_honored], source: "key_management_endpoints.py:4503-4560", rationale: "Old key valid during grace_period then revoked"}
- {id: other.key_mgmt.spend_reset.resets_to_value, module: other, tier: P1, area: auth, assertions: [resets_to_value], source: "key_management_endpoints.py:4841", rationale: "reset_spend resets accumulated spend"}

View file

@ -0,0 +1,26 @@
"""Load and validate the registry: the denominator, built in one shot from the YAMLs."""
from __future__ import annotations
from collections import Counter
from pathlib import Path
import yaml
from .schema import CELL_ADAPTER, Cell
REGISTRY_DIR = Path(__file__).resolve().parent
def load_registry(registry_dir: Path = REGISTRY_DIR) -> tuple[Cell, ...]:
"""Every cell across every `*.yaml`, validated. Raises on a schema violation or
a duplicate id, since either would corrupt the coverage denominator."""
cells = tuple(
CELL_ADAPTER.validate_python(row)
for path in sorted(registry_dir.glob("*.yaml"))
for row in (yaml.safe_load(path.read_text()) or ())
)
duplicates = sorted(cid for cid, n in Counter(c.id for c in cells).items() if n > 1)
if duplicates:
raise ValueError(f"duplicate cell ids in registry: {duplicates}")
return cells

View file

@ -0,0 +1,30 @@
# Reliability & Performance (behavior features). Grounded in litellm/router.py + router_strategy/ + router_utils/.
- {id: reliability.fallback.5xx.routes_to_fallback, module: reliability, tier: P0, behavior: fallback, variant: "5xx", assertions: [routes_to_fallback], exercised_on: [chat_completions, messages], source: "litellm/router.py:2024", rationale: "Reroute on provider 5xx to alternate deployment"}
- {id: reliability.fallback.context_window.routes_to_fallback, module: reliability, tier: P0, behavior: fallback, variant: context_window, assertions: [routes_to_fallback], exercised_on: [chat_completions, messages], source: "litellm/router.py:6108", rationale: "Fallback when model exceeds context limit"}
- {id: reliability.fallback.content_policy.routes_to_fallback, module: reliability, tier: P0, behavior: fallback, variant: content_policy, assertions: [routes_to_fallback], exercised_on: [chat_completions, messages], source: "litellm/router.py:6023", rationale: "Reroute on content-policy violation"}
- {id: reliability.fallback.timeout.routes_to_fallback, module: reliability, tier: P0, behavior: fallback, variant: "timeout", assertions: [routes_to_fallback], exercised_on: [chat_completions, messages], source: "litellm/router.py:2766", rationale: "Fallback on request timeout"}
- {id: reliability.retry.5xx.succeeds_within_retries, module: reliability, tier: P0, behavior: retry, variant: "5xx", assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "litellm/router.py:6414", rationale: "Transient 5xx often succeeds on retry"}
- {id: reliability.retry.timeout.succeeds_within_retries, module: reliability, tier: P0, behavior: retry, variant: timeout, assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:44", rationale: "Timeout retried per policy"}
- {id: reliability.retry.429.succeeds_within_retries, module: reliability, tier: P0, behavior: retry, variant: "429", assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:46", rationale: "429 retried per RateLimitErrorRetries policy"}
- {id: reliability.retry.auth.succeeds_within_retries, module: reliability, tier: P1, behavior: retry, variant: auth, assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:42", rationale: "Transient auth glitch retry"}
- {id: reliability.retry.context_window.succeeds_within_retries, module: reliability, tier: P1, behavior: retry, variant: context_window, assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:51", rationale: "Multi-attempt on context error"}
- {id: reliability.cooldown.5xx.trips_then_recovers, module: reliability, tier: P0, behavior: cooldown, variant: "5xx", assertions: [trips_then_recovers], exercised_on: [chat_completions, messages], source: "cooldown_handlers.py:40", rationale: "Deployment cools after repeated 5xx, recovers after cooldown_time"}
- {id: reliability.cooldown.429.trips_then_recovers, module: reliability, tier: P0, behavior: cooldown, variant: "429", assertions: [trips_then_recovers], exercised_on: [chat_completions, messages], source: "cooldown_handlers.py:69", rationale: "Cools on 429, avoids hammering exhausted provider"}
- {id: reliability.cooldown.auth.trips_then_recovers, module: reliability, tier: P1, behavior: cooldown, variant: auth, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages], source: "cooldown_handlers.py:74", rationale: "Cools on 401 auth error"}
- {id: reliability.cooldown.timeout.trips_then_recovers, module: reliability, tier: P1, behavior: cooldown, variant: timeout, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages], source: "cooldown_handlers.py:77", rationale: "Cools on 408 timeout"}
- {id: reliability.ratelimit.rpm.blocks_over_limit, module: reliability, tier: P0, behavior: ratelimit, variant: rpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py", rationale: "v3 limiter enforces RPM per key/team/model; 429 on breach"}
- {id: reliability.ratelimit.tpm.blocks_over_limit, module: reliability, tier: P0, behavior: ratelimit, variant: tpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py", rationale: "v3 limiter enforces TPM per key/team/model; 429 on breach"}
- {id: reliability.ratelimit.priority_generous.picks_under_tpm, module: reliability, tier: P1, behavior: ratelimit, variant: priority_generous, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py:36-52", rationale: "Generous mode (<80% sat) allows priority borrowing"}
- {id: reliability.ratelimit.priority_strict.picks_under_tpm, module: reliability, tier: P1, behavior: ratelimit, variant: priority_strict, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py:53-71", rationale: "Strict mode (>=80% sat) enforces priority fairness"}
- {id: reliability.routing.simple_shuffle.picks_healthy_deployment, module: reliability, tier: P1, behavior: routing, variant: simple_shuffle, assertions: [picks_healthy_deployment], exercised_on: [chat_completions, messages], source: "router_strategy/simple_shuffle.py", rationale: "Baseline weighted/uniform pick"}
- {id: reliability.routing.latency_based.picks_lowest_latency, module: reliability, tier: P1, behavior: routing, variant: latency_based, assertions: [picks_lowest_latency], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_latency.py", rationale: "Routes to lowest-latency deployment"}
- {id: reliability.routing.cost_based.picks_lowest_cost, module: reliability, tier: P1, behavior: routing, variant: cost_based, assertions: [picks_lowest_cost], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_cost.py", rationale: "Spend-aware routing"}
- {id: reliability.routing.usage_based.picks_under_tpm, module: reliability, tier: P0, behavior: routing, variant: usage_based, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_tpm_rpm_v2.py", rationale: "Routes to lowest-TPM deployment; prevents over-allocation"}
- {id: reliability.routing.least_busy.picks_lowest_traffic, module: reliability, tier: P1, behavior: routing, variant: least_busy, assertions: [picks_lowest_traffic], exercised_on: [chat_completions, messages], source: "router_strategy/least_busy.py", rationale: "Fewest in-flight requests"}
- {id: reliability.cache.exact.returns_cached, module: reliability, tier: P1, behavior: cache, variant: exact, assertions: [returns_cached], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/caching.py", rationale: "Response cache returns cached on exact match"}
- {id: reliability.cache.prompt_caching_model_select.returns_cached, module: reliability, tier: P1, behavior: cache, variant: prompt_caching_model_select, assertions: [returns_cached], exercised_on: [chat_completions], source: "router_utils/prompt_caching_cache.py", rationale: "Selects model supporting prompt caching for cacheable prefix"}
- {id: reliability.circuit_breaker.redis.trips_then_recovers, module: reliability, tier: P0, behavior: circuit_breaker, variant: redis, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/redis_cache.py:99", rationale: "Redis breaker CLOSED->OPEN->HALF_OPEN; guards all cache/rate-limit ops"}
- {id: reliability.timeout.request_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: request_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions, messages], source: "litellm/router.py:545-551", rationale: "Per-request timeout raises Timeout"}
- {id: reliability.timeout.stream_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: stream_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions], source: "litellm/router.py:551", rationale: "Streaming chunk-delivery timeout"}
- {id: reliability.perf.latency.under_slo, module: reliability, tier: P1, behavior: perf, variant: latency, assertions: [under_slo], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_latency.py", rationale: "Latency SLO (p50/p99) compliance"}
- {id: reliability.perf.throughput.under_slo, module: reliability, tier: P1, behavior: perf, variant: throughput, assertions: [under_slo], exercised_on: [chat_completions, messages], source: grammar, rationale: "Throughput SLO under load"}

View file

@ -0,0 +1,167 @@
"""Registry row schema: the contract every denominator cell validates against.
A cell is one customer-noticeable behavior a single e2e test can assert pass/fail
on. `module` is the id's segment-1 prefix (seven of them); dashboard rollups can
split or merge those prefixes. The union is discriminated on `module`, so an LLM
row cannot carry a guardrail field and vice versa.
"""
from __future__ import annotations
from enum import Enum
from typing import Annotated, Literal
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
class Tier(str, Enum):
P0 = "P0"
P1 = "P1"
P2 = "P2"
class FailBeforeFix(str, Enum):
proven = "proven"
unproven = "unproven"
LlmEndpoint = Literal[
"chat_completions",
"messages",
"responses",
"embeddings",
"batches",
"files",
"rerank",
"images_generations",
"audio_speech",
"audio_transcriptions",
"moderations",
"realtime",
]
LlmRoute = Literal[
"anthropic",
"azure_foundry",
"azure_openai",
"bedrock_converse",
"cohere",
"openai",
"together_ai",
"vertex",
]
LlmCapability = Literal[
"basic",
"prompt_cache_5m",
"service_tier",
"structured_output",
"thinking",
"tool_use",
"vision",
]
class _Base(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
id: str
tier: Tier
assertions: tuple[str, ...]
source: str
rationale: str = ""
fail_before_fix: FailBeforeFix = FailBeforeFix.unproven
supported: bool = True
class LlmCell(_Base):
module: Literal["llm"]
subject_endpoint: LlmEndpoint
route: LlmRoute
capability: LlmCapability
streaming: Literal["stream", "nonstream", "na"]
class MgmtCell(_Base):
module: Literal["mgmt"]
surface: Literal["api", "ui"]
class McpCell(_Base):
module: Literal["mcp"]
operation: str
auth_family: Literal["none", "api_key", "bearer", "oauth"]
class ReliabilityCell(_Base):
module: Literal["reliability"]
behavior: str
variant: str
exercised_on: tuple[str, ...]
class LoggingCell(_Base):
module: Literal["logging"]
event: str
exercised_on: tuple[str, ...]
class GuardrailCell(_Base):
module: Literal["guardrail"]
hook_point: str
exercised_on: tuple[str, ...]
class OtherCell(_Base):
module: Literal["other"]
area: str
Cell = Annotated[
LlmCell
| MgmtCell
| McpCell
| ReliabilityCell
| LoggingCell
| GuardrailCell
| OtherCell,
Field(discriminator="module"),
]
CELL_ADAPTER: TypeAdapter[Cell] = TypeAdapter(Cell)
CORE_LLM_ENDPOINTS: frozenset[str] = frozenset(
{
"chat_completions",
"messages",
"responses",
}
)
PREFIX_ROLLUP: dict[str, str] = {
"mcp": "MCPs",
"mgmt": "Management/UI",
"reliability": "Reliability & Performance",
"logging": "Logging & Guardrails",
"guardrail": "Logging & Guardrails",
"other": "Other",
}
MODULE_ORDER: tuple[str, ...] = (
"Core LLMs",
"Non-Core LLMs",
"MCPs",
"Management/UI",
"Reliability & Performance",
"Logging & Guardrails",
"Other",
)
def dashboard_module(cell: Cell) -> str:
"""Return the Grafana/reporting module for a registry cell."""
if isinstance(cell, LlmCell):
if cell.subject_endpoint in CORE_LLM_ENDPOINTS:
return "Core LLMs"
return "Non-Core LLMs"
return PREFIX_ROLLUP[cell.module]

View file

@ -0,0 +1,168 @@
"""Tests for the coverage-registry tooling: pure logic plus a registry canary.
No `e2e` marker, so these run without a proxy. They exercise the coverage math and
the registry loader, and guard the checked-in registry against schema drift and
duplicate ids.
"""
from __future__ import annotations
from pathlib import Path
import pytest
from coverage_registry.collector import (
compute_coverage,
render,
render_json,
render_prometheus,
)
from coverage_registry.registry import load_registry
from coverage_registry.schema import (
GuardrailCell,
LlmCell,
LlmEndpoint,
LoggingCell,
Tier,
)
def _llm(
cell_id: str, tier: Tier, subject_endpoint: LlmEndpoint = "chat_completions"
) -> LlmCell:
return LlmCell(
id=cell_id,
module="llm",
tier=tier,
assertions=("works",),
source="test",
subject_endpoint=subject_endpoint,
route="openai",
capability="basic",
streaming="nonstream",
)
def test_compute_coverage_counts_covered_p0_and_gaps() -> None:
cells = (_llm("llm.a", Tier.P0), _llm("llm.b", Tier.P0), _llm("llm.c", Tier.P1))
report = compute_coverage(cells, frozenset({"llm.a"}))
assert (report.total, report.covered) == (3, 1)
assert (report.p0_total, report.p0_covered) == (2, 1)
assert report.p0_gaps == ("llm.b",)
assert report.orphan_markers == ()
def test_orphan_marker_is_reported_not_counted() -> None:
cells = (_llm("llm.a", Tier.P0),)
report = compute_coverage(cells, frozenset({"llm.a", "llm.ghost"}))
assert report.covered == 1
assert report.orphan_markers == ("llm.ghost",)
def test_logging_and_guardrail_roll_up_into_one_module() -> None:
cells = (
LoggingCell(
id="logging.x",
module="logging",
tier=Tier.P0,
assertions=("logs_spend",),
source="t",
event="success",
exercised_on=("chat_completions",),
),
GuardrailCell(
id="guardrail.y",
module="guardrail",
tier=Tier.P1,
assertions=("blocks",),
source="t",
hook_point="pre_call",
exercised_on=("chat_completions",),
),
)
report = compute_coverage(cells, frozenset())
logging_and_guardrails = next(
m for m in report.modules if m.module == "Logging & Guardrails"
)
assert logging_and_guardrails.total == 2
def test_llm_cells_roll_up_by_core_endpoint() -> None:
cells = (
_llm("llm.chat", Tier.P0, "chat_completions"),
_llm("llm.messages", Tier.P0, "messages"),
_llm("llm.responses", Tier.P1, "responses"),
_llm("llm.batches", Tier.P0, "batches"),
_llm("llm.realtime", Tier.P1, "realtime"),
)
report = compute_coverage(cells, frozenset({"llm.chat", "llm.batches"}))
core = next(m for m in report.modules if m.module == "Core LLMs")
non_core = next(m for m in report.modules if m.module == "Non-Core LLMs")
assert (core.total, core.covered, core.p0_total, core.p0_covered) == (3, 1, 2, 1)
assert (
non_core.total,
non_core.covered,
non_core.p0_total,
non_core.p0_covered,
) == (2, 1, 1, 1)
def test_text_render_uses_plain_coverage_language() -> None:
report = compute_coverage(
(_llm("llm.chat", Tier.P0), _llm("llm.batches", Tier.P0, "batches")),
frozenset({"llm.chat"}),
)
text = render(report)
assert "COVERAGE" in text
assert "Headline coverage: 1/2 (50.0%)" in text
assert "P0 COVERED" not in text
def test_json_render_exposes_module_coverage_for_grafana_jobs() -> None:
report = compute_coverage(
(_llm("llm.chat", Tier.P0), _llm("llm.batches", Tier.P0, "batches")),
frozenset({"llm.chat"}),
)
payload = render_json(report)
assert '"coverage_percent": 50.0' in payload
assert '"module": "Core LLMs"' in payload
assert '"module": "Non-Core LLMs"' in payload
def test_prometheus_render_exposes_module_coverage_timeseries() -> None:
report = compute_coverage(
(_llm("llm.chat", Tier.P0), _llm("llm.batches", Tier.P0, "batches")),
frozenset({"llm.chat"}),
)
metrics = render_prometheus(report)
assert 'litellm_e2e_coverage_cells{module="Core LLMs",state="covered"} 1' in metrics
assert 'litellm_e2e_coverage_percent{module="Core LLMs"} 100.000000' in metrics
assert 'litellm_e2e_coverage_percent{module="Non-Core LLMs"} 0.000000' in metrics
assert "litellm_e2e_coverage_orphan_markers 0" in metrics
def test_real_registry_loads_and_ids_are_unique() -> None:
cells = load_registry()
ids = [c.id for c in cells]
assert len(cells) > 250
assert len(ids) == len(set(ids))
assert any(c.id == "logging.prometheus.success.exports_metric" for c in cells)
def test_load_registry_rejects_duplicate_ids(tmp_path: Path) -> None:
row = (
"- {id: llm.dup, module: llm, tier: P0, assertions: [works], source: t, "
"subject_endpoint: chat_completions, route: openai, capability: basic, streaming: nonstream}\n"
)
(tmp_path / "a.yaml").write_text(row)
(tmp_path / "b.yaml").write_text(row)
with pytest.raises(ValueError, match="duplicate cell ids"):
load_registry(tmp_path)

View file

@ -58,6 +58,8 @@ services:
environment:
LITELLM_MASTER_KEY: sk-1234
DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm
UI_USERNAME: admin
UI_PASSWORD: sk-1234
ports:
- "4000:4000"
configs:

View file

@ -21,6 +21,9 @@ CONTROL_PLANE_BASE_URL = os.environ.get(
"LITELLM_CONTROL_PLANE_URL", PROXY_BASE_URL
).rstrip("/")
UI_USERNAME = os.environ.get("E2E_UI_USERNAME", "admin")
UI_PASSWORD = os.environ.get("E2E_UI_PASSWORD", MASTER_KEY)
# Writes on the proxy are eventually consistent (e.g. spend rows flush on
# proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once.
POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120"))

View file

@ -33,7 +33,12 @@ class TestChatCompletionsRegression:
CHAT_MODELS,
ids=[f"{model}-{route}" for model, route in CHAT_MODELS],
)
@pytest.mark.covers("llm.chat_completions.provider.basic.nonstream.works", exercised_on=[])
@pytest.mark.covers(
"llm.chat_completions.openai.basic.nonstream.works",
"llm.chat_completions.anthropic.basic.nonstream.works",
"llm.chat_completions.vertex.basic.nonstream.works",
exercised_on=[],
)
def test_chat_returns_real_completion(
self, client: PassthroughClient, scoped_key: str, model: str, route: str
) -> None:
@ -43,16 +48,23 @@ class TestChatCompletionsRegression:
ChatBody(
model=model,
messages=[
ChatMessage(role="user", content=f"reply with one word {unique_marker()}")
ChatMessage(
role="user",
content=f"reply with one word {unique_marker()}",
)
],
max_tokens=512,
),
)
)
assert response.model, f"{model} ({route}): response carried no model name: {response}"
assert response.choices, f"{model} ({route}): response had no choices: {response}"
assert (
response.model
), f"{model} ({route}): response carried no model name: {response}"
assert (
response.choices
), f"{model} ({route}): response had no choices: {response}"
message = response.choices[0].message
assert message is not None and message.content and message.content.strip(), (
f"{model} ({route}): 200 with an empty completion (#28991): {response}"
)
assert (
message is not None and message.content and message.content.strip()
), f"{model} ({route}): 200 with an empty completion (#28991): {response}"

View file

@ -76,13 +76,18 @@ def post_chat(client: PassthroughClient, key: str, body: BaseModel) -> ChatRespo
class TestServiceTier:
@pytest.mark.covers("llm.chat_completions.openai.service_tier.works", exercised_on=[])
@pytest.mark.covers(
"llm.chat_completions.openai.service_tier.nonstream.works", exercised_on=[]
)
def test_openai_service_tier_is_echoed(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
model = f"e2e-service-tier-{unique_marker()}"
model_id = client.gateway.create_model(
model, LiteLLMParamsBody(model="openai/gpt-5.5", api_key="os.environ/OPENAI_API_KEY")
model,
LiteLLMParamsBody(
model="openai/gpt-5.5", api_key="os.environ/OPENAI_API_KEY"
),
)
resources.defer(lambda: client.gateway.delete_model(model_id))
key = resources.key()
@ -106,7 +111,8 @@ class TestServiceTier:
class TestPromptCaching:
@pytest.mark.covers(
"llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.cache_hit", exercised_on=[]
"llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.works",
exercised_on=[],
)
def test_bedrock_cache_control_produces_cache_read(
self, client: PassthroughClient, resources: ResourceManager
@ -129,7 +135,9 @@ class TestPromptCaching:
RichMessage(
role="user",
content=[
CacheTextBlock(text=cacheable_prefix(), cache_control=CacheControl()),
CacheTextBlock(
text=cacheable_prefix(), cache_control=CacheControl()
),
CacheTextBlock(text="Answer in one word: acknowledged?"),
],
)

View file

@ -1,9 +1,23 @@
"""Management suite client fixture; lifecycle/skip/marker live in the parent conftest."""
"""Management suite fixtures: the client plus a logged-in dashboard page.
Lifecycle/skip/marker live in the parent conftest. The browser fixtures drive
the dashboard the proxy serves at /ui, so browser tests exercise exactly what an
end user sees. playwright is an optional dependency loaded behind importorskip
inside the fixture, so the API tests in this suite collect and run without it:
uv pip install playwright && uv run playwright install chromium
"""
from typing import TYPE_CHECKING, Iterator
import pytest
from e2e_config import PROXY_BASE_URL, UI_PASSWORD, UI_USERNAME
from management_client import ManagementClient, build_client
if TYPE_CHECKING:
from playwright.sync_api import Browser, Page
def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
@ -15,3 +29,29 @@ def pytest_configure(config: pytest.Config) -> None:
@pytest.fixture(scope="session")
def client() -> ManagementClient:
return build_client()
@pytest.fixture(scope="session")
def browser() -> "Iterator[Browser]":
pytest.importorskip("playwright.sync_api", reason="playwright not installed")
from playwright.sync_api import sync_playwright
with sync_playwright() as playwright:
launched = playwright.chromium.launch()
yield launched
launched.close()
@pytest.fixture
def ui_page(browser: "Browser") -> "Iterator[Page]":
context = browser.new_context()
try:
page = context.new_page()
page.goto(f"{PROXY_BASE_URL}/ui/")
page.fill("#username", UI_USERNAME)
page.fill("#password", UI_PASSWORD)
page.click('input[type="submit"]')
page.wait_for_url("**/ui/**")
yield page
finally:
context.close()

View file

@ -0,0 +1,161 @@
"""The dashboard's key create/edit Models dropdown scopes its options to the key's team.
A teamless key offers All Proxy Models but not the all-team-models sentinel (the
backend expands the latter to the full proxy model list when no team is attached),
and a team key offers all-team-models plus the team's own models but never the
all-proxy-models sentinel, even when the team's model list carries it. The create
cases also walk the full product path: submit the modal with the offered sentinel
and read the persisted key back through /key/info.
The tests drive gpt-5.5, one of the example models prewired in the proxy config in
tests/e2e/docker-compose.yml; the dropdown wait fails with a pointer there when the
proxy under test does not serve it.
"""
import pytest
from e2e_config import PROXY_BASE_URL, unique_marker
from lifecycle import ResourceManager
from management_client import ManagementClient
from models import KeyGenerateBody, TeamNewBody
pytest.importorskip("playwright.sync_api", reason="playwright not installed")
from playwright.sync_api import Locator, Page, expect # noqa: E402 # import must follow the importorskip guard above
def _form_item(page: Page, label: str) -> Locator:
return page.locator(".ant-form-item").filter(has=page.get_by_text(label, exact=True)).first
def _open_dropdown(page: Page, label: str) -> Locator:
_form_item(page, label).locator(".ant-select-selector").first.click()
dropdown = page.locator(".ant-select-dropdown:not(.ant-select-dropdown-hidden)").last
expect(dropdown).to_be_visible()
return dropdown
def _models_dropdown_texts(page: Page, must_contain: str) -> list[str]:
dropdown = _open_dropdown(page, "Models")
expect(
dropdown.locator(".ant-select-item-option-content", has_text=must_contain).first,
f"{must_contain!r} never appeared in the Models dropdown; the proxy must serve it "
f"(see the model_list in tests/e2e/docker-compose.yml)",
).to_be_visible()
return dropdown.locator(".ant-select-item-option-content").all_inner_texts()
def _open_create_key_modal(page: Page) -> None:
page.goto(f"{PROXY_BASE_URL}/ui/api-keys/?create=true")
expect(page.locator(".ant-modal").first).to_be_visible()
def _select_team(page: Page, alias: str) -> None:
dropdown = _open_dropdown(page, "Team")
dropdown.get_by_text(alias).first.click()
def _submit_create_modal(page: Page, sentinel_label: str) -> str:
dropdown = page.locator(".ant-select-dropdown:not(.ant-select-dropdown-hidden)").last
dropdown.locator(".ant-select-item-option-content", has_text=sentinel_label).first.click()
page.keyboard.press("Escape")
_form_item(page, "Key Name").locator("input").first.fill(f"e2e-ui-key-{unique_marker()}")
page.get_by_role("button", name="Create Key", exact=True).click()
expect(page.get_by_text("Save your Key")).to_be_visible()
key = page.locator(".ant-modal pre").last.inner_text().strip()
assert key.startswith("sk-"), f"expected the created key in the success modal, got {key!r}"
return key
def _open_key_edit_form(page: Page, key_alias: str) -> None:
page.goto(f"{PROXY_BASE_URL}/ui/api-keys/")
page.get_by_text(key_alias).first.click()
page.get_by_role("tab", name="Settings").click()
page.get_by_role("button", name="Edit Settings").click()
expect(_form_item(page, "Models")).to_be_visible()
def _provision_team(client: ManagementClient, resources: ResourceManager, alias: str) -> str:
team_id = client.create_team(TeamNewBody(team_alias=alias, models=["all-proxy-models", "gpt-5.5"]))
resources.defer(lambda: client.delete_team(team_id))
return team_id
def _provision_key(
client: ManagementClient, resources: ResourceManager, alias: str, team_id: str | None = None
) -> str:
key = client.gateway.generate_key(KeyGenerateBody(key_alias=alias, models=["gpt-5.5"], team_id=team_id))
resources.defer(lambda: client.gateway.delete_key(key))
return key
@pytest.mark.e2e
class TestKeyModelsDropdownUI:
@pytest.mark.covers("mgmt.key.generate.happy_path", exercised_on=[])
def test_create_teamless_key_offers_proxy_scope_and_persists(
self, ui_page: Page, client: ManagementClient, resources: ResourceManager
) -> None:
_open_create_key_modal(ui_page)
options = _models_dropdown_texts(ui_page, must_contain="gpt-5.5")
assert "All Proxy Models" in options, f"teamless create lost 'All Proxy Models': {options}"
assert "All Team Models" not in options, f"teamless create offered 'All Team Models': {options}"
key = _submit_create_modal(ui_page, sentinel_label="All Proxy Models")
resources.defer(lambda: client.gateway.delete_key(key))
info = client.gateway.key_info(key)
assert info.models == ["all-proxy-models"], f"persisted models {info.models}"
assert info.team_id is None, f"teamless key persisted with team {info.team_id}"
@pytest.mark.covers("mgmt.key.generate.happy_path", exercised_on=[])
def test_create_team_key_offers_team_scope_and_persists(
self, ui_page: Page, client: ManagementClient, resources: ResourceManager
) -> None:
team_alias = f"e2e-ui-team-{unique_marker()}"
team_id = _provision_team(client, resources, team_alias)
_open_create_key_modal(ui_page)
_select_team(ui_page, team_alias)
options = _models_dropdown_texts(ui_page, must_contain="All Team Models")
assert "gpt-5.5" in options, f"team key create lost the team's own model: {options}"
assert "All Proxy Models" not in options, f"team key create offered 'All Proxy Models': {options}"
assert "all-proxy-models" not in options, f"team key create offered the raw sentinel: {options}"
key = _submit_create_modal(ui_page, sentinel_label="All Team Models")
resources.defer(lambda: client.gateway.delete_key(key))
info = client.gateway.key_info(key)
assert info.models == ["all-team-models"], f"persisted models {info.models}"
assert info.team_id == team_id, f"persisted team {info.team_id}, expected {team_id}"
@pytest.mark.covers("mgmt.key.update.happy_path", exercised_on=[])
def test_edit_teamless_key_offers_proxy_scope(
self, ui_page: Page, client: ManagementClient, resources: ResourceManager
) -> None:
key_alias = f"e2e-ui-teamless-{unique_marker()}"
_provision_key(client, resources, key_alias)
_open_key_edit_form(ui_page, key_alias)
options = _models_dropdown_texts(ui_page, must_contain="gpt-5.5")
assert "All Proxy Models" in options, f"teamless edit lost 'All Proxy Models': {options}"
assert "All Team Models" not in options, f"teamless edit offered 'All Team Models': {options}"
@pytest.mark.covers("mgmt.key.update.happy_path", exercised_on=[])
def test_edit_team_key_offers_team_scope_only(
self, ui_page: Page, client: ManagementClient, resources: ResourceManager
) -> None:
team_alias = f"e2e-ui-team-{unique_marker()}"
team_id = _provision_team(client, resources, team_alias)
key_alias = f"e2e-ui-teamkey-{unique_marker()}"
_provision_key(client, resources, key_alias, team_id=team_id)
_open_key_edit_form(ui_page, key_alias)
options = _models_dropdown_texts(ui_page, must_contain="All Team Models")
assert "gpt-5.5" in options, f"team key edit lost the team's own model: {options}"
assert "All Proxy Models" not in options, f"team key edit offered 'All Proxy Models': {options}"
assert "all-proxy-models" not in options, f"team key edit offered the raw sentinel: {options}"

View file

@ -167,7 +167,7 @@ class TestKeyRoutes:
class TestTeamRoutes:
@pytest.mark.covers("management.team.new.persists")
@pytest.mark.covers("mgmt.team.new.persists")
def test_new_persists_to_team_info_and_binds_keys(
self, client: ManagementClient, resources: ResourceManager
) -> None:
@ -212,7 +212,7 @@ class TestTeamRoutes:
class TestUserRoutes:
@pytest.mark.covers("mgmt.user.new.persists")
@pytest.mark.covers("mgmt.user.new.happy_path")
def test_new_persists_to_user_info(self, client: ManagementClient, resources: ResourceManager) -> None:
email = f"e2e-mgmt-{unique_marker()}@example.com"
user_id = _create_user(client, resources, UserNewBody(user_email=email, user_role="internal_user"))
@ -225,7 +225,7 @@ class TestUserRoutes:
class TestOrganizationRoutes:
@pytest.mark.covers("mgmt.organization.new.persists")
@pytest.mark.covers("mgmt.organization.new.happy_path")
def test_new_persists_to_organization_info(
self, client: ManagementClient, resources: ResourceManager
) -> None:
@ -252,7 +252,7 @@ def _assert_route_forbidden(route: str, outcome: StreamingResponse) -> None:
class TestManagementRoutePermissions:
@pytest.mark.covers("mgmt.key.generate.member_forbidden")
@pytest.mark.covers("other.auth.virtual_key.route_permission_enforced")
def test_llm_only_key_forbidden_from_management_writes(
self, client: ManagementClient, resources: ResourceManager
) -> None:

View file

@ -765,7 +765,10 @@ class BaseResponsesAPITest(ABC):
max_output_tokens=256,
tools=tools,
tool_choice="auto",
timeout=90,
)
except litellm.Timeout:
pytest.skip("Provider did not answer the shell tool request within 90s")
except litellm.InternalServerError:
pytest.skip("Skipping test due to litellm.InternalServerError")
except litellm.BadRequestError as e:

View file

@ -4,7 +4,8 @@ Integration tests for RealTimeStreaming guardrails against a live OpenAI backend
These tests require OPENAI_API_KEY and are skipped if not set.
They verify end-to-end that:
1. A text message blocked by a guardrail -> error event sent to client, NO AI response.
1. A text message blocked by a guardrail -> error event sent to client, the blocked
message never reaches OpenAI, and the client's response.create is not forwarded.
2. A voice transcript blocked by a guardrail -> error event sent, response.create NOT sent.
3. A clean text message passes through and triggers a real OpenAI response.
@ -55,9 +56,25 @@ def _make_guardrail(event_hook=GuardrailEventHooks.pre_call):
)
async def _wait_for_event(
client_events: List[dict], event_type: str, timeout: float = 15.0
) -> dict:
class RecordingBackendWebSocket:
"""Wraps a real backend WebSocket and records every frame sent to it."""
def __init__(self, backend_ws):
self._backend_ws = backend_ws
self.sent_messages: List[str] = []
async def send(self, message):
self.sent_messages.append(message)
await self._backend_ws.send(message)
async def recv(self, *args, **kwargs):
return await self._backend_ws.recv(*args, **kwargs)
async def close(self):
await self._backend_ws.close()
async def _wait_for_event(client_events: List[dict], event_type: str, timeout: float = 15.0) -> dict:
"""Poll client_events list until an event with matching type appears."""
deadline = asyncio.get_event_loop().time() + timeout
while asyncio.get_event_loop().time() < deadline:
@ -65,9 +82,7 @@ async def _wait_for_event(
if matching:
return matching[0]
await asyncio.sleep(0.05)
raise TimeoutError(
f"Timed out waiting for '{event_type}'. Got so far: {[e.get('type') for e in client_events]}"
)
raise TimeoutError(f"Timed out waiting for '{event_type}'. Got so far: {[e.get('type') for e in client_events]}")
async def _build_streaming(client_events: List[dict], backend_ws, request_data=None):
@ -99,12 +114,21 @@ async def _build_streaming(client_events: List[dict], backend_ws, request_data=N
@pytest.mark.asyncio
async def test_text_message_blocked_by_guardrail_no_ai_response():
"""
Send a text message containing the blocked phrase.
Send a text message containing the blocked phrase, immediately followed by
response.create (the reflexive client pattern).
Guardrail must:
- Send error event (guardrail_violation) to client.
- Send response.output_audio_transcript.delta (or beta-protocol
response.audio_transcript.delta) with the block message to client.
- NOT forward response.create to OpenAI (no AI response).
response.audio_transcript.delta) to client.
- NEVER forward the blocked message to OpenAI.
- Drop the client's response.create; the only response.create OpenAI sees
is the guardrail's own (which voices the block message), so the model
can never answer the blocked content.
Assertions are on the recorded backend wire traffic, not on the model's
reply wording: gpt-realtime phrases its voicing/refusal of the guardrail
prompt nondeterministically, which made wording-based assertions flaky
(see PRs #28191, #28200, #29477).
"""
import websockets
@ -119,21 +143,16 @@ async def test_text_message_blocked_by_guardrail_no_ai_response():
additional_headers={
"Authorization": f"Bearer {OPENAI_API_KEY}",
},
) as backend_ws:
) as raw_backend_ws:
backend_ws = RecordingBackendWebSocket(raw_backend_ws)
streaming, input_queue = await _build_streaming(client_events, backend_ws)
# Start backend -> client forwarding
backend_task = asyncio.create_task(
streaming.backend_to_client_send_messages()
)
# Start client -> backend forwarding (reads from input_queue)
backend_task = asyncio.create_task(streaming.backend_to_client_send_messages())
client_task = asyncio.create_task(streaming.client_ack_messages())
try:
# Wait until session is ready
await _wait_for_event(client_events, "session.created", timeout=15)
# Send the blocked message + response.create
blocked_item = json.dumps(
{
"type": "conversation.item.create",
@ -149,34 +168,23 @@ async def test_text_message_blocked_by_guardrail_no_ai_response():
}
)
await input_queue.put(blocked_item)
# Give guardrail time to process before the follow-up response.create
await asyncio.sleep(0.3)
await input_queue.put(json.dumps({"type": "response.create"}))
# Allow time for guardrail round-trip
await asyncio.sleep(3.0)
await _wait_for_event(client_events, "response.done", timeout=30)
finally:
backend_task.cancel()
client_task.cancel()
await asyncio.gather(backend_task, client_task, return_exceptions=True)
# --- Assertions ---
event_types = [e.get("type") for e in client_events]
# 1. Must have received guardrail error (may not be the first error event
# if the OpenAI session emits other errors, e.g. missing parameters)
error_events = [e for e in client_events if e.get("type") == "error"]
guardrail_errors = [
e
for e in error_events
if e.get("error", {}).get("type") == "guardrail_violation"
]
assert (
len(guardrail_errors) >= 1
), f"Expected at least one guardrail_violation error but got: {[e.get('error', {}).get('type') for e in error_events]}"
guardrail_errors = [e for e in error_events if e.get("error", {}).get("type") == "guardrail_violation"]
assert len(guardrail_errors) >= 1, (
f"Expected at least one guardrail_violation error but got: {[e.get('error', {}).get('type') for e in error_events]}"
)
# 2. Must have the guardrail message surfaced as an AI transcript delta
transcript_deltas = [
e
for e in client_events
@ -186,64 +194,30 @@ async def test_text_message_blocked_by_guardrail_no_ai_response():
"response.audio_transcript.delta",
)
]
assert (
len(transcript_deltas) >= 1
), f"Expected guardrail message in transcript delta, got: {event_types}"
assert len(transcript_deltas) >= 1, f"Expected guardrail message in transcript delta, got: {event_types}"
# 3. No *real* AI response to the blocked content should have been
# generated. The original user message is blocked BEFORE it is
# forwarded to OpenAI, so the only thing the model ever sees is the
# guardrail's "say exactly: <block message>" prompt
# (see realtime_streaming.py). Two safe outcomes are possible:
# - the model voices the block message verbatim (older realtime
# snapshots did this -> text contains "blocked"), or
# - the model declines to repeat it (gpt-realtime tends to refuse
# verbatim-repeat instructions, e.g. "I'm sorry, but I can't
# repeat that message.").
# Both mean the blocked prompt itself was never answered, so we
# accept either. The hard invariant is that the blocked phrase must
# never leak into AI output, and the model must not have produced a
# normal answer to the user (which would have neither a block nor a
# refusal marker).
safe_markers = (
"block",
"guardrail",
"content filter",
"policy",
"can't repeat",
"cannot repeat",
"can't say",
"cannot say",
"won't repeat",
"can't assist",
"can't help",
"unable to",
"i'm sorry",
"i am sorry",
sent_frames = backend_ws.sent_messages
assert all(BLOCKED_PHRASE not in frame for frame in sent_frames), (
f"Blocked message was forwarded to OpenAI: {sent_frames}"
)
sent_types = [json.loads(frame).get("type") for frame in sent_frames]
assert sent_types.count("response.create") == 1, (
f"Expected only the guardrail's response.create to reach OpenAI, got backend frames: {sent_types}"
)
assert sent_types.count("conversation.item.create") == 1, (
f"Expected only the guardrail's conversation.item.create to reach OpenAI, got backend frames: {sent_types}"
)
done_events = [e for e in client_events if e.get("type") == "response.done"]
assert len(done_events) >= 1, f"Expected response.done, got: {event_types}"
for done in done_events:
output = done.get("response", {}).get("output", [])
ai_texts = [
c.get("text", "") or c.get("transcript", "")
for item in output
for c in item.get("content", [])
c.get("text", "") or c.get("transcript", "") for item in output for c in item.get("content", [])
]
real_ai_text = " ".join(ai_texts).strip()
if real_ai_text:
assert (
BLOCKED_PHRASE not in real_ai_text
), f"Blocked phrase leaked into AI response: {real_ai_text!r}"
normalized_ai_text = (
real_ai_text.lower()
.replace("\u2019", "'")
.replace("\u2018", "'")
.replace("\u201c", '"')
.replace("\u201d", '"')
)
assert any(
marker in normalized_ai_text for marker in safe_markers
), f"AI responded with non-guardrail content even though message was blocked: {real_ai_text!r}"
assert BLOCKED_PHRASE not in real_ai_text, f"Blocked phrase leaked into AI response: {real_ai_text!r}"
finally:
litellm.callbacks = []
@ -289,9 +263,7 @@ async def test_voice_transcript_blocked_by_guardrail():
# 1. Error event must be sent to client
error_events = [e for e in client_events if e.get("type") == "error"]
assert (
len(error_events) >= 1
), f"Expected guardrail error event, got: {event_types}"
assert len(error_events) >= 1, f"Expected guardrail error event, got: {event_types}"
assert error_events[0]["error"]["type"] == "guardrail_violation"
# 2. Check what was sent to backend.
@ -299,16 +271,12 @@ async def test_voice_transcript_blocked_by_guardrail():
# + response.create (to speak the block message). That's acceptable.
# What we assert is that a response.cancel was sent (blocking the original).
sent_to_backend = [
json.loads(c.args[0])
for c in backend_ws.send.call_args_list
if c.args and isinstance(c.args[0], str)
json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args and isinstance(c.args[0], str)
]
response_cancels = [
e for e in sent_to_backend if e.get("type") == "response.cancel"
]
assert (
len(response_cancels) >= 1 or len(sent_to_backend) == 0
), f"Guardrail should have sent response.cancel or nothing, got: {sent_to_backend}"
response_cancels = [e for e in sent_to_backend if e.get("type") == "response.cancel"]
assert len(response_cancels) >= 1 or len(sent_to_backend) == 0, (
f"Guardrail should have sent response.cancel or nothing, got: {sent_to_backend}"
)
# Note: The guardrail may or may not send transcript deltas; the error event
# (assertion #1) is the primary signal that the blocked content was handled.
@ -339,9 +307,7 @@ async def test_clean_text_message_passes_through_to_openai():
) as backend_ws:
streaming, input_queue = await _build_streaming(client_events, backend_ws)
backend_task = asyncio.create_task(
streaming.backend_to_client_send_messages()
)
backend_task = asyncio.create_task(streaming.backend_to_client_send_messages())
client_task = asyncio.create_task(streaming.client_ack_messages())
try:
@ -353,9 +319,7 @@ async def test_clean_text_message_passes_through_to_openai():
"type": "conversation.item.create",
"item": {
"role": "user",
"content": [
{"type": "input_text", "text": "Reply with just: OK"}
],
"content": [{"type": "input_text", "text": "Reply with just: OK"}],
},
}
)
@ -373,20 +337,14 @@ async def test_clean_text_message_passes_through_to_openai():
# No guardrail error should have been sent
error_events = [e for e in client_events if e.get("type") == "error"]
guardrail_errors = [
e
for e in error_events
if e.get("error", {}).get("type") == "guardrail_violation"
]
assert (
len(guardrail_errors) == 0
), f"Clean message should not trigger guardrail, got: {guardrail_errors}"
guardrail_errors = [e for e in error_events if e.get("error", {}).get("type") == "guardrail_violation"]
assert len(guardrail_errors) == 0, f"Clean message should not trigger guardrail, got: {guardrail_errors}"
# AI response must be present
done_events = [e for e in client_events if e.get("type") == "response.done"]
assert (
len(done_events) >= 1
), f"Expected response.done from OpenAI, got: {[e.get('type') for e in client_events]}"
assert len(done_events) >= 1, (
f"Expected response.done from OpenAI, got: {[e.get('type') for e in client_events]}"
)
finally:
litellm.callbacks = []

View file

@ -22,10 +22,8 @@ from litellm.proxy.proxy_server import initialize_pass_through_endpoints
# Mock the async_client used in the pass_through_request function
async def mock_request(*args, **kwargs):
mock_response = httpx.Response(200, json={"message": "Mocked response"})
mock_response.request = Mock(spec=httpx.Request)
return mock_response
async def mock_request(self, request, **kwargs):
return httpx.Response(200, json={"message": "Mocked response"}, request=request)
def remove_rerank_route(app):
@ -49,8 +47,8 @@ def client():
@pytest.mark.asyncio
async def test_pass_through_endpoint_no_headers(client, monkeypatch):
# Mock the httpx.AsyncClient.request method
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
# Mock the httpx.AsyncClient.send method
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
# Define a pass-through endpoint
@ -79,8 +77,8 @@ async def test_pass_through_endpoint_no_headers(client, monkeypatch):
@pytest.mark.asyncio
async def test_pass_through_endpoint(client, monkeypatch):
# Mock the httpx.AsyncClient.request method
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
# Mock the httpx.AsyncClient.send method
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
# Define a pass-through endpoint
@ -181,7 +179,7 @@ async def test_pass_through_endpoint_rpm_limit(
expected_status_codes,
num_users,
):
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache
@ -285,7 +283,7 @@ async def test_pass_through_endpoint_rpm_limit(
async def test_pass_through_endpoint_sequential_rpm_limit(
client, monkeypatch, auth, rpm_limit, requests_to_make, expected_status_codes
):
monkeypatch.setattr("httpx.AsyncClient.request", mock_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_request)
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache
@ -504,10 +502,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
captured_requests = []
async def mock_bing_request(*args, **kwargs):
async def mock_bing_request(self, request, **kwargs):
captured_requests.append((args, kwargs))
mock_response = httpx.Response(
captured_requests.append(request)
return httpx.Response(
200,
json={
"_type": "SearchResponse",
@ -518,11 +516,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
"value": [],
},
},
request=request,
)
mock_response.request = Mock(spec=httpx.Request)
return mock_response
monkeypatch.setattr("httpx.AsyncClient.request", mock_bing_request)
monkeypatch.setattr("httpx.AsyncClient.send", mock_bing_request)
# Define a pass-through endpoint
pass_through_endpoints = [
@ -555,8 +552,8 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
client.get("/bing/search?q=bob+barker")
client.get("/bing/search-no-merge-params?q=bob+barker")
first_transformed_url = captured_requests[0][1]["url"]
second_transformed_url = captured_requests[1][1]["url"]
first_transformed_url = captured_requests[0].url
second_transformed_url = captured_requests[1].url
# Parse URLs to compare query params order-independently
# Parse first URL
@ -573,7 +570,7 @@ async def test_pass_through_endpoint_bing(client, monkeypatch):
"setLang": ["en-US"],
"mkt": ["en-US"],
}
expected_second_params = {"setLang": ["en-US"], "mkt": ["en-US"]}
expected_second_params = {"q": ["bob barker"]}
# Assert the response - compare base URL and params separately
assert (

View file

@ -1556,6 +1556,10 @@ async def test_add_update_server_with_alias():
mock_mcp_server.registration_url = None
mock_mcp_server.token_url = None
mock_mcp_server.oauth2_flow = None
mock_mcp_server.token_exchange_endpoint = None
mock_mcp_server.audience = None
mock_mcp_server.subject_token_type = None
mock_mcp_server.token_exchange_profile = None
# Additional fields used by build_mcp_server_from_table
mock_mcp_server.extra_headers = None
mock_mcp_server.allow_all_keys = False
@ -1615,6 +1619,10 @@ async def test_add_update_server_without_alias():
mock_mcp_server.registration_url = None
mock_mcp_server.token_url = None
mock_mcp_server.oauth2_flow = None
mock_mcp_server.token_exchange_endpoint = None
mock_mcp_server.audience = None
mock_mcp_server.subject_token_type = None
mock_mcp_server.token_exchange_profile = None
# Additional fields used by build_mcp_server_from_table
mock_mcp_server.extra_headers = None
mock_mcp_server.allow_all_keys = False
@ -1674,6 +1682,10 @@ async def test_add_update_server_fallback_to_server_id():
mock_mcp_server.registration_url = None
mock_mcp_server.token_url = None
mock_mcp_server.oauth2_flow = None
mock_mcp_server.token_exchange_endpoint = None
mock_mcp_server.audience = None
mock_mcp_server.subject_token_type = None
mock_mcp_server.token_exchange_profile = None
# Additional fields used by build_mcp_server_from_table - set explicitly
# to avoid MagicMock objects being passed to Pydantic MCPServer constructor
mock_mcp_server.extra_headers = None

View file

@ -0,0 +1,93 @@
"""
Unit tests for PrometheusLogger._assemble_key_object DB access.
The post-request budget metrics run for every LLM API request. Auth has
already cached the key object for any real key in the same request, so the
metrics path must read the cache only. Falling through to the DB turns every
request whose token has no DB row (e.g. master-key requests, whose token is
an alias hash that never matches a stored key) into per-request
LiteLLM_VerificationToken and LiteLLM_DeprecatedVerificationToken queries.
"""
import datetime
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from prometheus_client import REGISTRY
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy._types import UserAPIKeyAuth
@pytest.fixture(autouse=True)
def cleanup_prometheus_registry():
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
yield
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
@pytest.fixture
def prometheus_logger():
return PrometheusLogger()
@pytest.mark.asyncio
async def test_assemble_key_object_does_not_query_db_on_cache_miss(prometheus_logger):
mock_prisma = MagicMock()
mock_prisma.get_data = AsyncMock()
cache = DualCache(in_memory_cache=InMemoryCache())
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", cache),
):
result = await prometheus_logger._assemble_key_object(
user_api_key="hashed-token-not-in-cache",
user_api_key_alias="",
key_max_budget=None,
key_spend=1.0,
response_cost=0.5,
)
mock_prisma.get_data.assert_not_called()
assert result.spend == 1.5
assert result.budget_reset_at is None
@pytest.mark.asyncio
async def test_assemble_key_object_reads_budget_reset_at_from_cache(prometheus_logger):
hashed_token = "hashed-token-in-cache"
reset_at = datetime.datetime(2026, 8, 1, tzinfo=datetime.timezone.utc)
cached_key = UserAPIKeyAuth(token=hashed_token, budget_reset_at=reset_at)
mock_prisma = MagicMock()
mock_prisma.get_data = AsyncMock()
cache = DualCache(in_memory_cache=InMemoryCache())
await cache.async_set_cache(key=hashed_token, value=cached_key)
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
patch("litellm.proxy.proxy_server.user_api_key_cache", cache),
):
result = await prometheus_logger._assemble_key_object(
user_api_key=hashed_token,
user_api_key_alias="alias",
key_max_budget=10.0,
key_spend=1.0,
response_cost=0.5,
)
mock_prisma.get_data.assert_not_called()
assert result.budget_reset_at == reset_at

View file

@ -1538,22 +1538,23 @@ def _local_model_cost_map():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = prev_env
@pytest.mark.parametrize("model", ["gpt-5.4", "gpt-realtime-2.1", "gpt-realtime-2.1-mini"])
@pytest.mark.parametrize("data_residency", ["eu", "us"])
def test_data_residency_applies_uplift(data_residency, _local_model_cost_map):
"""gpt-5.4 should apply the regional processing uplift multiplier when
data_residency is set. gpt-5.4+ (released 2026-03-05) carry the 10% uplift;
gpt-5 and older models do not."""
def test_data_residency_applies_uplift(data_residency, model, _local_model_cost_map):
"""Models released on/after 2026-03-05 (gpt-5.4/5.5 and gpt-realtime-2.1
series) apply the 10% regional processing uplift multiplier when
data_residency is set; gpt-5 and older models do not."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(
model="gpt-5.4",
model=model,
usage=usage,
custom_llm_provider="openai",
)
regional = generic_cost_per_token(
model="gpt-5.4",
model=model,
usage=usage,
custom_llm_provider="openai",
data_residency=data_residency,

View file

@ -11,7 +11,7 @@ sys.path.insert(
from datetime import datetime, timedelta, timezone
from typing import Any, Dict
from typing import Any, Dict, Optional
from unittest.mock import MagicMock, patch
from botocore.awsrequest import AWSPreparedRequest, AWSRequest
@ -2653,3 +2653,184 @@ class TestGetBedrockModelIdArnHandling:
"""invoke/ prefix stripping still works after the fix."""
model_id = self._call("invoke/anthropic.claude-3-sonnet-20240229-v1:0")
assert model_id == "anthropic.claude-3-sonnet-20240229-v1:0"
def _recomputed_sigv4_signature(url: str, secret_key: str, authorization: str, headers: Dict[str, Any], body) -> str:
import hashlib
import hmac
from urllib.parse import urlparse
parsed = urlparse(url)
credential_scope = authorization.split("Credential=")[1].split(",")[0].split("/", 1)[1]
signed_header_names = authorization.split("SignedHeaders=")[1].split(",")[0].split(";")
header_lookup = {name.lower(): str(value) for name, value in headers.items()}
header_lookup["host"] = parsed.netloc
body_bytes = body if isinstance(body, bytes) else str(body).encode()
canonical_request = "\n".join(
[
"POST",
parsed.path or "/",
"",
"".join(f"{name}:{header_lookup[name]}\n" for name in signed_header_names),
";".join(signed_header_names),
hashlib.sha256(body_bytes).hexdigest(),
]
)
string_to_sign = "\n".join(
[
"AWS4-HMAC-SHA256",
header_lookup["x-amz-date"],
credential_scope,
hashlib.sha256(canonical_request.encode()).hexdigest(),
]
)
key = f"AWS4{secret_key}".encode()
for scope_part in credential_scope.split("/"):
key = hmac.new(key, scope_part.encode(), hashlib.sha256).digest()
return hmac.new(key, string_to_sign.encode(), hashlib.sha256).hexdigest()
class TestSignRequestResign:
"""Regression: retrying a Bedrock request with headers from a previous SigV4 sign
(e.g. the /v1/messages strip-thinking-and-retry path) must produce a fresh
Authorization / X-Amz-Date for the new body, not inherit the stale ones and 403."""
URL = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/invoke"
ACCESS_KEY = "AKIAIOSFODNN7EXAMPLE"
SECRET_KEY = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
@pytest.fixture(autouse=True)
def _clean_aws_env(self, monkeypatch):
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
monkeypatch.delenv(env_var, raising=False)
def _optional_params(self) -> Dict[str, Any]:
return {
"aws_access_key_id": self.ACCESS_KEY,
"aws_secret_access_key": self.SECRET_KEY,
"aws_region_name": "us-east-1",
}
def _sign(self, headers: Dict[str, Any], request_data: Dict[str, Any]):
return BaseAWSLLM()._sign_request(
service_name="bedrock",
headers=headers,
optional_params=self._optional_params(),
request_data=request_data,
api_base=self.URL,
)
def test_resign_with_previously_signed_headers_replaces_stale_sigv4_headers(self):
original_body = {
"messages": [
{
"role": "assistant",
"content": [{"type": "thinking", "thinking": "x", "signature": ""}],
}
]
}
first_headers, _ = self._sign(headers={"Content-Type": "application/json"}, request_data=original_body)
assert first_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
stale_headers = {**first_headers, "X-Amz-Date": "20200101T000000Z"}
stripped_body = {"messages": [{"role": "user", "content": "hi"}]}
second_headers, second_signed_body = self._sign(headers=stale_headers, request_data=stripped_body)
assert second_headers["X-Amz-Date"] != "20200101T000000Z"
assert second_headers["Authorization"] != stale_headers["Authorization"]
assert second_headers["Authorization"].split("Signature=")[1] == _recomputed_sigv4_signature(
url=self.URL,
secret_key=self.SECRET_KEY,
authorization=second_headers["Authorization"],
headers=second_headers,
body=second_signed_body,
)
def test_forwarded_headers_still_added_back_after_signing(self):
signed_headers, _ = self._sign(
headers={"Content-Type": "application/json", "anthropic-version": "bedrock-2023-05-31"},
request_data={"messages": []},
)
assert signed_headers["anthropic-version"] == "bedrock-2023-05-31"
assert signed_headers["Content-Type"] == "application/json"
def test_caller_supplied_bearer_authorization_survives_signing(self):
signed_headers, _ = self._sign(
headers={"Content-Type": "application/json", "Authorization": "Bearer caller-token"},
request_data={"messages": []},
)
assert signed_headers["Authorization"] == "Bearer caller-token"
class TestGetRequestHeadersResign:
"""Regression: get_request_headers (invoke/converse/embed/image paths) must not let
stale SigV4 values present in the input headers clobber the freshly computed signature."""
URL = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/converse"
ACCESS_KEY = "AKIAIOSFODNN7EXAMPLE"
SECRET_KEY = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
SESSION_TOKEN = "fresh-session-token"
@pytest.fixture(autouse=True)
def _clean_aws_env(self, monkeypatch):
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
monkeypatch.delenv(env_var, raising=False)
def _prepare(self, headers: Dict[str, Any], data: str, extra_headers: Optional[Dict[str, str]] = None):
return BaseAWSLLM().get_request_headers(
credentials=Credentials(self.ACCESS_KEY, self.SECRET_KEY, self.SESSION_TOKEN),
aws_region_name="us-east-1",
extra_headers=extra_headers,
endpoint_url=self.URL,
data=data,
headers=headers,
)
def test_stale_sigv4_headers_in_input_replaced_by_fresh_signature(self):
first_prepped = self._prepare(
headers={"Content-Type": "application/json"},
data=json.dumps({"messages": [{"role": "user", "content": "original"}]}),
)
stale_authorization = first_prepped.headers["Authorization"]
assert stale_authorization.startswith("AWS4-HMAC-SHA256")
stale_headers = {
"Content-Type": "application/json",
"Authorization": stale_authorization,
"X-Amz-Date": "20200101T000000Z",
"X-Amz-Security-Token": "stale-session-token",
}
retry_data = json.dumps({"messages": [{"role": "user", "content": "retry"}]})
second_prepped = self._prepare(headers=stale_headers, data=retry_data)
assert second_prepped.headers["X-Amz-Date"] != "20200101T000000Z"
assert second_prepped.headers["X-Amz-Security-Token"] == self.SESSION_TOKEN
assert second_prepped.headers["Authorization"] != stale_authorization
assert second_prepped.headers["Authorization"].split("Signature=")[1] == _recomputed_sigv4_signature(
url=self.URL,
secret_key=self.SECRET_KEY,
authorization=second_prepped.headers["Authorization"],
headers=dict(second_prepped.headers),
body=retry_data,
)
def test_forwarded_headers_still_added_back_after_signing(self):
prepped = self._prepare(
headers={
"Content-Type": "application/json",
"anthropic-version": "bedrock-2023-05-31",
"user-agent": "litellm-test-client",
},
data=json.dumps({"messages": []}),
)
assert prepped.headers["anthropic-version"] == "bedrock-2023-05-31"
assert prepped.headers["user-agent"] == "litellm-test-client"
assert prepped.headers["Content-Type"] == "application/json"
def test_extra_headers_bearer_authorization_still_overrides_signature(self):
prepped = self._prepare(
headers={"Content-Type": "application/json"},
data=json.dumps({"messages": []}),
extra_headers={"Authorization": "Bearer foo"},
)
assert prepped.headers["Authorization"] == "Bearer foo"

View file

@ -1837,3 +1837,94 @@ async def test_alist_input_items_surfaces_upstream_error_status():
)
assert excinfo.value.status_code == 404
@pytest.mark.asyncio
async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_request(monkeypatch):
"""Regression: after Bedrock rejects a replayed thinking block (400 invalid signature),
the strip-and-retry re-sign must not inherit attempt 1's SigV4 Authorization/X-Amz-Date;
reusing them over the new stripped body makes AWS return 403 SignatureDoesNotMatch."""
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
monkeypatch.delenv(env_var, raising=False)
handler = BaseLLMHTTPHandler()
provider_config = AmazonAnthropicClaudeMessagesConfig()
litellm_params = GenericLiteLLMParams(
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
aws_region_name="us-east-1",
)
request_url = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/invoke"
request_body = {
"anthropic_version": "bedrock-2023-05-31",
"max_tokens": 100,
"messages": [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": [
{"type": "thinking", "thinking": "x", "signature": ""},
{"type": "text", "text": "ok"},
],
},
{"role": "user", "content": "continue"},
],
}
first_attempt_headers, signed_json_body = provider_config.sign_request(
headers={"Content-Type": "application/json"},
optional_params=dict(litellm_params),
request_data=request_body,
api_base=request_url,
api_key=None,
stream=False,
fake_stream=False,
model="test-model",
)
posts: list = []
invalid_signature_response = httpx.Response(
400,
text='{"message": "messages.1.content.0: Invalid `signature` in `thinking` block"}',
request=httpx.Request("POST", request_url),
)
ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url))
class FakeAsyncClient:
async def post(self, url, headers, data, stream=False, logging_obj=None):
posts.append({"headers": dict(headers), "data": data})
return invalid_signature_response if len(posts) == 1 else ok_response
logging_obj = Mock()
logging_obj.model_call_details = {}
response = await handler._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=FakeAsyncClient(),
request_url=request_url,
headers=dict(first_attempt_headers),
signed_json_body=signed_json_body,
request_body=request_body,
stream=False,
logging_obj=logging_obj,
provider_config=provider_config,
litellm_params=litellm_params,
api_key=None,
model="test-model",
)
assert response.status_code == 200
assert len(posts) == 2
retry_payload = json.loads(posts[1]["data"])
retry_blocks = [
block
for message in retry_payload["messages"]
if isinstance(message.get("content"), list)
for block in message["content"]
]
assert retry_blocks and all(block["type"] != "thinking" for block in retry_blocks)
retry_authorization = posts[1]["headers"]["Authorization"]
assert retry_authorization.startswith("AWS4-HMAC-SHA256")
assert retry_authorization != first_attempt_headers["Authorization"]

View file

@ -152,9 +152,9 @@ class TestVertexAIPSCEndpointSupport:
assert url == expected_url, f"Expected {expected_url}, but got {url}"
def test_standard_proxy_with_googleapis(self):
"""Test that standard proxies with googleapis.com in URL use simple format"""
"""Test that standard proxies with a path in the URL use simple format"""
vertex_base = VertexBase()
proxy_api_base = "https://my-proxy.googleapis.com"
proxy_api_base = "https://my-proxy.googleapis.com/vertex-proxy"
endpoint_id = "gemini-pro" # Not numeric
project_id = "test-project"
location = "us-central1"

View file

@ -15,6 +15,7 @@ sys.path.insert(
import litellm
from litellm.llms.vertex_ai.vertex_ai_aws_wif import VertexAIAwsWifAuth
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.types.llms.vertex_ai import VertexPartnerProvider
def run_sync(coro):
@ -774,7 +775,7 @@ class TestVertexBase:
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent",
"gemini-pro",
"Bearer token123",
"https://custom-vertex-api.com:generateContent",
"https://custom-vertex-api.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent",
),
# Test case 4: No API base provided (should return original values)
(
@ -930,6 +931,173 @@ class TestVertexBase:
result_url_no_streaming == expected_no_streaming_url
), f"Expected {expected_no_streaming_url}, got {result_url_no_streaming}"
def test_check_custom_proxy_vertex_bare_host_api_base_grafts_default_path(self):
vertex_base = VertexBase()
result_auth_header, result_url = vertex_base._check_custom_proxy(
api_base="https://aiplatform.googleapis.com",
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint="embedContent",
stream=None,
auth_header="Bearer token123",
url="https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-2:embedContent",
model="gemini-embedding-2",
)
assert result_auth_header == "Bearer token123"
assert (
result_url
== "https://aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-2:embedContent"
)
def test_check_custom_proxy_vertex_bare_host_api_base_with_trailing_slash(self):
vertex_base = VertexBase()
_, result_url = vertex_base._check_custom_proxy(
api_base="https://internal-gateway.example.com/",
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint="embedContent",
stream=None,
auth_header="Bearer token123",
url="https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-2:embedContent",
model="gemini-embedding-2",
)
assert (
result_url
== "https://internal-gateway.example.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-2:embedContent"
)
def test_check_custom_proxy_vertex_api_base_with_path_keeps_endpoint_append(self):
vertex_base = VertexBase()
gateway_api_base = "https://gateway.ai.cloudflare.com/v1/account-id/my-gateway/google-vertex-ai/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-2"
_, result_url = vertex_base._check_custom_proxy(
api_base=gateway_api_base,
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint="embedContent",
stream=None,
auth_header="Bearer token123",
url="https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-2:embedContent",
model="gemini-embedding-2",
)
assert result_url == f"{gateway_api_base}:embedContent"
def test_check_custom_proxy_vertex_bare_host_streaming_keeps_single_alt_sse(self):
vertex_base = VertexBase()
_, result_url = vertex_base._check_custom_proxy(
api_base="https://internal-gateway.example.com",
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint="streamGenerateContent",
stream=True,
auth_header="Bearer token123",
url="https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-2.5-pro:streamGenerateContent?alt=sse",
model="gemini-2.5-pro",
)
assert (
result_url
== "https://internal-gateway.example.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-2.5-pro:streamGenerateContent?alt=sse"
)
def test_check_custom_proxy_psc_endpoint_format_unaffected_by_bare_host(self):
vertex_base = VertexBase()
_, result_url = vertex_base._check_custom_proxy(
api_base="https://10.96.32.8",
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint="predict",
stream=None,
auth_header="Bearer token123",
url="",
model="1234567890",
vertex_project="test-project",
vertex_location="us-central1",
vertex_api_version="v1",
use_psc_endpoint_format=True,
)
assert result_url == "https://10.96.32.8/v1/projects/test-project/locations/us-central1/endpoints/1234567890:predict"
@pytest.mark.parametrize(
"custom_api_base, stream, expected_url",
[
(
"https://aiplatform-myendpoint.p.googleapis.com",
False,
"https://aiplatform-myendpoint.p.googleapis.com/v1/projects/test-project/locations/global/endpoints/openapi/chat/completions",
),
(
"https://aiplatform-myendpoint.p.googleapis.com",
True,
"https://aiplatform-myendpoint.p.googleapis.com/v1/projects/test-project/locations/global/endpoints/openapi/chat/completions",
),
(
"https://gateway.example.com/vertex-proxy",
False,
"https://gateway.example.com/vertex-proxy/v1/projects/test-project/locations/global/endpoints/openapi/chat/completions",
),
],
ids=["psc-host", "psc-host-streaming", "api-base-with-path"],
)
def test_get_complete_vertex_url_openai_path_partner_custom_api_base(
self, custom_api_base, stream, expected_url
):
vertex_base = VertexBase()
result = vertex_base.get_complete_vertex_url(
custom_api_base=custom_api_base,
vertex_location="global",
vertex_project="test-project",
project_id="test-project",
partner=VertexPartnerProvider.llama,
stream=stream,
model="minimaxai/minimax-m2-maas",
)
assert result == expected_url
assert result.count("://") == 1
def test_get_complete_vertex_url_openai_path_partner_default_api_base(self):
vertex_base = VertexBase()
result = vertex_base.get_complete_vertex_url(
custom_api_base=None,
vertex_location="us-central1",
vertex_project="test-project",
project_id="test-project",
partner=VertexPartnerProvider.llama,
stream=True,
model="meta/llama-3.1-405b-instruct-maas",
)
assert (
result
== "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/endpoints/openapi/chat/completions"
)
def test_get_complete_vertex_url_rawpredict_partner_custom_api_base_keeps_endpoint_format(self):
vertex_base = VertexBase()
result = vertex_base.get_complete_vertex_url(
custom_api_base="https://gateway.example.com/vertex-proxy",
vertex_location="us-central1",
vertex_project="test-project",
project_id="test-project",
partner=VertexPartnerProvider.mistralai,
stream=False,
model="mistral-large-2411",
)
assert result == "https://gateway.example.com/vertex-proxy:rawPredict"
@pytest.mark.parametrize(
"api_base, custom_llm_provider, gemini_api_key, endpoint, stream, auth_header, url, model, expected_auth_header, expected_url",
[

View file

@ -387,8 +387,9 @@ async def test_pass_through_request_stream_param_no_override(
# Create mocks for the async client
mock_async_client = AsyncMock()
# Mock request to return the non-streaming response
mock_async_client.request.return_value = mock_response
# Mock build_request/send to return the non-streaming response
mock_async_client.build_request = Mock(return_value=Mock())
mock_async_client.send.return_value = mock_response
# Mock get_async_httpx_client to return our mock client
mock_client_obj = Mock()
@ -420,20 +421,19 @@ async def test_pass_through_request_stream_param_no_override(
stream=False, # Should be used since no stream in request body
)
# Verify that build_request was NOT called (no streaming path)
mock_async_client.build_request.assert_not_called()
# Verify that send was NOT called (no streaming path)
mock_async_client.send.assert_not_called()
# Verify that the non-streaming request method WAS called
mock_async_client.request.assert_called_once_with(
method="POST",
url=httpx.URL("https://api.anthropic.com/v1/messages"),
# Non-SSE requests are sent with stream semantics so large bodies can
# be relayed without buffering; the JSON response below is still
# buffered into a plain Response.
mock_async_client.request.assert_not_called()
mock_async_client.build_request.assert_called_once_with(
"POST",
httpx.URL("https://api.anthropic.com/v1/messages"),
headers={"Authorization": "Bearer test-key"},
params={},
json=request_body,
)
mock_async_client.send.assert_called_once()
assert mock_async_client.send.call_args.kwargs.get("stream") is True
# Verify response is a regular Response (not StreamingResponse)
from fastapi.responses import Response, StreamingResponse

View file

@ -3,8 +3,8 @@ import asyncio
import pytest
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InvalidatableOAuthTokenStore,
OAuthToken,
OAuthTokenStore,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import (
LazyPerUserOAuthTokenStore,
@ -16,11 +16,15 @@ class _RecordingStore:
def __init__(self, access_token: str) -> None:
self._access_token = access_token
self.calls: list[tuple[str, str]] = []
self.invalidations: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
return OAuthToken(access_token=self._access_token)
async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))
class _BlockingStore:
def __init__(self, access_token: str) -> None:
@ -28,6 +32,7 @@ class _BlockingStore:
self.started = asyncio.Event()
self.release = asyncio.Event()
self.calls: list[tuple[str, str]] = []
self.invalidations: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None:
self.calls.append((user_id, server_id))
@ -35,6 +40,9 @@ class _BlockingStore:
await self.release.wait()
return OAuthToken(access_token=self._access_token)
async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))
class _RedisAvailability:
def __init__(self) -> None:
@ -59,7 +67,7 @@ async def test_lazy_store_rebuilds_when_redis_becomes_available() -> None:
redis_available = _RedisAvailability()
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
if redis_available.available:
@ -94,7 +102,7 @@ async def test_lazy_store_allows_concurrent_local_fetches_without_redis() -> Non
redis_available = _RedisAvailability()
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
@ -127,7 +135,7 @@ async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild()
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()
def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]:
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
if redis_available.available:
return redis_store, True
return local_store, False
@ -158,3 +166,83 @@ async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild()
assert second is not None and second.access_token == "redis"
assert local_store.calls == [("u", "s")]
assert redis_store.calls == [("u", "s")]
@pytest.mark.asyncio
async def test_lazy_store_invalidate_builds_chain_and_delegates() -> None:
local_store = _RecordingStore("local")
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=_RedisAvailability(),
)
await store.invalidate("u", "s")
assert build_calls == 1
assert local_store.invalidations == [("u", "s")]
@pytest.mark.asyncio
async def test_lazy_store_invalidate_reaches_the_store_fetch_reads() -> None:
local_store = _RecordingStore("local")
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return local_store, False
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=_RedisAvailability(),
)
await store.fetch("u", "s")
await store.invalidate("u", "s")
assert build_calls == 1
assert local_store.calls == [("u", "s")]
assert local_store.invalidations == [("u", "s")]
@pytest.mark.asyncio
async def test_lazy_store_invalidate_works_after_redis_chain_is_built() -> None:
redis_store = _RecordingStore("redis")
redis_available = _RedisAvailability()
redis_available.available = True
build_calls = 0
def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]:
nonlocal build_calls
build_calls += 1
return redis_store, True
def server_lookup(_server_id: str) -> None:
return None
store = LazyPerUserOAuthTokenStore(
server_lookup,
store_builder=build_store,
redis_available=redis_available,
)
await store.fetch("u", "s")
await store.invalidate("u", "s")
assert build_calls == 1
assert redis_store.invalidations == [("u", "s")]

View file

@ -17,6 +17,7 @@ import pytest
from litellm.proxy._experimental.mcp_server.db import (
_decode_user_credential,
_prepare_mcp_server_data,
get_user_credential,
get_user_oauth_credential,
is_oauth_credential_expired,
@ -27,10 +28,12 @@ from litellm.proxy._experimental.mcp_server.db import (
store_user_credential,
store_user_oauth_credential,
)
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.types.mcp import MCPAuth, MCPTransport
SALT_KEY = "test-salt-key-for-byok-credential-tests-1234"
@ -722,3 +725,50 @@ async def test_refresh_user_oauth_token_defaults_to_client_secret_post(monkeypat
assert "Authorization" not in kwargs["headers"]
assert kwargs["data"]["client_id"] == "cid"
assert kwargs["data"]["client_secret"] == "sec"
def test_prepare_mcp_server_data_create_carries_token_exchange_columns():
"""The create path (POST /v1/mcp/server) must emit token_exchange_endpoint/audience/
subject_token_type as top-level column values so an auth_type=oauth2_token_exchange server
persists via the REST API, not only via config.yaml. Dropping the fields from the request
model would leave them out of the prepared column data."""
request = NewMCPServerRequest(
server_name="te_write",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/oauth2/token",
audience="https://upstream.example.com",
subject_token_type="urn:ietf:params:oauth:token-type:jwt",
token_exchange_profile="entra_obo",
credentials={"client_id": "te-client", "client_secret": "te-secret"},
)
data = _prepare_mcp_server_data(request)
assert data["token_exchange_endpoint"] == "https://idp.example.com/oauth2/token"
assert data["audience"] == "https://upstream.example.com"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:jwt"
assert data["token_exchange_profile"] == "entra_obo"
def test_prepare_mcp_server_data_update_carries_token_exchange_columns():
"""The partial-update path (PUT /v1/mcp/server, exclude_unset) must carry the three
token-exchange columns when the caller provides them."""
request = UpdateMCPServerRequest(
server_id="te-update",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/oauth2/token",
audience="https://upstream.example.com",
subject_token_type="urn:ietf:params:oauth:token-type:jwt",
token_exchange_profile="entra_obo",
)
data = _prepare_mcp_server_data(request, exclude_unset=True)
assert data["token_exchange_endpoint"] == "https://idp.example.com/oauth2/token"
assert data["audience"] == "https://upstream.example.com"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:jwt"
assert data["token_exchange_profile"] == "entra_obo"

View file

@ -4055,3 +4055,104 @@ async def test_oauth_authorization_server_404_for_unknown_server_name():
mcp_server_name="does_not_exist",
)
assert exc_info.value.status_code == 404
@pytest.mark.asyncio
async def test_store_per_user_token_server_side_invalidates_v2_token_cache():
"""A token stored by the OAuth callback (code exchange or refresh) drops the v2 per-user
token cache entry, so egress stops serving the replaced token immediately instead of
until its TTL."""
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_store_per_user_token_server_side,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="srv-cb-1",
name="cb_server",
url="https://upstream.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
)
invalidate_mock = AsyncMock(return_value=None)
cache_set_mock = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.utils.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch(
"litellm.proxy._experimental.mcp_server.db.store_user_oauth_credential",
new=AsyncMock(return_value=None),
),
patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_per_user_token_cache.set",
new=cache_set_mock,
),
patch.object(
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
await _store_per_user_token_server_side(
server=server,
user_id="user-cb-1",
token_response={"access_token": "fresh-tok", "expires_in": 3600},
)
invalidate_mock.assert_awaited_once_with("user-cb-1", "srv-cb-1")
cache_set_mock.assert_awaited_once()
@pytest.mark.asyncio
async def test_store_per_user_token_server_side_skips_invalidate_when_db_write_fails():
"""A failed DB write neither warms the v1 cache nor drops the v2 cache entry; the
previously stored token is still the truth."""
from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_store_per_user_token_server_side,
)
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
server = MCPServer(
server_id="srv-cb-2",
name="cb_server_2",
url="https://upstream.example/mcp",
transport="http",
auth_type=MCPAuth.oauth2,
)
invalidate_mock = AsyncMock(return_value=None)
cache_set_mock = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.utils.get_prisma_client_or_throw",
return_value=MagicMock(),
),
patch(
"litellm.proxy._experimental.mcp_server.db.store_user_oauth_credential",
new=AsyncMock(side_effect=RuntimeError("db down")),
),
patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_per_user_token_cache.set",
new=cache_set_mock,
),
patch.object(
manager_module.global_mcp_server_manager,
"invalidate_user_oauth_token_cache",
new=invalidate_mock,
),
):
await _store_per_user_token_server_side(
server=server,
user_id="user-cb-2",
token_response={"access_token": "fresh-tok", "expires_in": 3600},
)
invalidate_mock.assert_not_awaited()
cache_set_mock.assert_not_awaited()

View file

@ -164,6 +164,25 @@ async def test_openapi_backed_server_also_respects_the_cap():
assert tracker.peak_by_server["srv-openapi"] == 2
@pytest.mark.asyncio
async def test_edited_limit_takes_effect_without_restart():
"""Editing max_concurrent_requests must rebuild the cached semaphore so the
new cap applies to subsequent calls immediately, not only after a restart."""
manager = MCPServerManager()
server = _make_server("srv-edited", max_concurrent_requests=3)
before_edit = _ConcurrencyTracker()
with _patch_client_with_tracker(manager, before_edit):
await _fire(manager, server, n=6)
assert before_edit.peak_by_server["srv-edited"] == 3
server.max_concurrent_requests = 1
after_edit = _ConcurrencyTracker()
with _patch_client_with_tracker(manager, after_edit):
await _fire(manager, server, n=6)
assert after_edit.peak_by_server["srv-edited"] == 1
def test_semaphore_is_reused_per_server_and_distinct_across_servers():
manager = MCPServerManager()
server_a = _make_server("srv-a", max_concurrent_requests=3)

View file

@ -7,6 +7,7 @@ Omitting a field must NOT reset it to its Pydantic schema default (e.g.
would silently overwrite the existing DB row.
"""
import json
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -134,14 +135,10 @@ async def test_partial_update_writes_explicitly_provided_fields():
@pytest.mark.asyncio
async def test_partial_update_can_explicitly_reset_allow_all_keys():
"""Caller can still reset a field to its default by sending it explicitly."""
enabled = await _run_update(
UpdateMCPServerRequest(server_id="s", allow_all_keys=True)
)
enabled = await _run_update(UpdateMCPServerRequest(server_id="s", allow_all_keys=True))
assert enabled["allow_all_keys"] is True
disabled = await _run_update(
UpdateMCPServerRequest(server_id="s", allow_all_keys=False)
)
disabled = await _run_update(UpdateMCPServerRequest(server_id="s", allow_all_keys=False))
assert disabled["allow_all_keys"] is False
@ -178,6 +175,95 @@ async def test_partial_update_can_explicitly_clear_alias():
assert data_dict["alias"] is None
async def _run_update_with_existing(data: UpdateMCPServerRequest, existing_auth_type: str) -> dict:
mock_prisma = _mock_prisma()
existing = MagicMock()
existing.auth_type = existing_auth_type
existing.credentials = None
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
await update_mcp_server(mock_prisma, data, "test-user")
return mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
@pytest.mark.asyncio
async def test_auth_type_switch_clears_stale_flow_scoped_fields():
"""
Switching oauth2 -> oauth2_token_exchange must clear the previous flow's
endpoint config: a stale token_url would otherwise be picked up as the
token-exchange endpoint and suppress RFC 9728/8414 discovery.
"""
data = UpdateMCPServerRequest(server_id="my-test-server", auth_type="oauth2_token_exchange")
data_dict = await _run_update_with_existing(data, existing_auth_type="oauth2")
for stale_field in (
"authorization_url",
"token_url",
"registration_url",
"oauth2_flow",
"token_exchange_endpoint",
"audience",
"subject_token_type",
"token_exchange_profile",
):
assert data_dict[stale_field] is None, f"{stale_field} must be cleared on auth_type switch"
assert data_dict["credentials"] is None
@pytest.mark.asyncio
async def test_auth_type_switch_keeps_explicitly_provided_flow_fields():
"""Fields explicitly provided alongside the auth_type switch must survive it."""
data = UpdateMCPServerRequest(
server_id="my-test-server",
auth_type="oauth2_token_exchange",
token_exchange_endpoint="https://idp.example.com/oauth2/token",
)
data_dict = await _run_update_with_existing(data, existing_auth_type="oauth2")
assert data_dict["token_exchange_endpoint"] == "https://idp.example.com/oauth2/token"
assert data_dict["token_url"] is None
@pytest.mark.asyncio
async def test_auth_type_switch_back_to_oauth2_clears_token_exchange_fields():
"""The reverse switch must not leave token-exchange settings behind to
silently reactivate if the server is later switched back."""
data = UpdateMCPServerRequest(server_id="my-test-server", auth_type="oauth2")
data_dict = await _run_update_with_existing(data, existing_auth_type="oauth2_token_exchange")
assert data_dict["token_exchange_endpoint"] is None
assert data_dict["audience"] is None
assert data_dict["subject_token_type"] is None
assert data_dict["token_exchange_profile"] is None
@pytest.mark.asyncio
async def test_unchanged_auth_type_does_not_clear_flow_fields():
"""An update that keeps the auth_type must not touch flow-scoped fields, so a
legacy OBO server using token_url as its exchange endpoint keeps working."""
data = UpdateMCPServerRequest(
server_id="my-test-server",
auth_type="oauth2_token_exchange",
allowed_tools=["foo"],
)
data_dict = await _run_update_with_existing(data, existing_auth_type="oauth2_token_exchange")
for flow_field in (
"authorization_url",
"token_url",
"registration_url",
"oauth2_flow",
"token_exchange_endpoint",
"audience",
"subject_token_type",
"token_exchange_profile",
):
assert flow_field not in data_dict
@pytest.mark.asyncio
async def test_create_still_writes_defaults():
"""
@ -203,3 +289,256 @@ async def test_create_still_writes_defaults():
# audit fields set by create_mcp_server.
assert data_dict["created_by"] == "test-user"
assert data_dict["updated_by"] == "test-user"
# ── token-exchange blob → column normalization ────────────────────────────────
#
# token_exchange_endpoint / audience / subject_token_type have dedicated columns;
# their MCPCredentials copies are a legacy shape. Writes must lift blob values
# into the columns and strip them from the stored blob so the read-time
# ``column or blob`` fallback can never resurrect a stale blob value after the
# column is cleared.
@pytest.fixture(autouse=True)
def _salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234")
def _existing_row(auth_type: str, credentials: dict | None = None):
existing = MagicMock()
existing.auth_type = auth_type
existing.credentials = json.dumps(credentials) if credentials is not None else None
existing.token_exchange_endpoint = None
existing.audience = None
existing.subject_token_type = None
existing.token_exchange_profile = None
return existing
@pytest.mark.asyncio
async def test_create_lifts_blob_token_exchange_settings_into_columns():
"""The legacy REST shape (TE settings inside ``credentials``) must land in
the dedicated columns, and the stored blob must not keep a copy."""
mock_prisma = _mock_prisma()
data = NewMCPServerRequest(
server_id="te-server",
url="https://example.com/mcp",
transport="http",
auth_type="oauth2_token_exchange",
credentials={
"client_id": "cid",
"client_secret": "sec",
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://upstream",
"subject_token_type": "urn:ietf:params:oauth:token-type:jwt",
"token_exchange_profile": "entra_obo",
},
)
await create_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.create.call_args[1]["data"]
assert data_dict["token_exchange_endpoint"] == "https://idp.example.com/oauth2/token"
assert data_dict["audience"] == "api://upstream"
assert data_dict["subject_token_type"] == "urn:ietf:params:oauth:token-type:jwt"
assert data_dict["token_exchange_profile"] == "entra_obo"
stored_blob = json.loads(data_dict["credentials"])
for te_field in ("token_exchange_endpoint", "audience", "subject_token_type", "token_exchange_profile"):
assert te_field not in stored_blob
assert "client_id" in stored_blob
@pytest.mark.asyncio
async def test_create_explicit_column_wins_over_blob_copy():
mock_prisma = _mock_prisma()
data = NewMCPServerRequest(
server_id="te-server",
url="https://example.com/mcp",
transport="http",
auth_type="oauth2_token_exchange",
token_exchange_endpoint="https://top-level.example.com/token",
credentials={"client_id": "cid", "token_exchange_endpoint": "https://blob.example.com/token"},
)
await create_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.create.call_args[1]["data"]
assert data_dict["token_exchange_endpoint"] == "https://top-level.example.com/token"
assert "token_exchange_endpoint" not in json.loads(data_dict["credentials"])
@pytest.mark.asyncio
async def test_credentials_merge_migrates_legacy_blob_te_settings():
"""A same-auth credentials update on a legacy row (TE settings in the blob,
columns null) must move the settings to the columns and drop them from the
merged blob."""
mock_prisma = _mock_prisma()
existing = _existing_row(
"oauth2_token_exchange",
credentials={
"client_id": "enc-old-cid",
"token_exchange_endpoint": "https://legacy-idp.example.com/token",
"audience": "api://legacy",
"token_exchange_profile": "entra_obo",
},
)
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
data = UpdateMCPServerRequest(
server_id="te-server",
auth_type="oauth2_token_exchange",
credentials={"client_id": "new-cid"},
)
await update_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
assert data_dict["token_exchange_endpoint"] == "https://legacy-idp.example.com/token"
assert data_dict["audience"] == "api://legacy"
assert data_dict["token_exchange_profile"] == "entra_obo"
merged_blob = json.loads(data_dict["credentials"])
for te_field in ("token_exchange_endpoint", "audience", "subject_token_type", "token_exchange_profile"):
assert te_field not in merged_blob
@pytest.mark.asyncio
async def test_cleared_column_is_not_resurrected_by_legacy_blob_value():
"""The Greptile scenario: explicitly clearing the column (to re-enable
RFC 9728/8414 discovery) while the legacy blob still holds an endpoint must
NOT resurrect the blob value — the explicit null wins and the blob copy is
stripped."""
mock_prisma = _mock_prisma()
existing = _existing_row(
"oauth2_token_exchange",
credentials={"client_id": "enc-old-cid", "token_exchange_endpoint": "https://dead-idp.example.com/token"},
)
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
data = UpdateMCPServerRequest(
server_id="te-server",
auth_type="oauth2_token_exchange",
token_exchange_endpoint=None,
credentials={"client_id": "new-cid"},
)
await update_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
assert data_dict["token_exchange_endpoint"] is None
assert "token_exchange_endpoint" not in json.loads(data_dict["credentials"])
@pytest.mark.asyncio
async def test_merge_strips_blob_te_copy_when_column_already_set():
"""When the row already has a column value, the blob copy is shadowed at
read time anyway — the merge must strip it rather than carry it forward."""
mock_prisma = _mock_prisma()
existing = _existing_row(
"oauth2_token_exchange",
credentials={"client_id": "enc-old-cid", "token_exchange_endpoint": "https://blob-copy.example.com/token"},
)
existing.token_exchange_endpoint = "https://column.example.com/token"
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
data = UpdateMCPServerRequest(
server_id="te-server",
auth_type="oauth2_token_exchange",
credentials={"client_id": "new-cid"},
)
await update_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
# Column untouched by this update (not in payload), blob copy gone.
assert "token_exchange_endpoint" not in data_dict
assert "token_exchange_endpoint" not in json.loads(data_dict["credentials"])
@pytest.mark.asyncio
async def test_auth_type_switch_clears_flow_fields_with_external_fields_set():
"""The management endpoint passes ``fields_set`` explicitly (PUT
/v1/mcp/server). The auth-switch clearing must fire on that path too — it is
gated on ``data.auth_type``/the existing row, not on how fields_set arrives."""
mock_prisma = _mock_prisma()
existing = _existing_row("oauth2")
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
data = UpdateMCPServerRequest(server_id="te-server", auth_type="oauth2_token_exchange")
await update_mcp_server(mock_prisma, data, "test-user", fields_set=set(data.fields_set()))
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
for stale_field in (
"authorization_url",
"token_url",
"registration_url",
"oauth2_flow",
"token_exchange_endpoint",
"audience",
"subject_token_type",
"token_exchange_profile",
):
assert data_dict[stale_field] is None, f"{stale_field} must be cleared via the fields_set path"
@pytest.mark.asyncio
async def test_explicit_clear_without_credentials_purges_legacy_blob_copy():
"""Clearing a column in an update that does not touch credentials must strip
the legacy blob copy too — otherwise the next credentials update's
migrate-on-write would repopulate the column the admin just cleared."""
mock_prisma = _mock_prisma()
existing = _existing_row(
"oauth2_token_exchange",
credentials={"client_id": "enc-old-cid", "token_exchange_endpoint": "https://dead-idp.example.com/token"},
)
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
data = UpdateMCPServerRequest(server_id="te-server", token_exchange_endpoint=None)
await update_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
assert data_dict["token_exchange_endpoint"] is None
stored_blob = json.loads(data_dict["credentials"])
assert "token_exchange_endpoint" not in stored_blob
# Unrelated blob keys (encrypted secrets) survive untouched.
assert stored_blob["client_id"] == "enc-old-cid"
@pytest.mark.asyncio
async def test_explicit_te_write_without_credentials_migrates_other_legacy_fields():
"""A no-credentials update that writes one token-exchange column migrates the
whole row: untouched null columns are lifted from the blob, and every blob
copy is stripped."""
mock_prisma = _mock_prisma()
existing = _existing_row(
"oauth2_token_exchange",
credentials={
"client_id": "enc-old-cid",
"token_exchange_endpoint": "https://legacy-idp.example.com/token",
"audience": "api://legacy",
},
)
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
data = UpdateMCPServerRequest(server_id="te-server", audience="api://new")
await update_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
assert data_dict["audience"] == "api://new"
assert data_dict["token_exchange_endpoint"] == "https://legacy-idp.example.com/token"
stored_blob = json.loads(data_dict["credentials"])
for te_field in ("token_exchange_endpoint", "audience", "subject_token_type"):
assert te_field not in stored_blob
@pytest.mark.asyncio
async def test_te_update_without_blob_te_keys_leaves_credentials_untouched():
"""A no-credentials column write on a row whose blob has no legacy copies
must not rewrite the credentials blob at all."""
mock_prisma = _mock_prisma()
existing = _existing_row("oauth2_token_exchange", credentials={"client_id": "enc-old-cid"})
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
data = UpdateMCPServerRequest(server_id="te-server", token_exchange_endpoint="https://new.example.com/token")
await update_mcp_server(mock_prisma, data, "test-user")
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
assert data_dict["token_exchange_endpoint"] == "https://new.example.com/token"
assert "credentials" not in data_dict

View file

@ -5184,6 +5184,10 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow():
legacy_server.server_id = "legacy-m2m-id"
legacy_server.auth_type = MCPAuth.oauth2
legacy_server.oauth2_flow = None # Legacy: field not set in DB
legacy_server.token_exchange_endpoint = None
legacy_server.audience = None
legacy_server.subject_token_type = None
legacy_server.token_exchange_profile = None
legacy_server.token_url = "https://oauth.example.com/token"
legacy_server.authorization_url = None
legacy_server.client_id = "client-id"

View file

@ -1772,6 +1772,66 @@ class TestMCPServerManager:
assert isinstance(result, Error)
assert result.error.tag == "misconfigured"
@pytest.mark.asyncio
async def test_load_servers_from_config_reads_all_token_exchange_fields(self):
"""Every token-exchange setting is configurable through config.yaml as a top-level
key (the config counterpart of the REST/UI columns) and reaches the resolver spec;
omitted keys resolve to their documented defaults. token_exchange servers need no
oauth2_flow (that requirement is oauth2-only)."""
from litellm.types.mcp import DEFAULT_SUBJECT_TOKEN_TYPE
manager = MCPServerManager()
config = {
"te_full": {
"url": "https://up.example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2_token_exchange,
"token_exchange_endpoint": "https://idp.example.com/oauth2/token",
"audience": "api://upstream",
"subject_token_type": "urn:ietf:params:oauth:token-type:jwt",
"token_exchange_profile": "entra_obo",
"client_id": "cid",
"client_secret": "csec",
"scopes": ["api://upstream/.default"],
},
"te_minimal": {
"url": "https://up2.example.com/mcp",
"transport": MCPTransport.http,
"auth_type": MCPAuth.oauth2_token_exchange,
"token_exchange_endpoint": "https://idp2.example.com/oauth2/token",
"client_id": "cid2",
"client_secret": "csec2",
},
}
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)):
await manager.load_servers_from_config(config)
by_name = {s.server_name: s for s in manager.config_mcp_servers.values()}
full = by_name["te_full"]
assert full.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert full.audience == "api://upstream"
assert full.subject_token_type == "urn:ietf:params:oauth:token-type:jwt"
assert full.token_exchange_profile == "entra_obo"
minimal = by_name["te_minimal"]
assert minimal.audience is None
assert minimal.subject_token_type == DEFAULT_SUBJECT_TOKEN_TYPE
assert minimal.token_exchange_profile == "rfc8693"
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import to_server_spec
spec = to_server_spec(full)
assert spec is not None
assert spec.config.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert spec.config.profile == "entra_obo"
minimal_spec = to_server_spec(minimal)
assert minimal_spec is not None
assert minimal_spec.config.subject_token_type == DEFAULT_SUBJECT_TOKEN_TYPE
assert minimal_spec.config.profile == "rfc8693"
@pytest.mark.asyncio
async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self):
manager = MCPServerManager()
@ -2880,6 +2940,39 @@ class TestMCPServerManager:
assert await manager.has_user_oauth_token(server, user_auth) is False
assert calls == [] # short-circuited on the None spec, never hit the resolver
@pytest.mark.asyncio
async def test_invalidate_user_oauth_token_cache_delegates_to_store(self):
"""The write side's cache drop reaches the same per-user store the resolver reads."""
class _Store:
def __init__(self) -> None:
self.invalidations: list[tuple[str, str]] = []
async def fetch(self, user_id: str, server_id: str):
return None
async def invalidate(self, user_id: str, server_id: str) -> None:
self.invalidations.append((user_id, server_id))
store = _Store()
manager = MCPServerManager(per_user_oauth_token_store=store)
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
assert store.invalidations == [("alice", "srv-1")]
@pytest.mark.asyncio
async def test_invalidate_user_oauth_token_cache_swallows_store_errors(self):
"""A cache-drop failure must not fail the credential write that triggered it."""
class _Store:
async def fetch(self, user_id: str, server_id: str):
return None
async def invalidate(self, user_id: str, server_id: str) -> None:
raise RuntimeError("redis down")
manager = MCPServerManager(per_user_oauth_token_store=_Store())
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
@pytest.mark.asyncio
async def test_resolve_oauth2_headers_no_user_id(self):
"""Skip lookup entirely when user_api_key_auth has no user_id."""
@ -4387,6 +4480,177 @@ class TestMCPServerTimestamps:
assert "0.01s" in exc_info.value.detail["message"]
class TestMCPServerTokenExchangeColumns:
"""Token-exchange (RFC 8693) config persists through the dedicated columns added for the
create/update REST + DB path, mirroring how ``token_url`` is stored. The credentials JSON
blob is kept as a read-fallback so servers persisted before the columns existed still load."""
@pytest.mark.asyncio
async def test_build_mcp_server_from_table_reads_token_exchange_columns(self):
"""The DB->runtime loader must read the three fields from the dedicated columns. Before the
columns existed it only read the credentials blob, so column values would be dropped."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-cols",
server_name="te_cols",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/oauth2/token",
audience="https://upstream.example.com",
subject_token_type="urn:ietf:params:oauth:token-type:jwt",
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
assert mcp_server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert mcp_server.audience == "https://upstream.example.com"
assert mcp_server.subject_token_type == "urn:ietf:params:oauth:token-type:jwt"
@pytest.mark.asyncio
async def test_build_mcp_server_from_table_falls_back_to_credentials_blob(self):
"""Backwards compatibility: a server whose token-exchange config lives only in the
credentials blob (no columns) must still load with those values."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-blob",
server_name="te_blob",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
credentials={
"token_exchange_endpoint": "https://idp.example.com/legacy/token",
"audience": "legacy-audience",
"subject_token_type": "urn:ietf:params:oauth:token-type:saml2",
},
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
assert mcp_server.token_exchange_endpoint == "https://idp.example.com/legacy/token"
assert mcp_server.audience == "legacy-audience"
assert mcp_server.subject_token_type == "urn:ietf:params:oauth:token-type:saml2"
@pytest.mark.asyncio
async def test_build_mcp_server_from_table_subject_token_type_defaults(self):
"""subject_token_type falls back to the RFC 8693 access_token URN when unset."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-default",
server_name="te_default",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/oauth2/token",
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
assert mcp_server.subject_token_type == "urn:ietf:params:oauth:token-type:access_token"
@pytest.mark.asyncio
async def test_round_trip_token_exchange_columns_preserved(self):
"""The three fields survive LiteLLM_MCPServerTable -> MCPServer -> LiteLLM_MCPServerTable.
Before the table builder wrote them back, a registry round-trip dropped them."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-rt",
server_name="te_rt",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/oauth2/token",
audience="https://upstream.example.com",
subject_token_type="urn:ietf:params:oauth:token-type:jwt",
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
rebuilt_table = manager._build_mcp_server_table(mcp_server)
assert rebuilt_table.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert rebuilt_table.audience == "https://upstream.example.com"
assert rebuilt_table.subject_token_type == "urn:ietf:params:oauth:token-type:jwt"
@pytest.mark.asyncio
async def test_build_mcp_server_from_table_reads_token_exchange_profile_column(self):
"""The profile dialect selector (rfc8693 vs entra_obo) is read from its dedicated column
so a server created via the REST API/UI as entra_obo resolves to the Entra dialect."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-profile",
server_name="te_profile",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://login.microsoftonline.com/tenant/oauth2/v2.0/token",
token_exchange_profile="entra_obo",
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
assert mcp_server.token_exchange_profile == "entra_obo"
@pytest.mark.asyncio
async def test_build_mcp_server_from_table_token_exchange_profile_defaults_rfc8693(self):
"""token_exchange_profile falls back to rfc8693 when neither column nor blob sets it."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-profile-default",
server_name="te_profile_default",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/oauth2/token",
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
assert mcp_server.token_exchange_profile == "rfc8693"
@pytest.mark.asyncio
async def test_build_mcp_server_from_table_token_exchange_profile_blob_fallback(self):
"""Backwards compatibility: a server with the profile only in the credentials blob still loads."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-profile-blob",
server_name="te_profile_blob",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
credentials={"token_exchange_profile": "entra_obo"},
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
assert mcp_server.token_exchange_profile == "entra_obo"
@pytest.mark.asyncio
async def test_round_trip_token_exchange_profile_preserved(self):
"""token_exchange_profile survives LiteLLM_MCPServerTable -> MCPServer -> LiteLLM_MCPServerTable."""
manager = MCPServerManager()
table_record = LiteLLM_MCPServerTable(
server_id="te-profile-rt",
server_name="te_profile_rt",
url="https://upstream.example.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_profile="entra_obo",
)
mcp_server = await manager.build_mcp_server_from_table(table_record)
rebuilt_table = manager._build_mcp_server_table(mcp_server)
assert rebuilt_table.token_exchange_profile == "entra_obo"
class TestInternalDelegatePkceWarningLog:
@pytest.mark.asyncio
async def test_build_mcp_server_logs_on_internal_delegate_interactive(self, caplog):
@ -5974,6 +6238,83 @@ class TestOBOCallToolRetry:
assert first.attempts == 1 and retry.attempts == 1
class TestOBOConcurrencyLimit:
"""OBO (token_exchange) tool calls must honor the server's max_concurrent_requests.
Regression: the token_exchange dispatch built its coroutine outside
_limit_outbound_concurrency, so OBO calls skipped the per-server semaphore the
non-OBO path enforces and a caller could exceed the admin-configured cap.
"""
@pytest.mark.asyncio
async def test_obo_dispatch_respects_max_concurrent_requests(self):
max_concurrent = 2
overflow = 3
server = MCPServer(
server_id="obo-concurrency",
name="obo",
url="https://upstream.example/mcp",
transport=MCPTransport.sse,
auth_type=MCPAuth.oauth2_token_exchange,
token_exchange_endpoint="https://idp.example.com/token",
client_id="cid",
client_secret="csec",
max_concurrent_requests=max_concurrent,
)
release = asyncio.Event()
inflight = {"current": 0, "peak": 0}
class _ConcurrencyRecordingClient:
async def call_tool(self, params, host_progress_callback=None, raise_on_error=False):
inflight["current"] += 1
inflight["peak"] = max(inflight["peak"], inflight["current"])
try:
await release.wait()
finally:
inflight["current"] -= 1
return CallToolResult(content=[], isError=False)
manager = MCPServerManager()
manager._create_mcp_client = AsyncMock(return_value=_ConcurrencyRecordingClient())
async def _dispatch():
return await manager._call_regular_mcp_tool(
mcp_server=server,
original_tool_name="do_thing",
arguments={},
tasks=[],
mcp_auth_header=None,
mcp_server_auth_headers=None,
oauth2_headers={"Authorization": "Bearer subject-jwt"},
raw_headers=None,
proxy_logging_obj=None,
)
callers = [asyncio.create_task(_dispatch()) for _ in range(max_concurrent + overflow)]
stable = 0
previous = -1
for _ in range(1000):
await asyncio.sleep(0)
current = inflight["current"]
if current == previous:
stable += 1
if current > 0 and stable >= 10:
break
else:
stable = 0
previous = current
peak_while_blocked = inflight["peak"]
release.set()
results = await asyncio.gather(*callers)
assert peak_while_blocked == max_concurrent
assert inflight["current"] == 0
assert all(result.isError is False for result in results)
class TestOBOEndpointDiscovery:
"""An oauth2_token_exchange server with no configured token endpoint discovers it (RFC 9728 ->
RFC 8414) like the oauth2 flow does; an explicitly configured endpoint skips discovery."""

View file

@ -856,6 +856,10 @@ class TestSigV4BuildFromTable:
table_record.tool_name_to_description = None
table_record.byok_api_key_help_url = None
table_record.oauth2_flow = None
table_record.token_exchange_endpoint = None
table_record.audience = None
table_record.subject_token_type = None
table_record.token_exchange_profile = None
table_record.instructions = None
table_record.source_url = None
@ -915,6 +919,10 @@ class TestSigV4BuildFromTable:
table_record.tool_name_to_description = None
table_record.byok_api_key_help_url = None
table_record.oauth2_flow = None
table_record.token_exchange_endpoint = None
table_record.audience = None
table_record.subject_token_type = None
table_record.token_exchange_profile = None
table_record.instructions = None
table_record.source_url = None

View file

@ -789,6 +789,47 @@ class TestCaptureHostProgressCallback:
host.request_context.session = MagicMock()
assert callable(_capture_host_progress_callback(host))
def test_returns_callable_when_token_is_integer(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = 12345
host.request_context.session = MagicMock()
assert callable(_capture_host_progress_callback(host))
def test_returns_callable_when_token_is_zero(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = 0
host.request_context.session = MagicMock()
assert callable(_capture_host_progress_callback(host))
@pytest.mark.asyncio
async def test_forwarded_progress_token_preserves_integer_value(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = 12345
session = AsyncMock()
host.request_context.session = session
callback = _capture_host_progress_callback(host)
assert callback is not None
await callback(0.5, 1.0)
session.send_progress_notification.assert_awaited_once_with(
progress_token=12345,
progress=0.5,
total=1.0,
)
class TestHandleListToolsVirtual:
"""Covers the protocol list_tools early-return when the flag is enabled."""

View file

@ -516,3 +516,104 @@ async def test_full_hot_path_network_count():
assert (
summary["total_network_requests"] == 4
), f"Expected 4 total network requests on warm path, got {summary['total_network_requests']}"
# ============================================================================
# TEST: negative caching for entities that do not exist in the DB
# ============================================================================
@pytest.mark.asyncio
async def test_get_user_object_missing_user_negative_cache():
"""
A user_id with no DB row (e.g. the master key's default admin user_id)
must not trigger a DB query on every request. The first lookup hits the
DB; repeat lookups inside the db_cache_expiry window are throttled.
"""
user_id = "user-missing-negative-cache"
cache = DualCache(in_memory_cache=InMemoryCache())
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_usertable = MagicMock()
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
for _ in range(3):
with pytest.raises(ValueError):
await get_user_object(
user_id=user_id,
prisma_client=mock_prisma,
user_api_key_cache=cache,
parent_otel_span=None,
proxy_logging_obj=None,
user_id_upsert=False,
)
assert mock_prisma.db.litellm_usertable.find_unique.call_count == 1
@pytest.mark.asyncio
async def test_get_user_object_missing_user_rechecks_after_expiry():
"""
The negative cache must expire: a user created after a miss becomes
visible once the db_cache_expiry window has passed.
"""
from litellm.proxy.auth.auth_checks import db_cache_expiry, last_db_access_time
user_id = "user-missing-expiry-recheck"
cache = DualCache(in_memory_cache=InMemoryCache())
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_usertable = MagicMock()
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
with pytest.raises(ValueError):
await get_user_object(
user_id=user_id,
prisma_client=mock_prisma,
user_api_key_cache=cache,
parent_otel_span=None,
proxy_logging_obj=None,
user_id_upsert=False,
)
assert mock_prisma.db.litellm_usertable.find_unique.call_count == 1
last_db_access_time[f"user_id:{user_id}"] = (
None,
time.time() - (db_cache_expiry + 1),
)
with pytest.raises(ValueError):
await get_user_object(
user_id=user_id,
prisma_client=mock_prisma,
user_api_key_cache=cache,
parent_otel_span=None,
proxy_logging_obj=None,
user_id_upsert=False,
)
assert mock_prisma.db.litellm_usertable.find_unique.call_count == 2
def test_should_check_db_negative_entry_throttles_then_expires():
"""
A recorded miss (value=None) suppresses DB checks inside the expiry
window and allows them again after it. Exercises the timestamp element
of the stored (value, time) tuple directly.
"""
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.proxy.auth.auth_checks import (
_should_check_db,
_update_last_db_access_time,
)
tracker: LimitedSizeOrderedDict = LimitedSizeOrderedDict(max_size=10)
_update_last_db_access_time(key="k", value=None, last_db_access_time=tracker)
assert _should_check_db(key="k", last_db_access_time=tracker, db_cache_expiry=5) is False
tracker["k"] = (None, time.time() - 6)
assert _should_check_db(key="k", last_db_access_time=tracker, db_cache_expiry=5) is True

View file

@ -19,6 +19,7 @@ from litellm.proxy.auth.auth_utils import (
get_key_mcp_rpm_limit,
get_key_model_rpm_limit,
get_key_model_tpm_limit,
get_key_tag_rpm_limit,
get_model_from_request,
get_project_model_rpm_limit,
get_project_model_tpm_limit,
@ -2393,3 +2394,17 @@ class TestIsRequestBodySafeBlocksModelList:
)
is True
)
class TestGetKeyTagRateLimits:
"""Tests for get_key_tag_rpm_limit."""
def test_reads_tag_rpm_limit_from_metadata(self):
key = UserAPIKeyAuth(
api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}}
)
assert get_key_tag_rpm_limit(key) == {"cell-1": 5}
def test_returns_none_when_unset(self):
key = UserAPIKeyAuth(api_key="sk-123")
assert get_key_tag_rpm_limit(key) is None

View file

@ -3573,3 +3573,136 @@ async def test_pre_call_hook_skips_reservation_when_disabled(monkeypatch):
)
assert TPM_RESERVED_TOKENS_KEY not in (data.get("metadata") or {})
@pytest.mark.asyncio
async def test_per_tag_rate_limit_independent_counters_v3(monkeypatch):
"""
A single key with per-tag RPM limits tracks each tag independently: a tag
at its limit returns 429 while a different (unlimited) tag keeps flowing,
governed only by the generous key-level limit.
"""
monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60")
_api_key = hash_token("sk-per-tag-rpm")
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
rpm_limit=100,
metadata={"tag_rpm_limit": {"cell-1": 2}},
)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
async def call(tag: str) -> None:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"model": "gpt-3.5-turbo", "metadata": {"tags": [tag]}},
call_type="",
)
await call("cell-1")
await call("cell-1")
with pytest.raises(HTTPException) as exc_info:
await call("cell-1")
assert exc_info.value.status_code == 429
assert "tag_per_key" in str(exc_info.value.detail)
# cell-2 has no configured tag limit, so cell-1's exhausted counter must
# not block it; only the generous key-level limit applies.
for _ in range(5):
await call("cell-2")
@pytest.mark.asyncio
async def test_per_tag_descriptor_creation_v3():
"""
_create_rate_limit_descriptors emits a tag_per_key descriptor carrying the
configured RPM limit only for request tags present in the configured map.
"""
_api_key = hash_token("sk-per-tag-desc")
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
metadata={"tag_rpm_limit": {"cell-1": 5}},
)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = handler._create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1", "cell-2"]}},
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
tag_descriptors = [d for d in descriptors if d["key"] == "tag_per_key"]
assert len(tag_descriptors) == 1, "only the configured tag yields a descriptor"
descriptor = tag_descriptors[0]
assert descriptor["value"] == f"{_api_key}:cell-1"
assert descriptor["rate_limit"]["requests_per_unit"] == 5
@pytest.mark.asyncio
async def test_per_tag_descriptor_absent_without_config_v3():
"""No tag_per_key descriptor is created when the key has no tag limits."""
user_api_key_dict = UserAPIKeyAuth(
api_key=hash_token("sk-no-tag"),
rpm_limit=10,
)
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
descriptors = handler._create_rate_limit_descriptors(
user_api_key_dict=user_api_key_dict,
data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1"]}},
rpm_limit_type=None,
tpm_limit_type=None,
model_has_failures=False,
)
assert not [d for d in descriptors if d["key"] == "tag_per_key"]
@pytest.mark.asyncio
async def test_per_tag_untagged_request_governed_by_key_limit_v3(monkeypatch):
"""
Per-tag limits are opt-in sub-limits under the key-level ceiling, not a
standalone enforcement boundary: a request that carries no tag (or a tag
without a configured limit) is not rejected by any tag counter, but it is
still bounded by the key-level rpm_limit. This pins the documented
untagged-fallback behavior so a future "fail closed on missing tag" change
would fail here instead of silently breaking it.
"""
monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60")
_api_key = hash_token("sk-untagged-fallback")
user_api_key_dict = UserAPIKeyAuth(
api_key=_api_key,
rpm_limit=3,
metadata={"tag_rpm_limit": {"cell-1": 2}},
)
local_cache = DualCache()
handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
async def call(metadata: dict) -> None:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"model": "gpt-3.5-turbo", "metadata": metadata},
call_type="",
)
# Untagged and unconfigured-tag requests share the key-level budget of 3
# and never hit a tag_per_key counter.
await call({})
await call({"tags": ["cell-99"]})
await call({})
with pytest.raises(HTTPException) as exc_info:
await call({"tags": ["cell-99"]})
assert exc_info.value.status_code == 429
assert "tag_per_key" not in str(exc_info.value.detail)

View file

@ -15,8 +15,11 @@ from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import HTTPException
import inspect
from litellm.proxy._types import (
GenerateKeyRequest,
NewUserRequest,
LiteLLM_BudgetTable,
LiteLLM_OrganizationTable,
LiteLLM_TeamTableCachedObj,
@ -14480,3 +14483,23 @@ async def test_regenerate_key_non_admin_permissions_rejected_before_enterprise_g
assert int(exc.value.code) == 403
assert "permissions" in str(exc.value.message)
assert "Enterprise" not in str(exc.value.message)
def test_generate_key_helper_fn_accepts_per_tag_rate_limits():
"""
Regression: new_user / SSO sign-in forward NewUserRequest fields to
generate_key_helper_fn via `**data_json`. The per-tag limit field must be
an accepted kwarg, otherwise user creation 500s with
"generate_key_helper_fn() got an unexpected keyword argument 'tag_rpm_limit'".
"""
params = inspect.signature(generate_key_helper_fn).parameters
assert "tag_rpm_limit" in params
# The field exists on the request model that new_user forwards via **data_json.
assert "tag_rpm_limit" in NewUserRequest.model_fields
# Binding the per-tag kwarg must not raise an unexpected-keyword TypeError.
inspect.signature(generate_key_helper_fn).bind_partial(
request_type="user",
tag_rpm_limit={"cell-1": 5},
)

View file

@ -7131,3 +7131,54 @@ async def test_legacy_login_page_hides_credentials_hint_via_general_settings():
assert response.status_code == 200
assert "Default Credentials" not in body
assert "MASTER_KEY" not in body
@pytest.mark.asyncio
async def test_cli_poll_key_tolerates_missing_user_row():
"""The CLI poll must still mint the JWT when the user lookup raises,
e.g. the user row was created moments ago and a negative-cache window
from the pre-creation SSO existence check is still active on this pod."""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
session_key = "cli-session-missing-user"
session_data = {
"user_id": "just-created-user",
"user_role": "internal_user",
"teams": [],
"models": ["gpt-4"],
}
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.missing.user"
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch("litellm.proxy.proxy_server.prisma_client"),
patch(
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
return_value=mock_jwt_token,
),
patch(
"litellm.proxy.auth.auth_checks.get_user_object",
new=AsyncMock(side_effect=ValueError("User doesn't exist in db. 'user_id'=just-created-user")),
),
):
result = await cli_poll_key(
key_id=session_key,
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result["status"] == "ready"
assert result["key"] == mock_jwt_token
assert result["user_id"] == "just-created-user"

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