mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'BerriAI:litellm_internal_staging' into feature/improve-gigachat-provider
This commit is contained in:
commit
7a73138877
126 changed files with 6597 additions and 1031 deletions
|
|
@ -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
|
||||
|
|
|
|||
4
.github/workflows/test-unit-core-utils.yml
vendored
4
.github/workflows/test-unit-core-utils.yml
vendored
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
4
.github/workflows/test-unit-integrations.yml
vendored
4
.github/workflows/test-unit-integrations.yml
vendored
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
4
.github/workflows/test-unit-misc.yml
vendored
4
.github/workflows/test-unit-misc.yml
vendored
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
4
.github/workflows/test-unit-proxy-auth.yml
vendored
4
.github/workflows/test-unit-proxy-auth.yml
vendored
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
4
.github/workflows/test-unit-proxy-db.yml
vendored
4
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -5,6 +5,10 @@ on:
|
|||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
4
.github/workflows/test-unit-proxy-infra.yml
vendored
4
.github/workflows/test-unit-proxy-infra.yml
vendored
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
4
.github/workflows/test-unit-proxy-legacy.yml
vendored
4
.github/workflows/test-unit-proxy-legacy.yml
vendored
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
|
|
@ -7,6 +7,10 @@ on:
|
|||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths-ignore:
|
||||
- "ui/**"
|
||||
- "**.md"
|
||||
- "**.mdx"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
23
.github/workflows/test_server_root_path.yml
vendored
23
.github/workflows/test_server_root_path.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "token_exchange_profile" TEXT;
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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-*
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
72
tests/e2e/coverage_registry/README.md
Normal file
72
tests/e2e/coverage_registry/README.md
Normal 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
|
||||
8
tests/e2e/coverage_registry/__init__.py
Normal file
8
tests/e2e/coverage_registry/__init__.py
Normal 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.
|
||||
"""
|
||||
280
tests/e2e/coverage_registry/collector.py
Normal file
280
tests/e2e/coverage_registry/collector.py
Normal 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())
|
||||
29
tests/e2e/coverage_registry/guardrail.yaml
Normal file
29
tests/e2e/coverage_registry/guardrail.yaml
Normal 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"}
|
||||
54
tests/e2e/coverage_registry/llm_conversational.yaml
Normal file
54
tests/e2e/coverage_registry/llm_conversational.yaml
Normal 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"}
|
||||
45
tests/e2e/coverage_registry/llm_nonconversational.yaml
Normal file
45
tests/e2e/coverage_registry/llm_nonconversational.yaml
Normal 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)"}
|
||||
25
tests/e2e/coverage_registry/logging.yaml
Normal file
25
tests/e2e/coverage_registry/logging.yaml
Normal 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"}
|
||||
113
tests/e2e/coverage_registry/mcp.yaml
Normal file
113
tests/e2e/coverage_registry/mcp.yaml
Normal 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
|
||||
68
tests/e2e/coverage_registry/mgmt.yaml
Normal file
68
tests/e2e/coverage_registry/mgmt.yaml
Normal 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)"}
|
||||
28
tests/e2e/coverage_registry/other.yaml
Normal file
28
tests/e2e/coverage_registry/other.yaml
Normal 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"}
|
||||
26
tests/e2e/coverage_registry/registry.py
Normal file
26
tests/e2e/coverage_registry/registry.py
Normal 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
|
||||
30
tests/e2e/coverage_registry/reliability.yaml
Normal file
30
tests/e2e/coverage_registry/reliability.yaml
Normal 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"}
|
||||
167
tests/e2e/coverage_registry/schema.py
Normal file
167
tests/e2e/coverage_registry/schema.py
Normal 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]
|
||||
168
tests/e2e/coverage_registry/test_collector.py
Normal file
168
tests/e2e/coverage_registry/test_collector.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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?"),
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
161
tests/e2e/management/test_key_models_dropdown_e2e.py
Normal file
161
tests/e2e/management/test_key_models_dropdown_e2e.py
Normal 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}"
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
)
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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
Loading…
Add table
Reference in a new issue