fix(mcp): preserve current bridge auth in configuration scope

This commit is contained in:
Joshua Valluru 2026-09-22 18:16:06 -07:00
commit b0f1eb656c
139 changed files with 9189 additions and 1791 deletions

View file

@ -11,6 +11,7 @@ UNSUPPORTED: Final = re.compile(
r"|^tests/e2e/guardrails/test_presidio_masking_e2e\.py$"
r"|^tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e\.py$"
r"|^tests/e2e/logging/test_langsmith_batch_serialization_e2e\.py$"
r"|^tests/e2e/secret_manager/"
)
HARNESS: Final = re.compile(
r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$"

View file

@ -4,7 +4,13 @@ on:
workflow_call:
inputs:
test-path:
description: "Pytest path(s) to run"
description: >-
Space-separated pytest paths to run. A path that no longer exists is
dropped with a warning instead of being passed to pytest, because one
missing path makes pytest-xdist collect nothing and report exit 5, which
the step treats as a drained shard. Options are passed through as
written, so use the `--flag=value` form: a bare `--ignore path` would
have its path existence-checked like any other token.
required: true
type: string
workers:
@ -165,14 +171,22 @@ jobs:
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
run: |
found_path=false
for path in ${TEST_PATH}; do
if [ -e "${path%%::*}" ]; then
found_path=true
break
fi
pytest_args=()
existing_paths=0
for token in ${TEST_PATH:?}; do
case "${token}" in
-*) pytest_args+=("${token}") ;;
*)
if [ -e "${token%%::*}" ]; then
pytest_args+=("${token}")
existing_paths=$((existing_paths + 1))
else
echo "::warning::${token} does not exist; drop it from this shard's test-path"
fi
;;
esac
done
if [ "$found_path" = false ]; then
if [ "${existing_paths}" -eq 0 ]; then
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
exit 0
fi
@ -181,7 +195,7 @@ jobs:
xdist_args=(-n "${WORKERS}" --dist="${DIST}")
fi
set +e
uv run --no-sync pytest ${TEST_PATH:?} \
uv run --no-sync pytest "${pytest_args[@]}" \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
"${xdist_args[@]}" \

View file

@ -107,26 +107,18 @@ jobs:
tests/test_litellm/batches
tests/test_litellm/secret_managers
tests/test_litellm/a2a_protocol
tests/test_litellm/anthropic_interface
tests/test_litellm/chat_completions
tests/test_litellm/completion_extras
tests/test_litellm/compression
tests/test_litellm/containers
tests/test_litellm/endpoints
tests/test_litellm/models
tests/test_litellm/repositories
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/messages
tests/test_litellm/ocr
tests/test_litellm/passthrough
tests/test_litellm/rag
tests/test_litellm/realtime_api
tests/test_litellm/rerank_api
tests/test_litellm/rust_bridge
tests/test_litellm/sandbox
tests/test_litellm/skills
tests/test_litellm/test_router
tests/test_litellm/vector_stores
tests/test_litellm/videos
tests/test_litellm/test_*.py

View file

@ -277,7 +277,7 @@ def get_s3_object_key(
start_time: datetime,
s3_file_name: str,
) -> str:
sanitized_s3_file_name: Final = s3_file_name.replace("/", "_")
sanitized_s3_file_name: Final = s3_file_name.replace("/", "_").replace(":", "_")
configured_prefix: Final = (s3_path.rstrip("/") + "/" if s3_path else "") + prefix
date_segment: Final = start_time.strftime("%Y-%m-%d") + "/"
# we need the s3 key to include the time, so we log cache hits too

View file

@ -3,10 +3,30 @@ from typing import Final
import litellm
from litellm import verbose_logger
from ...litellm_core_utils.get_llm_provider_logic import get_llm_provider
from ...litellm_core_utils.get_llm_provider_logic import (
declared_authenticating_provider,
get_llm_provider,
)
from ...types.router import LiteLLM_Params
def _api_base_without_login(provider: str) -> str | None:
if provider == "github_copilot":
return litellm.GithubCopilotConfig().api_base_without_login()
if provider == "chatgpt":
return litellm.ChatGPTConfig().api_base_without_login()
return None
def _provider_default_api_base(model: str, custom_llm_provider: str | None, stream: bool) -> str | None:
if custom_llm_provider == "gemini":
action: Final = "streamGenerateContent" if stream else "generateContent"
return f"https://generativelanguage.googleapis.com/v1beta/models/{model}:{action}"
if custom_llm_provider == "openai":
return "https://api.openai.com"
return None
def get_api_base(model: str, optional_params: dict | LiteLLM_Params) -> str | None:
"""
Returns the api base used for calling the model.
@ -42,6 +62,9 @@ def get_api_base(model: str, optional_params: dict | LiteLLM_Params) -> str | No
if litellm.model_alias_map and model in litellm.model_alias_map:
model = litellm.model_alias_map[model]
declared: Final = declared_authenticating_provider(model, _optional_params.custom_llm_provider)
if declared is not None:
return _api_base_without_login(declared)
try:
(
model,
@ -83,16 +106,4 @@ def get_api_base(model: str, optional_params: dict | LiteLLM_Params) -> str | No
_api_base = f"{_optional_params.vertex_location}-aiplatform.googleapis.com/v1/projects/{_optional_params.vertex_project}/locations/{_optional_params.vertex_location}/publishers/google/models/{model}:generateContent"
return _api_base
if custom_llm_provider is None:
return None
if custom_llm_provider == "gemini":
if stream:
_api_base = f"https://generativelanguage.googleapis.com/v1beta/models/{model}:streamGenerateContent"
else:
_api_base = f"https://generativelanguage.googleapis.com/v1beta/models/{model}:generateContent"
return _api_base
elif custom_llm_provider == "openai":
_api_base = "https://api.openai.com"
return _api_base
return None
return _provider_default_api_base(model, custom_llm_provider, stream)

View file

@ -1593,7 +1593,7 @@ def _resolve_s3_setting(
source.get(param_name) for source in (litellm_params, optional_params) if source is not None
)
explicit: Final = next((value for value in candidates if isinstance(value, str) and value), None)
return explicit or get_secret_str(env_var)
return explicit or get_secret_str(env_var) or None
class CommonBatchFilesUtils:

View file

@ -23,6 +23,9 @@ class ChatGPTConfig(OpenAIConfig):
super().__init__()
self.authenticator = Authenticator()
def api_base_without_login(self) -> str:
return self.authenticator.get_api_base()
def _get_openai_compatible_provider_info(
self,
model: str,
@ -30,7 +33,7 @@ class ChatGPTConfig(OpenAIConfig):
api_key: str | None,
custom_llm_provider: str,
) -> tuple[str | None, str | None, str]:
dynamic_api_base: Final = self.authenticator.get_api_base()
dynamic_api_base: Final = self.api_base_without_login()
try:
dynamic_api_key: Final = self.authenticator.get_access_token()
except GetAccessTokenError as e:

View file

@ -31,6 +31,14 @@ class GithubCopilotConfig(OpenAIConfig):
super().__init__()
self.authenticator = Authenticator()
def api_base_without_login(self, api_base: str | None = None) -> str:
return (
api_base
or self.authenticator.get_api_base()
or os.getenv("GITHUB_COPILOT_API_BASE")
or DEFAULT_GITHUB_COPILOT_API_BASE
)
def _get_openai_compatible_provider_info(
self,
model: str,
@ -38,12 +46,7 @@ class GithubCopilotConfig(OpenAIConfig):
api_key: str | None,
custom_llm_provider: str,
) -> tuple[str | None, str | None, str]:
dynamic_api_base: Final = (
api_base
or self.authenticator.get_api_base()
or os.getenv("GITHUB_COPILOT_API_BASE")
or DEFAULT_GITHUB_COPILOT_API_BASE
)
dynamic_api_base: Final = self.api_base_without_login(api_base)
try:
dynamic_api_key: Final = self.authenticator.get_api_key()
except GetAPIKeyError as e:

View file

@ -1,3 +1,5 @@
import base64
import io
from typing import Any, Final
import httpx
@ -11,37 +13,35 @@ class OllamaError(BaseLLMException):
super().__init__(status_code=status_code, message=message, headers=headers)
def _convert_image(image):
"""
Convert image to base64 encoded image if not already in base64 format
_JPEG_AND_PNG_SIGNATURES: Final = (b"\xff\xd8\xff", b"\x89PNG\r\n\x1a\n")
If image is already in base64 format AND is a jpeg/png, return it
If image is not JPEG/PNG, convert it to JPEG base64 format
"""
import base64
import io
def _reencode_as_jpeg(raw_image: bytes, original: str) -> str:
try:
from PIL import Image
except Exception:
raise Exception("ollama image conversion failed please run `pip install Pillow`")
orig: Final = image
if image.startswith("data:"):
image = image.split(",")[-1]
try:
image_data: Final = Image.open(io.BytesIO(base64.b64decode(image)))
if image_data.format in ["JPEG", "PNG"]:
return image
picture: Final = Image.open(io.BytesIO(raw_image))
except Exception:
return orig
return original
jpeg_image: Final = io.BytesIO()
image_data.convert("RGB").save(jpeg_image, "JPEG")
jpeg_image.seek(0)
picture.convert("RGB").save(jpeg_image, "JPEG")
return base64.b64encode(jpeg_image.getvalue()).decode("utf-8")
def _convert_image(image: str) -> str:
payload: Final = image.split(",")[-1] if image.startswith("data:") else image
try:
raw_image: Final = base64.b64decode(payload)
except ValueError:
return image
if raw_image.startswith(_JPEG_AND_PNG_SIGNATURES):
return payload
return _reencode_as_jpeg(raw_image, original=image)
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo

View file

@ -25314,8 +25314,7 @@
"search_context_size_medium": 0.014,
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"cache_read_input_token_cost_batches": 1e-07
"web_search_billing_unit": "per_query"
},
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@ -25401,8 +25400,7 @@
"search_context_size_medium": 0.014,
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"cache_read_input_token_cost_batches": 2.5e-08
"web_search_billing_unit": "per_query"
},
"gemini-3.1-flash-image-preview": {
"input_cost_per_image": 0.00056,
@ -25482,8 +25480,7 @@
"supports_response_schema": false,
"supports_system_messages": true,
"supports_video_input": true,
"supports_vision": true,
"cache_read_input_token_cost_batches": 1.25e-08
"supports_vision": true
},
"gemini-3.1-flash-lite-preview": {
"cache_read_input_token_cost": 2.5e-08,
@ -25593,8 +25590,7 @@
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"input_cost_per_audio_token_batches": 2.5e-07,
"cache_read_input_token_cost_batches": 1.25e-08
"input_cost_per_audio_token_batches": 2.5e-07
},
"gemini-3.5-flash-lite": {
"deprecation_date": "2027-07-21",
@ -25652,8 +25648,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 1.5e-08
"google_maps_grounding_cost_per_query": 0.014
},
"deep-research-pro-preview-12-2025": {
"cache_read_input_token_cost": 2e-07,
@ -25689,8 +25684,7 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_vision": true,
"supports_web_search": true,
"cache_read_input_token_cost_batches": 1e-07
"supports_web_search": true
},
"gemini-2.5-flash-lite": {
"cache_read_input_audio_token_cost": 3e-08,
@ -26321,8 +26315,7 @@
"output_cost_per_token_batches": 4.5e-06,
"input_cost_per_token_flex": 7.5e-07,
"output_cost_per_token_flex": 4.5e-06,
"cache_read_input_token_cost_flex": 7.5e-08,
"cache_read_input_token_cost_batches": 7.5e-08
"cache_read_input_token_cost_flex": 7.5e-08
},
"vertex_ai/gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
@ -26380,8 +26373,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
@ -26440,8 +26432,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/gemini-3.8-flash": {
"prompt_cache_min_tokens": 4096,
@ -26500,8 +26491,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
@ -28250,8 +28240,7 @@
"output_cost_per_token_batches": 4.5e-06,
"input_cost_per_token_flex": 7.5e-07,
"output_cost_per_token_flex": 4.5e-06,
"cache_read_input_token_cost_flex": 7.5e-08,
"cache_read_input_token_cost_batches": 7.5e-08
"cache_read_input_token_cost_flex": 7.5e-08
},
"gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
@ -28309,8 +28298,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
@ -28369,8 +28357,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"gemini-3.8-flash": {
"prompt_cache_min_tokens": 4096,
@ -28429,8 +28416,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"gemini/gemini-2.5-pro-preview-tts": {
"cache_read_input_token_cost": 1.25e-07,
@ -33110,8 +33096,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true,
"cache_read_input_token_cost_batches": 6.25e-08
"supports_minimal_reasoning_effort": true
},
"gpt-5-chat": {
"cache_read_input_token_cost": 1.25e-07,
@ -33292,8 +33277,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true,
"cache_read_input_token_cost_batches": 1.25e-08
"supports_minimal_reasoning_effort": true
},
"gpt-5-nano": {
"cache_read_input_token_cost": 5e-09,
@ -33392,8 +33376,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true,
"cache_read_input_token_cost_batches": 2.5e-09
"supports_minimal_reasoning_effort": true
},
"gpt-image-1": {
"cache_read_input_token_cost": 1.25e-06,
@ -47531,8 +47514,7 @@
"output_cost_per_token_flex": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
"supports_reasoning": false,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"cache_read_input_token_cost_batches": 1e-07
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@ -47570,8 +47552,7 @@
"output_cost_per_token_batches": 1.5e-06,
"output_cost_per_token_flex": 1.5e-06,
"supports_reasoning": false,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"cache_read_input_token_cost_batches": 2.5e-08
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/gemini-3.1-flash-image-preview": {
"input_cost_per_image": 0.00056,
@ -47627,8 +47608,7 @@
"supports_response_schema": false,
"supports_system_messages": true,
"supports_video_input": true,
"supports_vision": true,
"cache_read_input_token_cost_batches": 1.25e-08
"supports_vision": true
},
"vertex_ai/gemini-3.1-flash-lite-preview": {
"cache_read_input_token_cost": 2.5e-08,
@ -47739,8 +47719,7 @@
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"input_cost_per_audio_token_batches": 2.5e-07,
"cache_read_input_token_cost_batches": 1.25e-08
"input_cost_per_audio_token_batches": 2.5e-07
},
"vertex_ai/gemini-3.5-flash-lite": {
"deprecation_date": "2027-07-21",
@ -47799,8 +47778,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 1.5e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/deep-research-pro-preview-12-2025": {
"cache_read_input_token_cost": 2e-07,
@ -47817,8 +47795,7 @@
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_batches": 6e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"cache_read_input_token_cost_batches": 1e-07
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/jamba-1.5": {
"input_cost_per_token": 2e-07,

View file

@ -43,6 +43,7 @@ from litellm.proxy._types import (
SpecialMCPServerName,
SpecialMCPServerNames,
UserAPIKeyAuth,
hash_token,
user_api_key_has_admin_view,
)
from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
@ -183,6 +184,20 @@ def _is_litellm_auth_admission_error(exc: Exception) -> bool:
return False
def _explicit_credential_matches_envelope(
explicit_auth: UserAPIKeyAuth,
presented_token: str,
identity: EnvelopeIdentity,
) -> bool:
"""Match the stored key hash or user ID, including token-only mapped JWT keys."""
match identity.subject_type:
case "key_hash":
return identity.subject in (hash_token(presented_token), explicit_auth.token)
case "user_id":
return explicit_auth.user_id is not None and explicit_auth.user_id == identity.subject
return assert_never(identity.subject_type)
def _has_client_supplied_mcp_auth(
mcp_auth_header: str | None,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
@ -479,6 +494,31 @@ class MCPRequestHandler:
# Only OAuth metadata routes registered under /.well-known/ are public.
if request_route.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
has_explicit_litellm_key
and oauth2_headers
and is_bridge_envelope_shaped(oauth2_headers["Authorization"])
and (
dual_bridge_target := MCPRequestHandler._single_dcr_bridge_delegate_target(
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
)
is not None
):
(
validated_user_api_key_auth,
mcp_server_auth_headers,
) = await MCPRequestHandler._admit_dcr_bridge_dual_credential(
server=dual_bridge_target.server,
requested_name=dual_bridge_target.requested_name,
authorization_value=oauth2_headers["Authorization"],
litellm_api_key=litellm_api_key,
mcp_server_auth_headers=mcp_server_auth_headers,
request=request,
route=request_route,
)
elif has_explicit_litellm_key:
# An explicit x-litellm-api-key is always a LiteLLM credential, even
# for a delegated server, so validate it: identity / spend / rate
@ -786,6 +826,45 @@ class MCPRequestHandler:
higher-priority alias slot, pairing the admitted identity with an attacker's upstream
credential; the alias-keyed injection overwrites any such caller value.
"""
result: Final = await MCPRequestHandler._open_dcr_bridge_envelope(
server=server,
requested_name=requested_name,
authorization_value=authorization_value,
request=request,
route=route,
)
header_key: Final = server.alias or server.server_name
if header_key is None:
raise HTTPException(status_code=500, detail="Server misconfigured: MCP server has no routable name")
admitted: Final = await MCPRequestHandler._reload_admitted_principal(result.identity)
await MCPRequestHandler._enforce_admitted_live_policy(admitted=admitted, request=request, route=route)
injected: Final = { # mutable-ok: mcp_server_auth_headers contract requires concrete dicts
header_key: { # mutable-ok: concrete dict header payload
"Authorization": result.upstream_authorization.get_secret_value()
}
}
new_headers: Final = { # mutable-ok: merged header map must stay a concrete dict
**(mcp_server_auth_headers or {}), # mutable-ok: empty-dict fallback for the merge
**injected,
}
return admitted, new_headers
@staticmethod
async def _open_dcr_bridge_envelope(
server: MCPServer,
requested_name: str,
authorization_value: str,
request: Request,
route: str,
) -> BridgeEnvelopeAdmitted:
"""Open a bridge envelope after the pre-DB gates, or fail closed with the scope's challenge.
Shared by the envelope-only arm (:meth:`_admit_dcr_bridge_delegate`) and the dual-credential
arm (:meth:`_admit_dcr_bridge_dual_credential`): both require master_key, run the same
proxy-wide pre-DB checks the standard pipeline applies before any key lookup, and resolve
the envelope's crypto. Returns only the ``BridgeEnvelopeAdmitted`` result; an invalid,
expired, tampered, or non-envelope value raises the requested scope's ``invalid_token``
challenge instead."""
from litellm.proxy.proxy_server import master_key
if not master_key:
@ -797,20 +876,67 @@ class MCPRequestHandler:
result: Final = resolve_bridge_envelope(authorization_value, keys, datetime.now(timezone.utc), server.server_id)
match result:
case BridgeEnvelopeAdmitted():
header_key: Final = server.alias or server.server_name
if header_key is None:
raise HTTPException(status_code=500, detail="Server misconfigured: MCP server has no routable name")
admitted: Final = await MCPRequestHandler._reload_admitted_principal(result.identity)
await MCPRequestHandler._enforce_admitted_live_policy(admitted=admitted, request=request, route=route)
injected: Final = {header_key: {"Authorization": result.upstream_authorization.get_secret_value()}}
new_headers: Final = {**(mcp_server_auth_headers or {}), **injected}
return admitted, new_headers
return result
case BridgeEnvelopeInvalid() | NotBridgeEnvelope():
raise MCPRequestHandler._dcr_bridge_invalid_token_challenge(
requested_name=requested_name, request=request
)
case _:
assert_never(result)
return assert_never(result)
@staticmethod
async def _admit_dcr_bridge_dual_credential(
server: MCPServer,
requested_name: str,
authorization_value: str,
litellm_api_key: str,
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
request: Request,
route: str,
) -> tuple[UserAPIKeyAuth, dict[str, dict[str, str]] | None]:
"""Admit a request carrying BOTH an explicit litellm credential and a bridge envelope.
MCP clients send ``x-litellm-api-key`` on every request, including the ``tools/list`` that
follows the ``/{server}/token`` mint, so the envelope arrives alongside the key rather than
alone. The explicit credential is validated first (its own pipeline, so a bad key keeps the
normal 401/403), then the envelope is opened and its sealed identity must match the explicit
credential's principal — a mismatch is a 403, never a fallback onto either credential alone.
On a match the explicit credential's ``UserAPIKeyAuth`` is the admission context (key
permissions, budgets, rate limits) and the sealed upstream token is injected under the
server's per-server auth-header key, while the leak-defense chokepoint strips the envelope
``Authorization`` itself from egress."""
presented_token: Final = _get_bearer_token_or_received_api_key(litellm_api_key)
explicit_auth: Final = await user_api_key_auth(api_key=f"Bearer {presented_token}", request=request)
result: Final = await MCPRequestHandler._open_dcr_bridge_envelope(
server=server,
requested_name=requested_name,
authorization_value=authorization_value,
request=request,
route=route,
)
if not _explicit_credential_matches_envelope(
explicit_auth=explicit_auth,
presented_token=presented_token,
identity=result.identity,
):
raise HTTPException(
status_code=403,
detail={ # mutable-ok: HTTPException detail payload requires a concrete dict
"error": "oauth_principal_mismatch"
},
)
header_key: Final = server.alias or server.server_name
if header_key is None:
raise HTTPException(status_code=500, detail="Server misconfigured: MCP server has no routable name")
injected: Final = { # mutable-ok: mcp_server_auth_headers contract requires concrete dicts
header_key: { # mutable-ok: concrete dict header payload
"Authorization": result.upstream_authorization.get_secret_value()
}
}
new_headers: Final = { # mutable-ok: merged header map must stay a concrete dict
**(mcp_server_auth_headers or {}), # mutable-ok: empty-dict fallback for the merge
**injected,
}
return explicit_auth, new_headers
@staticmethod
async def _admit_dcr_bridge_authorization(
@ -1095,9 +1221,15 @@ class MCPRequestHandler:
project, org, and budget state are NOT re-checked here; the caller runs the admitted
identity through ``_enforce_admitted_live_policy`` for those.
"""
from litellm.proxy._experimental.mcp_server.bridge_token_flow import (
master_key_admin_auth, # noqa: PLC0415 # inline import avoids a module-load circular import
)
from litellm.proxy.auth.auth_checks import get_key_object
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
admin: Final = master_key_admin_auth(key_hash)
if admin is not None:
return admin
if prisma_client is None:
raise HTTPException(status_code=500, detail="Server misconfigured: no database connection")
try:

View file

@ -209,6 +209,31 @@ async def _resolve_active_litellm_key(request: Request) -> "_ResolvedKey | _KeyR
return await _reload_active_key_by_hash(hash_token(token))
def master_key_admin_auth(key_hash: str) -> "UserAPIKeyAuth | None":
from litellm.constants import ( # noqa: PLC0415 # inline import avoids a module-load circular import
LITELLM_PROXY_MASTER_KEY_ALIAS,
)
from litellm.proxy._types import ( # noqa: PLC0415 # inline import avoids a module-load circular import
LitellmUserRoles,
UserAPIKeyAuth,
hash_token,
)
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
litellm_proxy_admin_name,
master_key,
)
if not master_key or not secrets.compare_digest(key_hash, hash_token(master_key)):
return None
auth: Final = UserAPIKeyAuth(
api_key=LITELLM_PROXY_MASTER_KEY_ALIAS,
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id=litellm_proxy_admin_name,
)
auth.via_virtual_key = True
return auth
async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResolutionFailure":
"""Reload the live key record for ``key_hash`` (cache first, then DB) and gate it on active state,
returning the resolved key or a precise failure. Shared by the token request's presented-key
@ -233,6 +258,8 @@ async def _reload_active_key_by_hash(key_hash: str) -> "_ResolvedKey | _KeyResol
user_api_key_cache,
)
if (admin := master_key_admin_auth(key_hash)) is not None:
return _ResolvedKey(key_hash=key_hash, key=admin)
if prisma_client is None:
return "unresolvable"
try:
@ -475,7 +502,7 @@ async def _resolve_jwt_auth(
proxy_logging_obj=proxy_logging_obj,
)
if isinstance(mapped, UserAPIKeyAuth):
return None if await _key_owner_scim_deactivated(mapped) or not _active_key_user_id(mapped) else mapped
return None if await _key_owner_scim_deactivated(mapped) or not _key_is_active(mapped) else mapped
if mapped is not None:
return None
if write_route is None:
@ -593,6 +620,7 @@ def _bridge_grant_from_token_response(token_response: object) -> "UpstreamTokenG
_BridgeMintError = Literal[
"no_identity",
"jwt_client_policy_unsupported",
"invalid_refresh",
"identity_unavailable",
"identity_faulted",
@ -633,6 +661,13 @@ def _bridge_mint_error_response(error: _BridgeMintError) -> JSONResponse:
"this server issues a gateway-bound credential; complete the interactive sign-in, or "
"send a litellm credential (x-litellm-api-key or Authorization) on the token request",
)
case "jwt_client_policy_unsupported":
status, code, desc = (
400,
"invalid_request",
"JWT bridge minting is not supported with a claim-based MCP client allowlist; "
"the bridge credential cannot preserve the signed client identity",
)
case "invalid_refresh":
status, code, desc = (
400,
@ -736,11 +771,18 @@ async def _prepare_bridge_mint(
Two identity sources, one envelope. The interactive DCR client authenticates via SSO at the bridged
authorize, so its identity arrives as ``bridge_identity`` (the user recovered from the gateway
authorization code) and mints a user subject. The scripted two-header client presents a litellm key
on the token request instead, so its identity is the active key's hash and mints a key_hash subject.
A missing or invalid presented key keeps its resolution origin so the mapper statuses it truthfully;
neither source present is ``no_identity``. The refresh_token grant has its own phase-1
authorization code) and mints a user subject. The scripted two-header client presents a litellm
credential (a virtual key or a JWT) on the token request instead: a key mints a key_hash subject,
while a JWT resolves through the same auth path as admission and mints a key_hash subject when it
maps to a virtual key. An unmapped JWT is rejected because a user subject cannot preserve its
JWT-specific authorization restrictions. A JWT client-claim allowlist also prevents JWT minting:
the envelope cannot retain the signed client identity for subsequent allowlist checks. A missing or invalid
presented key keeps its resolution origin so the mapper statuses it truthfully; neither source
present is ``no_identity``. The refresh_token grant has its own phase-1
(:func:`_prepare_bridge_refresh`), which recovers identity from the presented refresh envelope."""
from litellm.proxy._experimental.mcp_server.client_allowlist import ( # noqa: PLC0415 # keep mint policy dependencies local
load_mcp_client_allowlist,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( # noqa: PLC0415 # inline import avoids a module-load circular import
envelope_keys_from_master_key,
)
@ -748,7 +790,10 @@ async def _prepare_bridge_mint(
key_hash_identity,
user_identity,
)
from litellm.proxy._types import UserAPIKeyAuth # noqa: PLC0415 # proxy import cycle
from litellm.proxy.auth.handle_jwt import JWTHandler # noqa: PLC0415 # proxy import cycle
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # inline import avoids a module-load circular import
general_settings,
master_key,
)
@ -758,6 +803,16 @@ async def _prepare_bridge_mint(
if bridge_identity is not None:
identity = user_identity(server_id=mcp_server.server_id, user_id=bridge_identity.litellm_user_id)
return _BridgeMintReady(identity=identity, keys=keys)
presented_token: Final = _litellm_key_from_request(request)
if presented_token is not None and JWTHandler.is_jwt(presented_token):
client_allowlist: Final = load_mcp_client_allowlist(general_settings)
if client_allowlist is not None and client_allowlist.jwt_field is not None:
return "jwt_client_policy_unsupported"
resolved_jwt: Final = await _resolve_jwt_auth(request, presented_token, None)
if isinstance(resolved_jwt, UserAPIKeyAuth) and resolved_jwt.token:
identity = key_hash_identity(server_id=mcp_server.server_id, key_hash=resolved_jwt.token)
return _BridgeMintReady(identity=identity, keys=keys)
return "no_identity"
resolved: Final = await _resolve_active_litellm_key(request)
if not isinstance(resolved, _ResolvedKey):
return _key_resolution_failure_to_mint_error(resolved)

View file

@ -937,6 +937,7 @@ class LiteLLMRoutes(enum.Enum):
# proxy admin, or team admin naming their own team via team_id
"/auto_router/test_routing",
"/auto_router/validate_complexity_router_config",
"/auto_router/availability",
# Per-session auto-router read - the endpoint scopes the row to the caller's own key hash
"/auto_router/session",
"/cost/predict-cache",

View file

@ -1606,13 +1606,21 @@ class DBSpendUpdateWriter:
proxy_logging_obj: ProxyLogging,
) -> None:
transactions: Final = await queue.flush_and_get_aggregated_daily_spend_update_transactions()
try:
await commit(
commit_task: Final = asyncio.ensure_future(
commit(
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=cast(dict[str, _DailySpendTransactionT], transactions),
)
)
try:
await asyncio.shield(commit_task)
except asyncio.CancelledError:
commit_task.cancel()
if transactions:
await queue.add_update(transactions)
raise
except Exception as e: # noqa: BLE001 # whatever failed here, the other tables must still flush
if not transactions:
return
@ -1839,14 +1847,18 @@ class DBSpendUpdateWriter:
if not daily_tag_spend_update_transactions:
return
try:
await DBSpendUpdateWriter.update_daily_tag_spend(
commit_task: Final = asyncio.ensure_future(
DBSpendUpdateWriter.update_daily_tag_spend(
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_tag_spend_update_transactions,
)
except Exception:
)
try:
await asyncio.shield(commit_task)
except BaseException: # noqa: BLE001 # a cancel must restore the drained rows before its rollback returns
commit_task.cancel()
await self.redis_update_buffer.restore_transactions_to_redis(
daily_tag_spend_update_transactions=daily_tag_spend_update_transactions,
)
@ -2368,7 +2380,8 @@ class DBSpendUpdateWriter:
table=table, transactions=tuple(transactions_to_process.values())
)
sql, params = build_bulk_upsert(table=table, batch=merged_batch)
await prisma_client.db.execute_raw(sql, *params)
async with _spend_update_tx(prisma_client) as transaction:
await transaction.execute_raw(sql, *params)
except Exception as batch_error:
if _spend_commit_failure_is_requeue_safe(batch_error):
spend_log_error(

View file

@ -58,6 +58,8 @@ from litellm.router_utils.auto_router_model_naming import (
)
from litellm.types.management_endpoints.auto_router_endpoints import (
SHADOW_EVAL_TURN_VALVE,
AutoRouterAvailabilityRequest,
AutoRouterAvailabilityResponse,
AutoRouterBenchmarkGroup,
AutoRouterBenchmarksResponse,
AutoRouterBenchmarkTotals,
@ -391,6 +393,54 @@ async def validate_complexity_router_config(
return ComplexityRouterConfigValidationResponse(valid=error is None, error=error)
@router.post(
"/auto_router/availability",
tags=["model management"], # mutable-ok: FastAPI requires a list
response_model=AutoRouterAvailabilityResponse,
)
async def get_auto_router_availability(
data: AutoRouterAvailabilityRequest,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> AutoRouterAvailabilityResponse:
from litellm.proxy.management_helpers.auto_router_availability import auto_router_availability
from litellm.proxy.proxy_server import (
_license_check, # pyright: ignore[reportPrivateUsage] # same entitlement owner as the model write gate
heuristic_v1_tuning_baselines,
llm_router,
proxy_config,
)
member_team: Final = await _authorize_router_dry_run(user_api_key_dict, data.team_id)
rows: Final = proxy_config.auto_router_db_catalog
if rows is None or llm_router is None:
raise HTTPException(status_code=503, detail="Auto-router availability is unavailable")
saved: Final = next((row for row in rows if row.model_id == data.saved_model_id), None)
if data.saved_model_id is not None:
if saved is None:
raise HTTPException(status_code=404, detail="Saved auto router is unavailable")
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN and (
saved.team_id != data.team_id or (member_team is not None and saved.created_by != user_api_key_dict.user_id)
):
raise HTTPException(status_code=403, detail="Cannot check another user's auto router")
existing: Final = saved.deployment if saved is not None else None
others: Final = tuple(row.deployment for row in rows if row is not saved) + tuple(llm_router.config_deployments())
candidate: Final = MappingProxyType(
{
"litellm_params": MappingProxyType(
{"model": "auto_router/complexity_router", "complexity_router_config": data.complexity_router_config}
),
"model_info": MappingProxyType({"id": data.saved_model_id or "availability-new-router", "db_model": True}),
}
)
return auto_router_availability(
others=others,
existing=existing,
candidate=candidate,
baselines=heuristic_v1_tuning_baselines,
limit=_license_check.auto_router_capability_limit(),
)
async def _resolve_saved_routing_test(
data: AutoRouterRoutingTestRequest,
user_api_key_dict: UserAPIKeyAuth,

View file

@ -1782,10 +1782,11 @@ async def delete_team_models(
# Under MODEL_RECONCILE_LOCK, for the same reason as delete_model: the rows are
# gone, but a reconcile holding a pre-delete snapshot would upsert these ids back
# onto this pod. The lock orders the eviction after any in-flight reconcile.
if llm_router is not None:
from litellm.proxy.proxy_server import MODEL_RECONCILE_LOCK
from litellm.proxy.proxy_server import MODEL_RECONCILE_LOCK, proxy_config
async with MODEL_RECONCILE_LOCK:
async with MODEL_RECONCILE_LOCK:
proxy_config.remove_auto_router_catalog_entries(frozenset(deleted_model_ids))
if llm_router is not None:
for model_id in deleted_model_ids:
llm_router.delete_deployment(id=model_id)
@ -2194,6 +2195,7 @@ async def delete_model(
llm_router,
premium_user,
prisma_client,
proxy_config,
proxy_logging_obj,
store_model_in_db,
user_api_key_cache,
@ -2245,8 +2247,9 @@ async def delete_model(
# this pod serving a model the database no longer has, until the next
# reconcile. Taking the lock orders this eviction after any such in-flight
# reconcile's re-add, so the eviction is the last word.
if llm_router is not None:
async with MODEL_RECONCILE_LOCK:
async with MODEL_RECONCILE_LOCK:
proxy_config.remove_auto_router_catalog_entries(frozenset({model_info.id}))
if llm_router is not None:
llm_router.delete_deployment(id=model_info.id)
# Runs after the row delete so the sibling check sees post-delete state.

View file

@ -0,0 +1,148 @@
from collections.abc import Mapping, Sequence
from copy import deepcopy
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, Json, TypeAdapter, ValidationError
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.router_utils.auto_router_model_naming import (
GATED_AUTO_ROUTER_CAPABILITIES,
capability_limit_violation,
classify_strategy_router_model,
count_capability_routers,
gated_capability_of,
)
from litellm.router_utils.auto_router_tuning_baseline import (
is_mutable_tuned_candidate,
mutable_tuned_identities,
tuning_quota_violation,
)
from litellm.types.management_endpoints.auto_router_endpoints import (
AutoRouterAllowance,
AutoRouterAvailabilityResponse,
)
class _CatalogModelInfo(BaseModel):
team_id: str | None = None
class _CatalogSource(BaseModel):
model_id: str
created_by: str | None = None
litellm_params: Json[dict[str, object]] | dict[str, object]
model_info: Json[_CatalogModelInfo] | _CatalogModelInfo | None = None
@dataclass(frozen=True, slots=True)
class AutoRouterCatalogEntry:
model_id: str
team_id: str | None
created_by: str | None
deployment: Mapping[str, object]
def _catalog_field(value: object, key: str) -> object:
if not isinstance(value, str):
return deepcopy(value)
return decrypt_value_helper(value, key=key, exception_type="debug", return_original_value=True)
def build_auto_router_catalog(rows: Sequence[object]) -> tuple[AutoRouterCatalogEntry, ...] | None:
try:
sources: Final = TypeAdapter(tuple[_CatalogSource, ...]).validate_python(rows, from_attributes=True)
except ValidationError:
return None
return tuple(
AutoRouterCatalogEntry(
model_id=row.model_id,
team_id=row.model_info.team_id if row.model_info is not None else None,
created_by=row.created_by,
deployment=MappingProxyType(
{
"litellm_params": MappingProxyType(
{
"model": model,
"complexity_router_config": _catalog_field(
row.litellm_params.get("complexity_router_config"), "complexity_router_config"
),
}
),
"model_info": MappingProxyType({"id": row.model_id, "db_model": True}),
}
),
)
for row in sources
if isinstance(model := _catalog_field(row.litellm_params.get("model"), "model"), str)
and classify_strategy_router_model(model) == "complexity"
)
def auto_router_availability(
*,
others: Sequence[Mapping[str, object]],
existing: Mapping[str, object] | None,
candidate: Mapping[str, object],
baselines: Mapping[str, str] | None,
limit: int | None,
) -> AutoRouterAvailabilityResponse:
existing_params: Final = None if existing is None else existing.get("litellm_params")
candidate_params: Final = candidate.get("litellm_params")
owned: Final = gated_capability_of(existing_params) if isinstance(existing_params, Mapping) else None
claimed: Final = gated_capability_of(candidate_params) if isinstance(candidate_params, Mapping) else None
counts: Final = tuple(
(capability, count_capability_routers(others, capability=capability))
for capability in GATED_AUTO_ROUTER_CAPABILITIES
)
tuned_count: Final = len(mutable_tuned_identities(others, baselines)) if baselines is not None else 0
allowances: Final = tuple(
AutoRouterAllowance(
key=capability.key,
limit=limit,
remaining=None if limit is None else max(0, limit - held),
used_by_this_router=owned is capability,
)
for capability, held in counts
)
capability_error: Final = next(
(
capability_limit_violation(capability=capability, held=held + 1, limit=limit)
for capability, held in counts
if capability is claimed
),
None,
)
tuning_error: Final = (
tuning_quota_violation(candidate=candidate, others=others, baselines=baselines, limit=limit)
if baselines is not None
else None
)
capability_labels: Final = {
"heuristic_v2": "Heuristic v2",
"capability": "Capability",
"llm_v2": "Fuse v2",
"tier_or_classifier_prompt": "Custom tiers or classifier instructions",
}
return AutoRouterAvailabilityResponse(
allowances=(
*allowances,
AutoRouterAllowance(
key="heuristic_tuning",
limit=limit,
remaining=None if limit is None or baselines is None else max(0, limit - tuned_count),
available=limit is None or baselines is not None,
used_by_this_router=bool(
existing is not None and baselines is not None and is_mutable_tuned_candidate(existing, baselines)
),
),
),
error=(
f"{capability_labels[claimed.key]} has no available allowance. Choose another option or free an existing allowance."
if capability_error is not None and claimed is not None
else "These scoring rules need an available Rule-based tuning allowance. Check the weights, thresholds, keywords, and custom dimensions in Advanced settings. Model choices do not use this allowance."
if tuning_error is not None
else None
),
)

View file

@ -132,6 +132,7 @@ from litellm.proxy.common_utils.callback_utils import (
strip_callback_config,
)
from litellm.proxy.common_utils.realtime_utils import _realtime_request_body
from litellm.proxy.management_helpers.auto_router_availability import AutoRouterCatalogEntry, build_auto_router_catalog
from litellm.router_utils.access_windows import access_windows_config_error
from litellm.router_utils.add_retry_fallback_headers import (
get_fallback_errors_from_headers,
@ -5069,6 +5070,7 @@ class ProxyConfig:
def __init__(self) -> None:
self.config: Mapping[str, object] = MappingProxyType({})
self.auto_router_db_catalog: tuple[AutoRouterCatalogEntry, ...] | None = None
self._last_semantic_filter_config: dict[str, object] | None = None
self._last_websearch_interception_config: dict[str, object] | None = None
self._last_hashicorp_vault_config: dict[str, object] | None = None
@ -7692,6 +7694,12 @@ class ProxyConfig:
def _should_load_db_object(self, object_type: str | SupportedDBObjectType) -> bool:
return should_load_db_object(object_type=object_type)
def remove_auto_router_catalog_entries(self, model_ids: frozenset[str]) -> None:
if self.auto_router_db_catalog is not None:
self.auto_router_db_catalog = tuple(
row for row in self.auto_router_db_catalog if row.model_id not in model_ids
)
async def _get_models_from_db(self, prisma_client: PrismaClient) -> Sequence[_ProxyModelRow] | None:
"""
Fetch all model deployments from the DB.
@ -7712,6 +7720,7 @@ class ProxyConfig:
new_models: Final[Sequence[_ProxyModelRow]] = await ModelRepository(
WriterPinnedClient(prisma_client.db)
).table.find_many()
self.auto_router_db_catalog = build_auto_router_catalog(new_models)
return new_models
except Exception as e:
verbose_proxy_logger.exception(

View file

@ -1,18 +1,18 @@
{
"1m_context": {
"label": "1M Context",
"description": "Routes across models with 1M-token context windows: Luna for simple queries, Terra for medium, Sol for complex, Opus 5 at high thinking for reasoning.",
"description": "Routes across models with 1M-token context windows: GPT-6 Luna for simple queries, GPT-5.6 Terra for medium, GPT-6 Sol for complex, Opus 5.5 at high thinking for reasoning.",
"complexity_router_config": {
"tiers": {
"SIMPLE": ["gpt-5.6-luna"],
"SIMPLE": ["gpt-6-luna"],
"MEDIUM": ["gpt-5.6-terra"],
"COMPLEX": ["gpt-5.6-sol"],
"REASONING": ["claude-opus-5"]
"COMPLEX": ["gpt-6-sol"],
"REASONING": ["claude-opus-5-5"]
},
"tier_model_configs": {
"REASONING": [
{
"model_name": "claude-opus-5",
"model_name": "claude-opus-5-5",
"litellm_params": { "reasoning_effort": "high" }
}
]
@ -28,12 +28,12 @@
},
"anthropic_family": {
"label": "Anthropic Family",
"description": "Routes across the Claude model family: Haiku for simple queries, Sonnet for medium, Opus for complex, Fable 5.1 at high thinking for reasoning.",
"description": "Routes across the Claude model family: Haiku for simple queries, Sonnet for medium, Opus 5.5 for complex, Fable 5.1 at high thinking for reasoning.",
"complexity_router_config": {
"tiers": {
"SIMPLE": ["claude-haiku-4-5"],
"MEDIUM": ["claude-sonnet-5"],
"COMPLEX": ["claude-opus-5"],
"COMPLEX": ["claude-opus-5-5"],
"REASONING": ["claude-fable-5-1"]
},
"tier_model_configs": {
@ -55,12 +55,12 @@
},
"gemini_family": {
"label": "Gemini Family",
"description": "Routes across the Gemini model family: Flash Lite 2.5 for simple queries, Flash Lite 3.1 for medium, Flash 3.7 for complex, Pro 3.1 for reasoning-heavy requests.",
"description": "Routes across the Gemini model family: Flash Lite 3.5 for simple queries, Flash 3.8 for medium and complex queries, Pro 3.1 for reasoning-heavy requests.",
"complexity_router_config": {
"tiers": {
"SIMPLE": ["gemini-2.5-flash-lite"],
"MEDIUM": ["gemini-3.1-flash-lite"],
"COMPLEX": ["gemini-3.7-flash"],
"SIMPLE": ["gemini-3.5-flash-lite"],
"MEDIUM": ["gemini-3.8-flash"],
"COMPLEX": ["gemini-3.8-flash"],
"REASONING": ["gemini-3.1-pro-preview"]
},
"classifier_type": "heuristic",
@ -74,18 +74,18 @@
},
"lite": {
"label": "Lite",
"description": "Cost-optimized routing across providers: DeepSeek V4 Flash for simple queries, Muse Spark 1.2 at xhigh for medium, Kimi K3 at max for complex, Claude Opus 5 for reasoning. An LLM classifier with the agentic rubric assigns tiers.",
"description": "Cost-optimized routing across providers: DeepSeek V4 Flash for simple queries, Muse Spark 1.3 at xhigh for medium, Kimi K3 at max for complex, Claude Opus 5.5 for reasoning. An LLM classifier with the agentic rubric assigns tiers.",
"complexity_router_config": {
"tiers": {
"SIMPLE": ["deepseek-v4-flash"],
"MEDIUM": ["muse-spark-1.2"],
"MEDIUM": ["muse-spark-1.3"],
"COMPLEX": ["kimi-k3"],
"REASONING": ["claude-opus-5"]
"REASONING": ["claude-opus-5-5"]
},
"tier_model_configs": {
"MEDIUM": [
{
"model_name": "muse-spark-1.2",
"model_name": "muse-spark-1.3",
"litellm_params": { "reasoning_effort": "xhigh" }
}
],
@ -113,12 +113,12 @@
},
"openai_family": {
"label": "OpenAI Family",
"description": "Routes across the GPT model family: Luna for simple queries, Terra for medium, Sol for complex, Astra at xhigh thinking for reasoning.",
"description": "Routes across the GPT model family: GPT-6 Luna for simple queries, GPT-5.6 Terra for medium, GPT-6 Sol for complex, GPT-6 Astra at xhigh thinking for reasoning.",
"complexity_router_config": {
"tiers": {
"SIMPLE": ["gpt-5.6-luna"],
"SIMPLE": ["gpt-6-luna"],
"MEDIUM": ["gpt-5.6-terra"],
"COMPLEX": ["gpt-5.6-sol"],
"COMPLEX": ["gpt-6-sol"],
"REASONING": ["gpt-6-astra"]
},
"tier_model_configs": {

View file

@ -1,5 +1,5 @@
{
"version": "2026-09-17-v1",
"version": "2026-09-22-v1",
"models": [
{
"id": "gpt-6-astra-v1",
@ -8,6 +8,20 @@
"text": "OpenAI model for demanding end-to-end work, including reasoning, coding, research, and document tasks",
"sources": ["https://developers.openai.com/api/docs/models/gpt-6-astra"]
},
{
"id": "gpt-6-sol-v1",
"label": "GPT-6 Sol",
"model": "gpt-6-sol",
"text": "OpenAI model for complex coding and agentic workflows, supporting reasoning and tool calling through the Responses API",
"sources": ["https://developers.openai.com/api/docs/models/gpt-6-sol"]
},
{
"id": "gpt-6-luna-v1",
"label": "GPT-6 Luna",
"model": "gpt-6-luna",
"text": "OpenAI model for efficient, high-volume workloads, supporting reasoning and tool calling through the Responses API",
"sources": ["https://developers.openai.com/api/docs/models/gpt-6-luna"]
},
{
"id": "gpt-5.6-sol-v1",
"label": "GPT-5.6 Sol",
@ -63,6 +77,13 @@
"model": "claude-fable-5-1",
"text": "Anthropic model for demanding reasoning, long-running agentic coding, and multistep research, with always-on adaptive thinking",
"sources": ["https://platform.claude.com/docs/en/models/fable-5-1/overview"]
},
{
"id": "claude-opus-5-5-v1",
"label": "Claude Opus 5.5",
"model": "claude-opus-5-5",
"text": "Anthropic model for complex reasoning and agentic work, supporting adaptive thinking and tool use",
"sources": ["https://platform.claude.com/docs/en/models/opus-5-5/overview"]
}
],
"harnesses": [

View file

@ -10,12 +10,10 @@ from pydantic import ValidationError
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig
TUNING_BASELINE_PARAM_NAME: Final = "auto_router_tuning_baseline_v2"
# v2 hashes combine models and scoring rules; a new snapshot is required to separate them.
TUNING_BASELINE_PARAM_NAME: Final = "auto_router_tuning_baseline_v3"
HEURISTIC_V1_TUNING_FIELDS: Final = (
"tiers",
"tier_model_configs",
"classifier_type",
"tier_boundaries",
"reasoning_override_min_score",
"token_thresholds",
@ -49,8 +47,10 @@ def tuning_fingerprint(complexity_router_config: object) -> str | None:
validated: Final = ComplexityRouterConfig.model_validate(raw)
except ValidationError:
return None
supplied: Final = ((_TUNING_FIELD_SET - frozenset(("tier_model_configs",))) & frozenset(raw)) | (
frozenset(("tier_model_configs",)) if validated.tier_model_configs else frozenset()
# The UI always writes this built-in marker. Freeze its spelling so future defaults cannot change recorded hashes.
default_escalation: Final = validated.escalation_keywords in (None, ["LITELLM ESCALATE"])
supplied: Final = (_TUNING_FIELD_SET & frozenset(raw)) - (
frozenset(("escalation_keywords",)) if default_escalation else frozenset()
)
payload: Final = validated.model_dump(
mode="json",
@ -148,9 +148,9 @@ def tuning_limit_violation(*, held: int, limit: int | None) -> str | None:
if limit is None or held <= limit:
return None
return (
f"At most {limit} auto-router(s) with changed heuristic scorer settings or tier models can be modified "
f"At most {limit} auto-router(s) with changed heuristic scoring rules can be modified "
"without an auto-router license. Keep this router on its recorded settings, or revert the other changed "
"router to its baseline, or remove one of them."
"router to its baseline, or remove one of them. Selecting models does not use this allowance."
)

View file

@ -44,6 +44,25 @@ class ComplexityRouterConfigValidationResponse(BaseModel):
error: str | None = None
class AutoRouterAvailabilityRequest(BaseModel):
team_id: str | None = None
saved_model_id: str | None = None
complexity_router_config: Mapping[str, object] | None = None
class AutoRouterAllowance(BaseModel):
key: str
limit: int | None
remaining: int | None
used_by_this_router: bool = False
available: bool = True
class AutoRouterAvailabilityResponse(BaseModel):
allowances: tuple[AutoRouterAllowance, ...]
error: str | None = None
class AutoRouterRoutingTestRequest(BaseModel):
"""A single request to classify against a complexity-router config that need not be saved yet.

View file

@ -1440,11 +1440,22 @@ async def async_post_call_success_deployment_hook(
modified_response = response
CustomLogger: Final = _get_cached_custom_logger()
CustomGuardrail: Final = _get_cached_custom_guardrail()
for callback in litellm.callbacks:
if isinstance(callback, CustomLogger):
result = await callback.async_post_call_success_deployment_hook(
request_data, cast(LLMResponseTypes, modified_response), typed_call_type
)
try:
result = await callback.async_post_call_success_deployment_hook(
request_data, cast(LLMResponseTypes, modified_response), typed_call_type
)
except Exception: # noqa: BLE001 # a broken callback must not fail a completed request
if isinstance(callback, CustomGuardrail):
raise
verbose_logger.exception(
"async_post_call_success_deployment_hook error in %s for call_type=%s",
type(callback).__name__,
typed_call_type,
)
continue
if result is not None:
modified_response = result

View file

@ -25314,8 +25314,7 @@
"search_context_size_medium": 0.014,
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"cache_read_input_token_cost_batches": 1e-07
"web_search_billing_unit": "per_query"
},
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@ -25401,8 +25400,7 @@
"search_context_size_medium": 0.014,
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"cache_read_input_token_cost_batches": 2.5e-08
"web_search_billing_unit": "per_query"
},
"gemini-3.1-flash-image-preview": {
"input_cost_per_image": 0.00056,
@ -25482,8 +25480,7 @@
"supports_response_schema": false,
"supports_system_messages": true,
"supports_video_input": true,
"supports_vision": true,
"cache_read_input_token_cost_batches": 1.25e-08
"supports_vision": true
},
"gemini-3.1-flash-lite-preview": {
"cache_read_input_token_cost": 2.5e-08,
@ -25593,8 +25590,7 @@
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"input_cost_per_audio_token_batches": 2.5e-07,
"cache_read_input_token_cost_batches": 1.25e-08
"input_cost_per_audio_token_batches": 2.5e-07
},
"gemini-3.5-flash-lite": {
"deprecation_date": "2027-07-21",
@ -25652,8 +25648,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 1.5e-08
"google_maps_grounding_cost_per_query": 0.014
},
"deep-research-pro-preview-12-2025": {
"cache_read_input_token_cost": 2e-07,
@ -25689,8 +25684,7 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_vision": true,
"supports_web_search": true,
"cache_read_input_token_cost_batches": 1e-07
"supports_web_search": true
},
"gemini-2.5-flash-lite": {
"cache_read_input_audio_token_cost": 3e-08,
@ -26321,8 +26315,7 @@
"output_cost_per_token_batches": 4.5e-06,
"input_cost_per_token_flex": 7.5e-07,
"output_cost_per_token_flex": 4.5e-06,
"cache_read_input_token_cost_flex": 7.5e-08,
"cache_read_input_token_cost_batches": 7.5e-08
"cache_read_input_token_cost_flex": 7.5e-08
},
"vertex_ai/gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
@ -26380,8 +26373,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
@ -26440,8 +26432,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/gemini-3.8-flash": {
"prompt_cache_min_tokens": 4096,
@ -26500,8 +26491,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
@ -28250,8 +28240,7 @@
"output_cost_per_token_batches": 4.5e-06,
"input_cost_per_token_flex": 7.5e-07,
"output_cost_per_token_flex": 4.5e-06,
"cache_read_input_token_cost_flex": 7.5e-08,
"cache_read_input_token_cost_batches": 7.5e-08
"cache_read_input_token_cost_flex": 7.5e-08
},
"gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
@ -28309,8 +28298,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
@ -28369,8 +28357,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"gemini-3.8-flash": {
"prompt_cache_min_tokens": 4096,
@ -28429,8 +28416,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 3.75e-08
"google_maps_grounding_cost_per_query": 0.014
},
"gemini/gemini-2.5-pro-preview-tts": {
"cache_read_input_token_cost": 1.25e-07,
@ -33110,8 +33096,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true,
"cache_read_input_token_cost_batches": 6.25e-08
"supports_minimal_reasoning_effort": true
},
"gpt-5-chat": {
"cache_read_input_token_cost": 1.25e-07,
@ -33292,8 +33277,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true,
"cache_read_input_token_cost_batches": 1.25e-08
"supports_minimal_reasoning_effort": true
},
"gpt-5-nano": {
"cache_read_input_token_cost": 5e-09,
@ -33392,8 +33376,7 @@
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true,
"cache_read_input_token_cost_batches": 2.5e-09
"supports_minimal_reasoning_effort": true
},
"gpt-image-1": {
"cache_read_input_token_cost": 1.25e-06,
@ -47531,8 +47514,7 @@
"output_cost_per_token_flex": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
"supports_reasoning": false,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"cache_read_input_token_cost_batches": 1e-07
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@ -47570,8 +47552,7 @@
"output_cost_per_token_batches": 1.5e-06,
"output_cost_per_token_flex": 1.5e-06,
"supports_reasoning": false,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"cache_read_input_token_cost_batches": 2.5e-08
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/gemini-3.1-flash-image-preview": {
"input_cost_per_image": 0.00056,
@ -47627,8 +47608,7 @@
"supports_response_schema": false,
"supports_system_messages": true,
"supports_video_input": true,
"supports_vision": true,
"cache_read_input_token_cost_batches": 1.25e-08
"supports_vision": true
},
"vertex_ai/gemini-3.1-flash-lite-preview": {
"cache_read_input_token_cost": 2.5e-08,
@ -47739,8 +47719,7 @@
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"input_cost_per_audio_token_batches": 2.5e-07,
"cache_read_input_token_cost_batches": 1.25e-08
"input_cost_per_audio_token_batches": 2.5e-07
},
"vertex_ai/gemini-3.5-flash-lite": {
"deprecation_date": "2027-07-21",
@ -47799,8 +47778,7 @@
"search_context_size_high": 0.014
},
"web_search_billing_unit": "per_query",
"google_maps_grounding_cost_per_query": 0.014,
"cache_read_input_token_cost_batches": 1.5e-08
"google_maps_grounding_cost_per_query": 0.014
},
"vertex_ai/deep-research-pro-preview-12-2025": {
"cache_read_input_token_cost": 2e-07,
@ -47817,8 +47795,7 @@
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_batches": 6e-06,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing",
"cache_read_input_token_cost_batches": 1e-07
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/jamba-1.5": {
"input_cost_per_token": 2e-07,

View file

@ -46,6 +46,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family
- `router/` - routing and reliability behavior (fallbacks, cooldowns) plus the memory tests (`test_reliability_memory_e2e.py`: every worker's RSS as read at collection time, before any test traffic, must sit under a fixed idle budget, the release-gate check for a DB-backed boot that idles near the pod limit the way v1.100.x did; and a few hundred failing requests with retries and fallbacks must not grow proxy RSS past a fixed budget nor store a request snapshot past a fixed size, the release-gate check for the v1.100.0 retry-breadcrumb leak)
- `load/` - performance-category tests, kept OUT of the main suite: throughput/load SLO tests are a different testing category from functional e2e (variance-driven, historically flaky) and live outside this suite until re-implemented as their own pipeline (LIT-5163); do not add a live load test that runs in the default collection. What lives here: the weekly session-anomaly test (`test_weekly_session_anomaly_e2e.py`, Claude Code-shaped multi-turn sessions against real providers with ceilings on error rate, cache read/write, turn time, and spend; marked `weekly` and deselected unless `E2E_WEEKLY_ANOMALY` is set, driven by `.github/workflows/weekly_load_anomaly.yml`), the Redis chaos test (`test_redis_chaos_e2e.py`, locust load against mock deployments split round robin over `/chat/completions` and `/v1/messages`, one endpoint per simulated user, with `CLIENT PAUSE ALL` on the proxy's Redis mid-run to simulate it being down outright, asserting zero failed requests on every endpoint, budgeting RSS and CPU-per-request as ratios against the same run's healthy phase, and holding p50/p90/p99 latency and log-bytes-per-request to flat ceilings (a ratio cannot bound those two: an open breaker skips Redis instead of waiting on it, so the chaos phase can measure cheaper than baseline while still being far slower than a user should see); needs a proxy booted from `gateway/redis_chaos_ci_config.yml` on the same host with `E2E_PROXY_PID` and `E2E_PROXY_LOG` set, marked `redis_chaos`, deselected unless `E2E_REDIS_CHAOS` is set and excluded from the per-PR selector like the rest of `load/`, driven by `.github/workflows/test-e2e-redis-chaos.yml` and by the Buildkite `e2e-redis-chaos` step in project-releaser, which runs the proxy, Postgres and Valkey co-located with pytest in one pod and sets the opt-in), and markerless harness unit tests for the locust, process-usage, and session-anomaly aggregation logic
- `other/` - the holding-pen suite for the `other.*` registry cluster with no home of its own yet: the master-key auth gate, JWT auth (access tokens issued by a real Keycloak realm, `idp.py` plus `idp_realm.json`, whose JWKS the proxy's `JWT_PUBLIC_KEY_URL` points at; see CONTRIBUTING.md for the start command and config block), and the process-lifecycle health probes (liveness, public readiness, authenticated readiness diagnostics). Promote a cluster out once it is large/stable enough for its own suite
- `secret_manager/` - the gateway's `key_management_system` against a real secret manager: deployment keys resolved from it (`os.environ/<name>` where the name exists only in the manager) and virtual keys written to and deleted from it. The tests are backend-agnostic and each backend is its own lane, because the setting is global to the proxy: `E2E_SECRET_MANAGER=<system>` opts in and picks the backend from `secret_backends.BACKENDS`, the proxy is booted from `gateway/secret_manager_<system>_ci_config.yml` against the live manager, and the tests reach that manager through the backend's `SecretStore` (`secret_store_<system>.py`). A test needing something not every backend does carries `requires_capability(...)` and is deselected on lanes that lack it. `secret_manager/backend.sh up <system>` runs a backend in Docker and writes the proxy's and the tests' env. Marked `secret_manager`, deselected unless `E2E_SECRET_MANAGER` is set, and kept out of the per-PR selector. Backends today: `hashicorp_vault` and `cyberark` (CyberArk Conjur, which cannot delete, so the delete test is Vault-only)
- `gateway/` - proxy configuration only (`litellm-config.yml`); no tests
- `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees. The HTTP probes ride the shared transport (`ProxyClient.count_tokens` / `ProxyClient.messages`); the CLI-driving path stays bespoke
- `ui/` - the Admin UI browser suite: Playwright in TypeScript, driving the dashboard served by a live proxy on port 4000 (seeded postgres + mock LLM upstream; see its `run_e2e.sh`). It is a self-contained npm package with its own lockfile and does not use the Python harness, pytest markers, or the shared transport; the Python rules in this file (typed models, `Result` unions, basedpyright zero-error gate) do not apply inside it. Its only Python file, `fixtures/mock_llm_server/server.py`, is excluded from the e2e basedpyright gate via the root `pyrightconfig.json`

View file

@ -105,7 +105,7 @@ A couple of logging destinations are configured on the proxy rather than by the
### The pull request check
Every same-repository PR that adds, modifies, or renames a `tests/e2e/**/test_*.py` file runs those changed files three times. A change to the harness itself, meaning a root-level `tests/e2e/*.py` file or `pytest.ini`, `tests/e2e/gateway/`, `.github/e2e-stack/`, or the workflow, also runs the `access_control` suite and both JWT suites as canaries, because those files have no test of their own that exercises the stack. `.github/e2e-stack/select_tests.py` applies both rules. The stack config at `tests/e2e/gateway/stage_mirror_ci_config.yml` must declare every model the selected suites use; a missing one shows up as a failed test id in the public log. The suite's own single rerun for network errors and 5xx responses (see `pytest.ini`) applies on every pass, so a transport blip does not fail the check while a race inside a test still does. The stage-mirror stack has a control-plane backend, two gateways behind nginx, Postgres, Keycloak, Jaeger, and TLS cluster-mode Valkey. Realm-only edits also trigger these canaries. The stack exports every gateway address in `LITELLM_PROXY_REPLICA_URLS`, so model registration waits until each gateway lists the new model rather than whichever one the load balancer answered from. Documentation, deleted-file, and application-only changes do not start the stack or request environment approval. The `ui/`, `claude_code/`, and `load/` directories, `batches/test_managed_files_enforcement_e2e.py`, `llm_translation/realtime/test_realtime_pipecat_audio_e2e.py`, and `guardrails/test_presidio_masking_e2e.py` remain outside this check because they use separate tooling or need a differently configured stack: the pipecat audio suite skips itself at import time unless the NLTK `punkt_tab` data is installed, and the presidio suite fails without the analyzer and anonymizer services this stack does not start. `logging/test_otel_v2_langfuse_generation_output_e2e.py` is marked `otel_v2` and deselects itself unless `E2E_OTEL_V2` is set, because it needs a gateway booted with `LITELLM_OTEL_V2=true` and Langfuse credentials, neither of which this stack provides, so run it with `E2E_OTEL_V2=1` against a local OTel v2 proxy. The Redis chaos test under `load/` needs a proxy it can pause the Redis of on the same host (`gateway/redis_chaos_ci_config.yml`), which `.github/workflows/test-e2e-redis-chaos.yml` boots, and which the Buildkite `e2e-redis-chaos` step in project-releaser runs co-located with Postgres and Valkey in one pod; it is deselected unless `E2E_REDIS_CHAOS` is set
Every same-repository PR that adds, modifies, or renames a `tests/e2e/**/test_*.py` file runs those changed files three times. A change to the harness itself, meaning a root-level `tests/e2e/*.py` file or `pytest.ini`, `tests/e2e/gateway/`, `.github/e2e-stack/`, or the workflow, also runs the `access_control` suite and both JWT suites as canaries, because those files have no test of their own that exercises the stack. `.github/e2e-stack/select_tests.py` applies both rules. The stack config at `tests/e2e/gateway/stage_mirror_ci_config.yml` must declare every model the selected suites use; a missing one shows up as a failed test id in the public log. The suite's own single rerun for network errors and 5xx responses (see `pytest.ini`) applies on every pass, so a transport blip does not fail the check while a race inside a test still does. The stage-mirror stack has a control-plane backend, two gateways behind nginx, Postgres, Keycloak, Jaeger, and TLS cluster-mode Valkey. Realm-only edits also trigger these canaries. The stack exports every gateway address in `LITELLM_PROXY_REPLICA_URLS`, so model registration waits until each gateway lists the new model rather than whichever one the load balancer answered from. Documentation, deleted-file, and application-only changes do not start the stack or request environment approval. The `ui/`, `claude_code/`, `load/`, and `secret_manager/` directories, `batches/test_managed_files_enforcement_e2e.py`, `llm_translation/realtime/test_realtime_pipecat_audio_e2e.py`, and `guardrails/test_presidio_masking_e2e.py` remain outside this check because they use separate tooling or need a differently configured stack: the pipecat audio suite skips itself at import time unless the NLTK `punkt_tab` data is installed, and the presidio suite fails without the analyzer and anonymizer services this stack does not start. `logging/test_otel_v2_langfuse_generation_output_e2e.py` is marked `otel_v2` and deselects itself unless `E2E_OTEL_V2` is set, because it needs a gateway booted with `LITELLM_OTEL_V2=true` and Langfuse credentials, neither of which this stack provides, so run it with `E2E_OTEL_V2=1` against a local OTel v2 proxy. The Redis chaos test under `load/` needs a proxy it can pause the Redis of on the same host (`gateway/redis_chaos_ci_config.yml`), which `.github/workflows/test-e2e-redis-chaos.yml` boots, and which the Buildkite `e2e-redis-chaos` step in project-releaser runs co-located with Postgres and Valkey in one pod; it is deselected unless `E2E_REDIS_CHAOS` is set. The `secret_manager/` lanes each need a proxy configured against their own secret manager (see Secret manager lanes below)
Every selected file must execute at least one passing test in each pass, and any test failure, collection error, or entirely skipped or deselected file fails the check. A file whose tests are all marked skip therefore cannot pass this check, so unskip at least one of them, or add the file to `UNSUPPORTED` in `select_tests.py` with the reason, before changing one. A failed pass stops the run. The public log prints pytest's one-line summary for each pass, including the rerun count, and names each failed or errored test as `classname::name`, so a retried network error or a failing test is visible without the raw output. The final `e2e-changed-tests` job succeeds only when no supported test files changed or the approved run completed all three passes. Fork PRs with selected tests fail this gate until a maintainer brings the reviewed change onto a same-repository branch
@ -117,6 +117,31 @@ Fetched values of eight characters or more are masked before use, while shorter
To reproduce the CI topology on a dedicated machine, `bash .github/e2e-stack/up.sh` reads `tests/e2e/.env`, writes `stack.env` under `${E2E_STACK_DIR:-/tmp/litellm-e2e-stack}`, and `bash .github/e2e-stack/down.sh` stops it. Keep this directory private and remove its credential files and logs after use
### Secret manager lanes
`key_management_system` is global to the proxy, so the `secret_manager/` tests run once per backend, each against its own proxy. The backends are `hashicorp_vault` and `cyberark` (CyberArk Conjur). `E2E_SECRET_MANAGER` opts in and names the backend (a key of `secret_backends.BACKENDS`). The proxy boots from `gateway/secret_manager_<system>_ci_config.yml`, and the tests reach the same manager through that backend's `SecretStore`. The managers are enterprise features, so the proxy needs a license. `secret_manager/backend.sh` runs any backend in Docker and writes its env, so every lane runs the same way locally:
```bash
bash tests/e2e/secret_manager/backend.sh up cyberark
(set -a; . ~/.cache/litellm-e2e-secret-manager/cyberark/proxy.env; set +a; env -u OPENAI_API_KEY LITELLM_LICENSE=... \
LITELLM_MASTER_KEY=sk-1234 DATABASE_URL=... uv run litellm --config tests/e2e/gateway/secret_manager_cyberark_ci_config.yml --port 4000)
(set -a; . ~/.cache/litellm-e2e-secret-manager/cyberark/tests.env; set +a; OPENAI_API_KEY=... \
uv run --group e2e-dev pytest tests/e2e/secret_manager/ -v)
bash tests/e2e/secret_manager/backend.sh down cyberark
```
`E2E_SECRET_MANAGER_PORT` moves the manager off its usual port (8200 for Vault, 8080 for Conjur), and `E2E_SECRET_MANAGER_DIR` moves the env files. Keep that directory private, because both files hold a working admin credential. Keep `OPENAI_API_KEY` out of the proxy's environment. The tests copy the runner's key into the manager under a fresh name per test, so a passing call proves the key came through the manager rather than the `os.environ` fallback `get_secret` takes when the manager errors
A backend declares what it supports in its `SecretBackend.capabilities`, and a test that needs something not every backend does carries `@pytest.mark.requires_capability(...)`, so it is deselected, not failed or skipped, on the lanes that lack it. CyberArk has no `deletes_stored_keys`, because the proxy's delete answers `not_supported` and Conjur keeps the key, so the delete test runs only on the Vault lane
To add a backend, leave the tests and markers alone and add:
1. `secret_manager/secret_store_<system>.py`: a `SecretStore` (`write`, `read` returning None when absent, idempotent `destroy`) over the manager's own API through `e2e_http`'s external helpers, read from `E2E_<SYSTEM>_*` env vars, and a `SecretBackend` whose `system` is the litellm `KeyManagementSystem` value and whose `capabilities` lists what it supports
2. its entry in `secret_backends.BACKENDS`
3. `gateway/secret_manager_<system>_ci_config.yml`, a copy of an existing lane's with only `key_management_system` changed
4. an `up_<system>` function in `secret_manager/backend.sh` that starts the manager and writes `proxy.env` and `tests.env`
5. a CI step that runs `backend.sh up <system>` (or the same containers as sidecars), boots the proxy with `proxy.env` and a license, and runs pytest with `tests.env`
### Record and replay
Record/replay scopes to the proxy's provider-bound traffic only. In `E2E_FIXTURE_MODE=record` the harness boots a local provider-edge server, edge-wired tests register their deployments with an `api_base` pointing at it, and every provider call the proxy makes is forwarded verbatim and written to a fixture bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`). `E2E_FIXTURE_MODE=replay` runs the same tests against the same live proxy and database, but the edge answers the proxy's provider calls from the bundle instead of the provider, so the run makes zero provider calls and spends nothing while key auth, routing, cost calculation, and spend-log writes all still execute for real. Unset (or `live`) behaves exactly as before the knob existed. Both record and replay need the proxy up; only the provider is taken out of the loop

View file

@ -22,6 +22,7 @@ failures are hard test failures (see `tests/e2e/AGENTS.md`).
| Bedrock | yes (unified only) | yes | yes | yes (unfiltered managed list) | yes (provider-transformed) | S3 (`s3_bucket_name` + `aws_*` + `AWS_BATCH_ROLE_ARN` on model) |
| Bedrock GovCloud (`us-gov-west-1`) | yes (unified only) | yes | no | no | yes (provider-transformed) | S3 (`s3_bucket_name` + `aws_*` on model, resolved from `AWS_GOVCLOUD_ACCESS_KEY_ID` / `AWS_GOVCLOUD_SECRET_ACCESS_KEY` / `AWS_GOVCLOUD_BATCH_S3_BUCKET` / `AWS_GOVCLOUD_BATCH_ROLE_ARN`) |
| Bedrock split S3 identity | no | no | no | no | yes (file upload, content, delete) | S3 signed with `s3_access_key_id` / `s3_secret_access_key` (`AWS_S3_ONLY_ACCESS_KEY_ID` / `AWS_S3_ONLY_SECRET_ACCESS_KEY`, object rights on `AWS_BATCH_S3_BUCKET` only) while `aws_*` is `AWS_BEDROCK_ONLY_ACCESS_KEY_ID` / `AWS_BEDROCK_ONLY_SECRET_ACCESS_KEY`, an identity with no S3 rights on that bucket |
| Bedrock blank S3 env | yes (unified only, on an owned gateway exporting `AWS_S3_ENCRYPTION_KEY_ID` / `AWS_S3_BUCKET_OWNER` as empty strings) | no | no | no | no | S3 (`s3_bucket_name` + `aws_*` + `AWS_BATCH_ROLE_ARN` in the gateway config); blank env vars must be treated as unset, not serialized |
Bedrock cancel maps to `StopModelInvocationJob` and comes back `cancelling`; the
lifecycle asserts it the same way it does for OpenAI (`_CANCEL_ASSERTED_PROVIDERS`).

View file

@ -0,0 +1,145 @@
"""An owned, source-built proxy whose process env exports AWS_S3_* vars blank.
The shared fixture proxy inherits the harness env, which cannot reproduce a user
shell that exports AWS_S3_ENCRYPTION_KEY_ID / AWS_S3_BUCKET_OWNER as empty
strings. This gateway boots a second proxy with both vars present but blank, so
a batch create through it proves blank means unset, not an empty string.
"""
from __future__ import annotations
import os
import shutil
import socket
import subprocess
import sys
import tempfile
import time
from collections.abc import Mapping
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
from e2e_config import unique_marker
from e2e_http import NoBody
from idp import stop_process_group
from proxy_client import ProxyClient, build_proxy_client
from pydantic import TypeAdapter
STARTUP_TIMEOUT_SECONDS: Final = 240
LOG_TAIL_BYTES: Final = 4000
REPO_ROOT: Final = Path(__file__).resolve().parents[3]
_CONFIG_YAML: Final = """model_list:
- model_name: bedrock-blank-s3-batch
litellm_params:
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: os.environ/AWS_REGION
s3_region_name: os.environ/AWS_REGION
s3_bucket_name: os.environ/AWS_BATCH_S3_BUCKET
s3_access_key_id: os.environ/AWS_ACCESS_KEY_ID
s3_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_batch_role_arn: os.environ/AWS_BATCH_ROLE_ARN
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
database_url: os.environ/DATABASE_URL
"""
def available_port() -> int:
with socket.socket() as listener:
listener.bind(("127.0.0.1", 0))
return TypeAdapter(tuple[str, int]).validate_python(listener.getsockname())[1]
@dataclass(slots=True)
class BedrockEnvGateway:
base_url: str
master_key: str
proxy: ProxyClient
_environment: Mapping[str, str] = field(repr=False)
_command: tuple[str, ...] = field(repr=False)
_log_path: Path
_child: subprocess.Popen[bytes] | None = field(default=None, init=False, repr=False)
@classmethod
def start(cls) -> BedrockEnvGateway:
assert os.environ.get("DATABASE_URL"), "DATABASE_URL is required for the blank-S3-env gateway"
port: Final = available_port()
base_url: Final = f"http://127.0.0.1:{port}"
master_key: Final = f"sk-e2e-blank-s3-{unique_marker()}"
directory: Final = Path(tempfile.mkdtemp(prefix="litellm-e2e-blank-s3-"))
config: Final = directory / "blank-s3-gateway.yaml"
config.write_text(_CONFIG_YAML)
environment: Final = {
**{key: value for key, value in os.environ.items() if not key.startswith("REDIS_")},
"DATABASE_URL": os.environ["DATABASE_URL"],
"LITELLM_MASTER_KEY": master_key,
"STORE_MODEL_IN_DB": "False",
"PYTHONPATH": str(REPO_ROOT),
"AWS_S3_ENCRYPTION_KEY_ID": "",
"AWS_S3_BUCKET_OWNER": "",
}
gateway: Final = cls(
base_url=base_url,
master_key=master_key,
proxy=build_proxy_client(
base_url=base_url,
control_plane_base_url=base_url,
replica_urls=(base_url,),
master_key=master_key,
),
_environment=environment,
_command=(
sys.executable,
"-m",
"litellm.proxy.proxy_cli",
"--config",
str(config),
"--port",
str(port),
"--host",
"127.0.0.1",
),
_log_path=directory / "blank-s3-gateway.log",
)
with gateway._log_path.open("ab") as log:
gateway._child = subprocess.Popen(
gateway._command,
env=dict(gateway._environment),
stdout=log,
stderr=log,
start_new_session=True,
cwd=REPO_ROOT,
)
deadline: Final = time.monotonic() + STARTUP_TIMEOUT_SECONDS
while time.monotonic() < deadline:
assert gateway._child.poll() is None, (
f"blank-S3-env gateway exited early; log tail:\n{gateway.log_tail()}"
)
result = gateway.proxy.transport.probe("/health/liveliness", params=NoBody())
if result.status_code == 200:
return gateway
time.sleep(0.5)
tail: Final = gateway.log_tail()
gateway.stop()
raise AssertionError(
f"blank-S3-env gateway did not become ready in {STARTUP_TIMEOUT_SECONDS}s; log tail:\n{tail}"
)
def log_tail(self) -> str:
if not self._log_path.exists():
return "<no log written>"
with self._log_path.open("rb") as log:
log.seek(0, 2)
size: Final = log.tell()
log.seek(max(0, size - LOG_TAIL_BYTES))
return log.read().decode("utf-8", errors="replace")
def stop(self) -> None:
if self._child is not None:
stop_process_group(self._child)
shutil.rmtree(self._log_path.parent, ignore_errors=True)

View file

@ -0,0 +1,109 @@
"""Live e2e pin for Bedrock batch create with blank AWS_S3_* env vars.
Owns its own file (not test_batches_e2e.py) so the PR changed-file e2e gate
stays a single tiny file: this class boots its own gateway with
AWS_S3_ENCRYPTION_KEY_ID and AWS_S3_BUCKET_OWNER exported empty, then runs the
unified target_model_names upload + batch create lifecycle against real Bedrock.
"""
from __future__ import annotations
import json
from typing import Final
import pytest
from batch_cleanup import cleanup_batch, cleanup_file
from batch_client import BatchClient, BatchCreateBody, BatchObject, FileObject
from bedrock_env_gateway import BedrockEnvGateway
from capabilities import is_managed_id
from e2e_http import FileUploadForm, require_successful_call, unwrap
from lifecycle import ResourceManager
from models import KeyGenerateBody
pytestmark = pytest.mark.e2e
CREATED_BATCH_STATUSES = {"validating", "in_progress", "finalizing"}
BLANK_S3_RAW_MODEL: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
def render_jsonl(model: str) -> bytes:
line = {
"custom_id": "req-1",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": model,
"messages": [{"role": "user", "content": "ping"}],
"max_tokens": 8,
},
}
return (json.dumps(line) + "\n").encode()
def assert_file_object(file: FileObject, *, provider: str) -> None:
assert file.object == "file", f"file.object={file.object!r}"
assert file.purpose == "batch", f"file.purpose={file.purpose!r}"
assert file.bytes is not None, f"file.bytes={file.bytes!r}"
if provider != "bedrock":
assert file.bytes > 0, f"file.bytes={file.bytes!r}"
assert file.status, "file.status missing"
assert file.created_at is not None and file.created_at > 0, "file.created_at missing"
def assert_batch_object(batch: BatchObject) -> None:
assert batch.object == "batch", f"batch.object={batch.object!r}"
if batch.endpoint:
assert batch.endpoint == "/v1/chat/completions", f"batch.endpoint={batch.endpoint!r}"
assert batch.completion_window == "24h", f"window={batch.completion_window!r}"
assert batch.input_file_id, "batch.input_file_id missing"
assert batch.created_at is not None and batch.created_at > 0, "batch.created_at missing"
class TestBedrockBatchBlankS3EnvVars:
"""Bedrock batch create with AWS_S3_* env vars exported but blank.
Regression: a blank AWS_S3_ENCRYPTION_KEY_ID or AWS_S3_BUCKET_OWNER env var
resolved to "" and was serialized into the create-job request, which Bedrock
rejects. The owned gateway exports both vars empty, so the unified lifecycle
only passes when blank is treated as unset.
"""
@pytest.mark.covers(
"llm.batches.bedrock.blank_s3_env.nonstream.works",
"llm.files.bedrock.upload.nonstream.works",
exercised_on=["batches", "files"],
)
def test_unified_batch_create_ignores_blank_s3_env_vars(self, resources: ResourceManager) -> None:
gateway: Final = BedrockEnvGateway.start()
resources.defer(gateway.stop)
client: Final = BatchClient(proxy=gateway.proxy)
key: Final = client.proxy.generate_key(KeyGenerateBody(models=[], user_id="e2e-test-user"))
resources.defer(lambda: client.proxy.delete_key(key))
file: Final = unwrap(
client.upload_file(
content=render_jsonl(BLANK_S3_RAW_MODEL),
form=FileUploadForm(purpose="batch", target_model_names="bedrock-blank-s3-batch"),
key=key,
)
)
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="bedrock")
created: Final = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
assert created.status_code < 400, (
f"blank AWS_S3_ENCRYPTION_KEY_ID / AWS_S3_BUCKET_OWNER must be treated as "
f"unset; Bedrock rejected the job: {created.body[:400]}"
)
require_successful_call(created)
batch: Final = BatchObject.model_validate_json(created.body)
resources.defer(lambda: cleanup_batch(client, batch.id, key=key))
assert is_managed_id(batch.id), (
f"blank-S3-env create via target_model_names must return a managed batch id, got {batch.id!r}"
)
assert batch.status in CREATED_BATCH_STATUSES, (
f"blank-S3-env batch has non-transitional status {batch.status!r}"
)
assert_batch_object(batch)

View file

@ -36,6 +36,7 @@ from e2e_config import (
PROVIDER_EDGE_HOST_OPT_IN_ENV,
PROXY_BASE_URL,
REDIS_CHAOS_OPT_IN_ENV,
SECRET_MANAGER_OPT_IN_ENV,
WEEKLY_ANOMALY_OPT_IN_ENV,
unique_marker,
)
@ -70,6 +71,7 @@ OPT_IN_MARKERS: Final = MappingProxyType(
"provider_edge_host": PROVIDER_EDGE_HOST_OPT_IN_ENV,
"otel_v2": OTEL_V2_OPT_IN_ENV,
"otel_tls": OTEL_TLS_OPT_IN_ENV,
"secret_manager": SECRET_MANAGER_OPT_IN_ENV,
}
)
@ -172,6 +174,11 @@ def pytest_configure(config: pytest.Config) -> None:
"markers",
"otel_tls: needs a stack whose gateway exports OTLP over TLS signed by the CA in SSL_CERT_FILE; deselected unless E2E_OTEL_EXPORTER_ENDPOINT is set",
)
config.addinivalue_line(
"markers",
"secret_manager: needs a proxy booted from gateway/secret_manager_<system>_ci_config.yml against that live "
"secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py)",
)
def pytest_sessionstart(session: pytest.Session) -> None:

View file

@ -24,6 +24,7 @@
- {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.bedrock.assume_role.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: assume_role, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create under STS assume-role credentials"}
- {id: llm.batches.bedrock.govcloud_partition.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: govcloud_partition, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create in the us-gov-west-1 partition"}
- {id: llm.batches.bedrock.blank_s3_env.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: blank_s3_env, streaming: nonstream, assertions: [works], source: "test_bedrock_blank_s3_env_e2e.py", rationale: "Bedrock batch create treats blank AWS_S3_ENCRYPTION_KEY_ID / AWS_S3_BUCKET_OWNER env vars as unset instead of serializing empty strings"}
- {id: llm.batches.bedrock.cancel.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch cancel (StopModelInvocationJob) returns the same id with a cancelling/cancelled status"}
- {id: llm.batches.bedrock.list.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "A Bedrock managed batch is present in the GET /v1/batches list envelope"}
- {id: llm.batches.hosted_vllm.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: batches, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible batch create"}

View file

@ -40,7 +40,10 @@
- {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.passthrough.headers_forwarded, module: other, tier: P0, area: config, assertions: [headers_forwarded], source: "passthrough/utils.py forward_headers_from_request", rationale: "Custom pass-through static headers and x-pass-* client headers reach the upstream"}
- {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.secret_resolution.kms_integration, module: other, tier: P1, area: config, assertions: [kms_integration], source: "secret_managers/main.py get_secret / secret_manager/test_secret_manager_e2e.py", rationale: "A deployment whose api_key is os.environ/<name> gets its key from the configured secret manager when that name exists only in the manager"}
- {id: other.config.secret_resolution.manager_value_used, module: other, tier: P1, area: config, assertions: [manager_value_used], source: "secret_managers/main.py get_secret / secret_manager/test_secret_manager_e2e.py", rationale: "The value the manager holds is what reaches the provider: a bogus key in the manager is rejected by the provider with 401, so a passing resolution test cannot be an env fallback"}
- {id: other.config.secret_manager.virtual_key_stored, module: other, tier: P1, area: config, assertions: [virtual_key_stored], source: "key_management_event_hooks.py _store_virtual_key_in_secret_manager", rationale: "With store_virtual_keys, /key/generate writes the new key under prefix_for_stored_virtual_keys + key_alias in the manager"}
- {id: other.config.secret_manager.virtual_key_deleted, module: other, tier: P1, area: config, assertions: [virtual_key_deleted], source: "key_management_event_hooks.py _delete_virtual_keys_from_secret_manager", rationale: "/key/delete removes the stored key from the manager, so a revoked key does not linger there"}
- {id: other.config.overrides.audit_logged, module: other, tier: P1, area: config, assertions: [audit_logged], source: "config_override_endpoints.py:67-100", rationale: "Config override mutations audit-logged, values redacted"}
- {id: other.key_mgmt.regenerate.grace_period_honored, module: other, tier: P1, area: auth, assertions: [grace_period_honored], source: "key_management_endpoints.py:4503-4560", rationale: "Old key valid during grace_period then revoked"}
- {id: other.key_mgmt.spend_reset.resets_to_value, module: other, tier: P1, area: auth, assertions: [resets_to_value], source: "key_management_endpoints.py:4841", rationale: "reset_spend resets accumulated spend"}

View file

@ -65,6 +65,7 @@ LlmCapability = Literal[
"assume_role",
"basic",
"batch_deployment",
"blank_s3_env",
"count_tokens",
"govcloud_partition",
"split_s3_credentials",

View file

@ -150,6 +150,7 @@ MCP_OAUTH_LIVE_OPT_IN_ENV: Final = "E2E_MCP_OAUTH_LIVE"
PROVIDER_EDGE_HOST_OPT_IN_ENV: Final = "E2E_PROVIDER_EDGE_HOST_REACHABLE"
OTEL_V2_OPT_IN_ENV: Final = "E2E_OTEL_V2"
OTEL_TLS_OPT_IN_ENV: Final = "E2E_OTEL_EXPORTER_ENDPOINT"
SECRET_MANAGER_OPT_IN_ENV: Final = "E2E_SECRET_MANAGER"
ANOMALY_SESSIONS = int(os.environ.get("E2E_ANOMALY_SESSIONS", "6"))
ANOMALY_TURNS_PER_SESSION = int(os.environ.get("E2E_ANOMALY_TURNS_PER_SESSION", "6"))
ANOMALY_TURN_ATTEMPTS = int(os.environ.get("E2E_ANOMALY_TURN_ATTEMPTS", "3"))

View file

@ -129,9 +129,9 @@ class ProbeResult(BaseModel):
class ExternalWrite(BaseModel):
"""Outcome of a write to a non-proxy API (an identity provider's admin API)
that answers with a status and, on create, a Location header naming the new
resource rather than a JSON body."""
"""Outcome of a call to a non-proxy API (an identity provider's admin API, a
secret manager) that answers with a status, on create a Location header naming
the new resource, and a body kept as text rather than parsed as JSON."""
status_code: int
location: str = ""
@ -491,6 +491,30 @@ def post_json_external(
)
def send_text_external(
method: Literal["GET", "POST", "PATCH"],
url: str,
*,
headers: BaseModel,
content: str | None = None,
timeout: float = 30.0,
) -> ExternalWrite:
"""Send an absolute URL outside the proxy a raw text body (or none) and keep the
answer as text, for an API that takes and returns neither JSON nor forms: CyberArk
Conjur takes a secret value or a YAML policy and returns a secret as its raw value."""
try:
resp = requests.request(
method,
url,
headers=_headers(headers),
data=content.encode() if content is not None else None,
timeout=timeout,
)
except requests.RequestException as exc:
return ExternalWrite(status_code=-1, body=str(exc))
return ExternalWrite(status_code=resp.status_code, body=resp.text)
def delete_external(url: str, *, headers: BaseModel, timeout: float = 30.0) -> ExternalWrite:
try:
resp = requests.delete(url, headers=_headers(headers), timeout=timeout)

View file

@ -0,0 +1,8 @@
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
store_model_in_db: true
key_management_system: cyberark
key_management_settings:
access_mode: read_and_write
store_virtual_keys: true
prefix_for_stored_virtual_keys: litellm-e2e/virtual-keys/

View file

@ -0,0 +1,8 @@
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
store_model_in_db: true
key_management_system: hashicorp_vault
key_management_settings:
access_mode: read_and_write
store_virtual_keys: true
prefix_for_stored_virtual_keys: litellm-e2e/virtual-keys/

View file

@ -17,3 +17,4 @@ markers =
provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set
otel_v2: needs a proxy running with LITELLM_OTEL_V2=true; deselected unless E2E_OTEL_V2 is set
otel_tls: needs a stack whose gateway exports OTLP over TLS signed by the CA in SSL_CERT_FILE; deselected unless E2E_OTEL_EXPORTER_ENDPOINT is set
secret_manager: needs a proxy booted from gateway/secret_manager_<system>_ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py)

View file

@ -0,0 +1,76 @@
#!/usr/bin/env bash
set -euo pipefail
umask 077
usage() {
local systems
systems=$(declare -F | sed -n 's/^declare -f up_//p' | paste -sd '|' -)
echo "usage: $0 up|down $systems" >&2
exit 2
}
action=${1:-}
system=${2:-}
dir=${E2E_SECRET_MANAGER_DIR:-$HOME/.cache/litellm-e2e-secret-manager}/$system
name=litellm-e2e-$system
wait_for() {
local url=$1
for _ in $(seq 1 90); do
if curl -sf -o /dev/null "$url"; then
return 0
fi
sleep 2
done
echo "$system did not answer at $url" >&2
return 1
}
down() {
docker rm -f "$name" "$name-db" >/dev/null 2>&1 || true
docker network rm "$name" >/dev/null 2>&1 || true
rm -rf "$dir"
}
up_hashicorp_vault() {
local port=${E2E_SECRET_MANAGER_PORT:-8200}
local token
token=e2e-$(openssl rand -hex 16)
docker run -d --name "$name" -p "127.0.0.1:$port:8200" --cap-add IPC_LOCK \
-e VAULT_DEV_ROOT_TOKEN_ID="$token" hashicorp/vault:1.20 >/dev/null
wait_for "http://127.0.0.1:$port/v1/sys/health"
printf 'HCP_VAULT_ADDR=http://127.0.0.1:%s\nHCP_VAULT_TOKEN=%s\n' "$port" "$token" >"$dir/proxy.env"
printf 'E2E_VAULT_ADDR=http://127.0.0.1:%s\nE2E_VAULT_TOKEN=%s\n' "$port" "$token" >"$dir/tests.env"
}
up_cyberark() {
local port=${E2E_SECRET_MANAGER_PORT:-8080}
local data_key api_key
docker network create "$name" >/dev/null
docker run -d --name "$name-db" --network "$name" -e POSTGRES_HOST_AUTH_METHOD=trust postgres:15 >/dev/null
data_key=$(docker run --rm cyberark/conjur:1.24 data-key generate)
docker run -d --name "$name" --network "$name" -p "127.0.0.1:$port:80" \
-e DATABASE_URL="postgres://postgres@$name-db/postgres" -e CONJUR_DATA_KEY="$data_key" \
-e CONJUR_AUTHENTICATORS=authn cyberark/conjur:1.24 server >/dev/null
wait_for "http://127.0.0.1:$port/"
docker exec "$name" conjurctl account create --name default >/dev/null
api_key=$(docker exec "$name" conjurctl role retrieve-key default:user:admin | tr -d '\r\n')
printf 'CYBERARK_API_BASE=http://127.0.0.1:%s\nCYBERARK_ACCOUNT=default\nCYBERARK_USERNAME=admin\nCYBERARK_API_KEY=%s\n' \
"$port" "$api_key" >"$dir/proxy.env"
printf 'E2E_CYBERARK_API_BASE=http://127.0.0.1:%s\nE2E_CYBERARK_ACCOUNT=default\nE2E_CYBERARK_USERNAME=admin\nE2E_CYBERARK_API_KEY=%s\n' \
"$port" "$api_key" >"$dir/tests.env"
}
[[ $# -eq 2 && -n $system ]] && declare -F "up_$system" >/dev/null || usage
case $action in
up)
down
mkdir -p "$dir"
"up_$system"
echo "E2E_SECRET_MANAGER=$system" >>"$dir/tests.env"
echo "$system is up; env in $dir/proxy.env (proxy) and $dir/tests.env (pytest)"
;;
down) down ;;
*) usage ;;
esac

View file

@ -0,0 +1,57 @@
from __future__ import annotations
import os
from dataclasses import dataclass
from typing import Final
import pytest
from e2e_config import SECRET_MANAGER_OPT_IN_ENV
from proxy_client import ProxyClient
from secret_backends import BACKENDS, selected_backend
from secret_store import SecretBackend, SecretStore
REQUIRES_CAPABILITY: Final = "requires_capability"
def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
"markers",
f"{REQUIRES_CAPABILITY}(capability): secret_manager test deselected when the backend "
f"{SECRET_MANAGER_OPT_IN_ENV} names lacks the capability (secret_store.Capability)",
)
def _lacks_capability(item: pytest.Item, backend: SecretBackend) -> bool:
marker: Final = item.get_closest_marker(REQUIRES_CAPABILITY)
return marker is not None and marker.args[0] not in backend.capabilities
def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item]) -> None:
backend: Final = BACKENDS.get(os.environ.get(SECRET_MANAGER_OPT_IN_ENV, "").strip())
if backend is None:
return
deselected: Final = [item for item in items if _lacks_capability(item, backend)]
if deselected:
config.hook.pytest_deselected(items=deselected)
items[:] = [item for item in items if not _lacks_capability(item, backend)]
@dataclass(frozen=True, slots=True)
class SecretManagerClient:
proxy: ProxyClient
@pytest.fixture(scope="session")
def client(proxy: ProxyClient) -> SecretManagerClient:
return SecretManagerClient(proxy)
@pytest.fixture(scope="session")
def backend() -> SecretBackend:
return selected_backend()
@pytest.fixture(scope="session")
def store(backend: SecretBackend) -> SecretStore:
return backend.from_env()

View file

@ -0,0 +1,25 @@
from __future__ import annotations
import os
from types import MappingProxyType
from typing import Final
import pytest
from e2e_config import SECRET_MANAGER_OPT_IN_ENV
from secret_store import SecretBackend
from secret_store_cyberark import CYBERARK
from secret_store_hashicorp_vault import HASHICORP_VAULT
BACKENDS: Final = MappingProxyType({backend.system: backend for backend in (HASHICORP_VAULT, CYBERARK)})
def selected_backend() -> SecretBackend:
system: Final = os.environ.get(SECRET_MANAGER_OPT_IN_ENV, "").strip()
backend: Final = BACKENDS.get(system)
if backend is None:
pytest.fail(
f"{SECRET_MANAGER_OPT_IN_ENV}={system!r} names no secret manager backend; "
f"set it to one of {sorted(BACKENDS)}"
)
return backend

View file

@ -0,0 +1,29 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from typing import Final, Literal, Protocol
SECRET_MANAGER_CONFIG_DIR: Final = "gateway"
class SecretStore(Protocol):
def write(self, name: str, value: str) -> None: ...
def read(self, name: str) -> str | None: ...
def destroy(self, name: str) -> None: ...
Capability = Literal["deletes_stored_keys"]
@dataclass(frozen=True, slots=True)
class SecretBackend:
system: str
from_env: Callable[[], SecretStore]
capabilities: frozenset[Capability]
@property
def proxy_config(self) -> str:
return f"{SECRET_MANAGER_CONFIG_DIR}/secret_manager_{self.system}_ci_config.yml"

View file

@ -0,0 +1,118 @@
from __future__ import annotations
import base64
import os
from dataclasses import dataclass, field
from typing import Final, Literal
from urllib.parse import quote
import pytest
import yaml
from e2e_http import ExternalWrite, Headers, send_text_external
from pydantic import Field
from secret_store import SecretBackend
CYBERARK_API_BASE_ENV: Final = "E2E_CYBERARK_API_BASE"
CYBERARK_ACCOUNT_ENV: Final = "E2E_CYBERARK_ACCOUNT"
CYBERARK_USERNAME_ENV: Final = "E2E_CYBERARK_USERNAME"
CYBERARK_API_KEY_ENV: Final = "E2E_CYBERARK_API_KEY"
# The same defaults CyberArkSecretManager falls back to for CYBERARK_*.
DEFAULT_API_BASE: Final = "http://127.0.0.1:8080"
DEFAULT_ACCOUNT: Final = "default"
DEFAULT_USERNAME: Final = "admin"
SYSTEM: Final = "cyberark"
_START_HINT: Final = (
f"Start one with `bash tests/e2e/secret_manager/backend.sh up {SYSTEM}`, which writes the env for "
f"the proxy (booted from gateway/secret_manager_{SYSTEM}_ci_config.yml) and for the tests"
)
class ConjurHeaders(Headers):
authorization: str = Field(repr=False)
content_type: str | None = Field(default=None, serialization_alias="Content-Type")
def _policy_scalar(name: str) -> str:
# Quoted the way CyberArkSecretManager._ensure_variable_exists quotes it.
return yaml.safe_dump(name, default_style='"').strip()
@dataclass(frozen=True, slots=True)
class Conjur:
base_url: str
account: str
username: str
api_key: str = field(repr=False)
def _fail_unless_reached(self, result: ExternalWrite, action: str) -> None:
if result.status_code == -1:
pytest.fail(f"No live Conjur at {self.base_url}: {result.body}. {_START_HINT}")
if result.status_code == 401:
pytest.fail(f"Conjur rejected {self.username}'s credentials while trying to {action}. {_START_HINT}")
def _headers(self, content_type: str | None = None) -> ConjurHeaders:
# Tokens last about eight minutes, so each call authenticates afresh rather than
# letting a long session outlive a cached one.
auth: Final = send_text_external(
"POST",
f"{self.base_url}/authn/{self.account}/{quote(self.username, safe='')}/authenticate",
headers=Headers(),
content=self.api_key,
)
self._fail_unless_reached(auth, "authenticate")
if not auth.ok:
pytest.fail(f"Conjur refused to authenticate {self.username}: HTTP {auth.status_code} {auth.body[:300]}")
token: Final = base64.b64encode(auth.body.encode()).decode()
return ConjurHeaders(authorization=f'Token token="{token}"', content_type=content_type)
def _secret_url(self, name: str) -> str:
return f"{self.base_url}/secrets/{self.account}/variable/{quote(name, safe='')}"
def _update_root_policy(self, method: Literal["POST", "PATCH"], policy: str, action: str) -> None:
result: Final = send_text_external(
method,
f"{self.base_url}/policies/{self.account}/policy/root",
headers=self._headers(content_type="application/x-yaml"),
content=policy,
)
self._fail_unless_reached(result, action)
if not result.ok:
pytest.fail(f"Conjur refused to {action}: HTTP {result.status_code} {result.body[:300]}")
def write(self, name: str, value: str) -> None:
self._update_root_policy("POST", f"- !variable {_policy_scalar(name)}\n", f"declare {name}")
result: Final = send_text_external("POST", self._secret_url(name), headers=self._headers(), content=value)
self._fail_unless_reached(result, f"write {name}")
if not result.ok:
pytest.fail(f"Conjur refused to write {name}: HTTP {result.status_code} {result.body[:300]}")
def read(self, name: str) -> str | None:
result: Final = send_text_external("GET", self._secret_url(name), headers=self._headers())
self._fail_unless_reached(result, f"read {name}")
if result.status_code == 404:
return None
if not result.ok:
pytest.fail(f"Conjur refused to read {name}: HTTP {result.status_code} {result.body[:300]}")
return result.body
def destroy(self, name: str) -> None:
self._update_root_policy("PATCH", f"- !delete\n record: !variable {_policy_scalar(name)}\n", f"destroy {name}")
def conjur_from_env() -> Conjur:
api_key: Final = os.environ.get(CYBERARK_API_KEY_ENV, "").strip()
if not api_key:
pytest.fail(f"The {SYSTEM} lane needs {CYBERARK_API_KEY_ENV} to reach its Conjur. {_START_HINT}")
return Conjur(
base_url=os.environ.get(CYBERARK_API_BASE_ENV, "").strip().rstrip("/") or DEFAULT_API_BASE,
account=os.environ.get(CYBERARK_ACCOUNT_ENV, "").strip() or DEFAULT_ACCOUNT,
username=os.environ.get(CYBERARK_USERNAME_ENV, "").strip() or DEFAULT_USERNAME,
api_key=api_key,
)
CYBERARK: Final = SecretBackend(system=SYSTEM, from_env=conjur_from_env, capabilities=frozenset())

View file

@ -0,0 +1,113 @@
from __future__ import annotations
import os
from dataclasses import dataclass, field
from typing import Final
import pytest
from e2e_http import (
Headers,
NetworkError,
Success,
UnknownApiError,
delete_external,
get_external,
post_json_external,
)
from pydantic import BaseModel, Field
from secret_store import SecretBackend
VAULT_ADDR_ENV: Final = "E2E_VAULT_ADDR"
VAULT_TOKEN_ENV: Final = "E2E_VAULT_TOKEN"
VAULT_MOUNT_ENV: Final = "E2E_VAULT_MOUNT_NAME"
DEFAULT_VAULT_ADDR: Final = "http://127.0.0.1:8200"
DEFAULT_MOUNT: Final = "secret"
SYSTEM: Final = "hashicorp_vault"
_START_HINT: Final = (
f"Start one with `bash tests/e2e/secret_manager/backend.sh up {SYSTEM}`, which writes the env for "
f"the proxy (booted from gateway/secret_manager_{SYSTEM}_ci_config.yml) and for the tests"
)
class VaultHeaders(Headers):
x_vault_token: str = Field(serialization_alias="X-Vault-Token", repr=False)
class KvData(BaseModel):
key: str = Field(repr=False)
class KvWriteBody(BaseModel):
data: KvData
class KvReadData(BaseModel):
data: KvData
class KvReadResponse(BaseModel):
data: KvReadData
@dataclass(frozen=True, slots=True)
class Vault:
base_url: str
token: str = field(repr=False)
mount: str = DEFAULT_MOUNT
def _headers(self) -> VaultHeaders:
return VaultHeaders(x_vault_token=self.token)
def _data_url(self, name: str) -> str:
return f"{self.base_url}/v1/{self.mount}/data/{name}"
def _metadata_url(self, name: str) -> str:
return f"{self.base_url}/v1/{self.mount}/metadata/{name}"
def write(self, name: str, value: str) -> None:
write: Final = post_json_external(
self._data_url(name), headers=self._headers(), json=KvWriteBody(data=KvData(key=value))
)
if write.status_code == -1:
pytest.fail(f"No live Vault at {self.base_url}: {write.body}. {_START_HINT}")
if not write.ok:
pytest.fail(f"Vault refused to write {name}: HTTP {write.status_code} {write.body[:300]}")
def read(self, name: str) -> str | None:
result: Final = get_external(self._data_url(name), headers=self._headers(), response_type=KvReadResponse)
match result:
case Success(data=body):
return body.data.data.key
case UnknownApiError(status_code=404):
return None
case NetworkError(message=message):
return pytest.fail(f"No live Vault at {self.base_url}: {message}. {_START_HINT}")
case _:
return pytest.fail(f"Vault refused to read {name}: {result}")
def destroy(self, name: str) -> None:
write: Final = delete_external(self._metadata_url(name), headers=self._headers())
if not write.ok and write.status_code != 404:
pytest.fail(f"Vault refused to destroy {name}: HTTP {write.status_code} {write.body[:300]}")
def vault_from_env() -> Vault:
token: Final = os.environ.get(VAULT_TOKEN_ENV, "").strip()
if not token:
pytest.fail(f"The hashicorp_vault lane needs {VAULT_TOKEN_ENV} to reach its Vault. {_START_HINT}")
return Vault(
base_url=os.environ.get(VAULT_ADDR_ENV, DEFAULT_VAULT_ADDR).rstrip("/"),
token=token,
mount=os.environ.get(VAULT_MOUNT_ENV, "").strip() or DEFAULT_MOUNT,
)
HASHICORP_VAULT: Final = SecretBackend(
system=SYSTEM,
from_env=vault_from_env,
capabilities=frozenset({"deletes_stored_keys"}),
)

View file

@ -0,0 +1,128 @@
from __future__ import annotations
import os
import time
from collections.abc import Callable
from typing import Final
import pytest
from e2e_config import unique_marker
from e2e_http import Result, Success, UnauthorizedError, unwrap
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, LiteLLMParamsBody
from proxy_client import ProxyClient
from secret_store import SecretStore
pytestmark = [pytest.mark.e2e, pytest.mark.secret_manager]
BACKEND_MODEL: Final = "openai/gpt-4o-mini"
VIRTUAL_KEY_PREFIX: Final = "litellm-e2e/virtual-keys/"
PROVIDER_KEY_ENV: Final = "OPENAI_API_KEY"
# The proxy's env never holds OPENAI_API_KEY and each test seeds it under a fresh name, so a passing
# call proves the key came from the manager and not get_secret's os.environ fallback.
def _provider_key() -> str:
key: Final = os.environ.get(PROVIDER_KEY_ENV, "").strip()
if not key:
pytest.fail(f"The secret manager suite seeds the manager with the runner's {PROVIDER_KEY_ENV}, which is unset")
return key
def _seed(store: SecretStore, resources: ResourceManager, value: str) -> str:
name: Final = f"litellm-e2e-openai-{unique_marker()}"
store.write(name, value)
resources.defer(lambda: store.destroy(name))
return name
def _deploy(proxy: ProxyClient, resources: ResourceManager, secret_name: str) -> str:
model_name: Final = f"secret-manager-backed-{unique_marker()}"
model_id: Final = proxy.create_model(
model_name,
LiteLLMParamsBody(model=BACKEND_MODEL, api_key=f"os.environ/{secret_name}"),
provider_live=True,
)
resources.defer(lambda: proxy.delete_model(model_id))
return model_name
def _chat(proxy: ProxyClient, key: str, model: str) -> Result[ChatResponse]:
return proxy.chat(
key,
ChatBody(
model=model,
messages=[ChatMessage(role="user", content=f"reply with one word {unique_marker()}")],
max_tokens=16,
),
)
def _eventually(proxy: ProxyClient, read: Callable[[], str | None], expected: str | None, context: str) -> None:
deadline: Final = time.monotonic() + proxy.poll_timeout
last: str | None = read()
while last != expected and time.monotonic() < deadline:
time.sleep(proxy.poll_interval)
last = read()
if last != expected:
pytest.fail(
f"{context}: the secret manager still holds {'a value' if last is not None else 'nothing'} after the deadline"
)
class TestSecretManager:
@pytest.mark.covers("other.config.secret_resolution.kms_integration")
def test_deployment_key_resolves_from_the_manager(
self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore, scoped_key: str
) -> None:
model: Final = _deploy(proxy, resources, _seed(store, resources, _provider_key()))
response: Final = unwrap(_chat(proxy, scoped_key, model))
assert response.choices, f"the manager-backed deployment answered with no choices: {response}"
@pytest.mark.covers("other.config.secret_resolution.manager_value_used")
def test_deployment_uses_the_value_the_manager_holds(
self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore, scoped_key: str
) -> None:
bogus: Final = f"sk-litellm-e2e-not-a-key-{unique_marker()}"
model: Final = _deploy(proxy, resources, _seed(store, resources, bogus))
result: Final = _chat(proxy, scoped_key, model)
match result:
case UnauthorizedError(body=body):
assert "AuthenticationError" in body, f"the 401 did not come from the provider: {body[:300]}"
case Success():
pytest.fail("a deployment whose managed secret is not a real key still reached the provider")
case _:
pytest.fail(f"expected the provider to reject the manager-held key with 401, got {result}")
@pytest.mark.covers("other.config.secret_manager.virtual_key_stored")
def test_generated_key_is_written_to_the_manager(
self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore
) -> None:
alias: Final = f"litellm-e2e-vk-{unique_marker()}"
secret_name: Final = f"{VIRTUAL_KEY_PREFIX}{alias}"
resources.defer(lambda: store.destroy(secret_name))
key: Final = proxy.generate_key(KeyGenerateBody(key_alias=alias))
resources.defer(lambda: proxy.delete_key(key))
_eventually(proxy, lambda: store.read(secret_name), key, f"the generated key {alias}")
@pytest.mark.requires_capability("deletes_stored_keys")
@pytest.mark.covers("other.config.secret_manager.virtual_key_deleted")
def test_deleted_key_is_removed_from_the_manager(
self, proxy: ProxyClient, resources: ResourceManager, store: SecretStore
) -> None:
alias: Final = f"litellm-e2e-vk-{unique_marker()}"
secret_name: Final = f"{VIRTUAL_KEY_PREFIX}{alias}"
resources.defer(lambda: store.destroy(secret_name))
key: Final = proxy.generate_key(KeyGenerateBody(key_alias=alias))
resources.defer(lambda: proxy.delete_key(key))
_eventually(proxy, lambda: store.read(secret_name), key, f"the generated key {alias}")
proxy.delete_key(key)
_eventually(proxy, lambda: store.read(secret_name), None, f"the deleted key {alias}")

View file

@ -7,6 +7,7 @@ import time
import uuid
from collections.abc import Iterator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
@ -45,8 +46,37 @@ def stop_root_process(process: subprocess.Popen[bytes]) -> bool:
return True
@dataclass(frozen=True, slots=True)
class OwnedProxy:
gateway: Gateway
process: subprocess.Popen[bytes]
log: Path
@contextmanager
def owned_proxy(gateway: Gateway, directory: Path, overrides: Mapping[str, str], *, config: Path | None = None, remove_environment: tuple[str, ...] = ()) -> Iterator[Gateway]:
def owned_proxy(
gateway: Gateway,
directory: Path,
overrides: Mapping[str, str],
*,
config: Path | None = None,
remove_environment: tuple[str, ...] = (),
) -> Iterator[Gateway]:
with owned_proxy_process(
gateway, directory, overrides, config=config, remove_environment=remove_environment
) as owned:
yield owned.gateway
@contextmanager
def owned_proxy_process(
gateway: Gateway,
directory: Path,
overrides: Mapping[str, str],
*,
config: Path | None = None,
remove_environment: tuple[str, ...] = (),
) -> Iterator[OwnedProxy]:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
port: Final = reserve.getsockname()[1]
@ -60,7 +90,8 @@ def owned_proxy(gateway: Gateway, directory: Path, overrides: Mapping[str, str],
}
output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory)))
output.mkdir(parents=True, exist_ok=True)
with (output / f"owned-proxy-{uuid.uuid4().hex}.log").open("w") as log:
log_path: Final = output / f"owned-proxy-{uuid.uuid4().hex}.log"
with log_path.open("w") as log:
process: Final = subprocess.Popen(
[
sys.executable,
@ -95,7 +126,7 @@ def owned_proxy(gateway: Gateway, directory: Path, overrides: Mapping[str, str],
pass
assert time.monotonic() < deadline, "Owned proxy readiness deadline exceeded"
time.sleep(0.1)
yield Gateway(client, gateway.key, gateway.upstream_url)
yield OwnedProxy(Gateway(client, gateway.key, gateway.upstream_url), process, log_path)
finally:
root_stopped: Final = stop_root_process(process)
residual: Final = group_members(process.pid)

View file

@ -281,6 +281,12 @@
"tests/integration/spend/test_filtered_ledger.py::test_rotated_keys_users_and_model_groups_preserve_success_failure_cache_ledger": [
"quota_management.spend_tracking.filtered_ledger_preserves_owner_identity_and_totals"
],
"tests/integration/spend/test_shutdown_flush.py::test_daily_spend_batch_cancelled_while_waiting_for_a_pool_connection_is_written_by_the_final_flush": [
"quota_management.spend_tracking.shutdown_cancel_keeps_in_flight_daily_batch"
],
"tests/integration/spend/test_shutdown_flush.py::test_daily_spend_batch_cancelled_while_waiting_for_a_row_lock_is_written_exactly_once": [
"quota_management.spend_tracking.shutdown_cancel_keeps_in_flight_daily_batch"
],
"tests/integration/spend/test_spend_calculate.py::test_spend_calculate_rejects_unpriced_model_with_400": [
"quota_management.spend_tracking.spend_calculate.rejects_unpriced_model"
],
@ -1776,6 +1782,18 @@
],
"tests/integration/mcp/test_mcp_lifecycle.py::test_same_url_server_grants_scope_discovery_and_direct_or_virtual_execution[bearer]": [
"other.mcp.permissions.same_url_servers_enforce_discovery_and_execution"
],
"tests/integration/observability/test_callback_delivery.py::test_response_survives_raising_success_deployment_hook[chat]": [
"other.observability.callbacks.raising_success_deployment_hook_keeps_response"
],
"tests/integration/observability/test_callback_delivery.py::test_response_survives_raising_success_deployment_hook[embeddings]": [
"other.observability.callbacks.raising_success_deployment_hook_keeps_response"
],
"tests/integration/observability/test_callback_delivery.py::test_response_survives_raising_success_deployment_hook[responses]": [
"other.observability.callbacks.raising_success_deployment_hook_keeps_response"
],
"tests/integration/observability/test_callback_delivery.py::test_response_survives_raising_success_deployment_hook[videos]": [
"other.observability.callbacks.raising_success_deployment_hook_keeps_response"
]
},
"browser": {

View file

@ -1,13 +1,15 @@
import base64
import json
import uuid
from collections.abc import Callable, Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.client import Gateway, JsonValue, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
@ -152,3 +154,163 @@ def test_concurrent_success_and_failure_join_callbacks_and_rows_without_credenti
assert rows[0]["prompt_tokens"] == event["prompt_tokens"]
else:
assert event["prompt_tokens"] == event["completion_tokens"] == rows[0]["completion_tokens"] == 0
_RAISING_HOOK: Final = """
from litellm.integrations.custom_logger import CustomLogger
class RaisingHook(CustomLogger):
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
raise RuntimeError(f"hook rejected {type(response).__name__} for {call_type}")
instance = RaisingHook()
"""
_VIDEO_JOB: Final = {
"id": "video_hook_isolation",
"object": "video",
"status": "queued",
"model": "sora-2",
"seconds": "4",
"size": "720x1280",
}
_UPSTREAM_REPLIES: Final[Mapping[str, Mapping[str, JsonValue]]] = {
"/v1/chat/completions": {
"id": "chatcmpl_hook_isolation",
"object": "chat.completion",
"created": 1,
"model": "gpt-5.6",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
},
"/v1/embeddings": {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
},
"/v1/responses": {
"id": "resp_hook_isolation",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-5.6",
"output": [
{
"type": "message",
"id": "msg_hook_isolation",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
"parallel_tool_calls": False,
"tool_choice": "auto",
"tools": [],
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
},
"/v1/videos": _VIDEO_JOB,
}
def _item(value: JsonValue, index: int) -> JsonValue:
assert isinstance(value, list), f"Expected a list, received {type(value).__name__}"
return value[index]
def _chat_text(body: dict[str, JsonValue]) -> str:
return string_value(object_value(object_value(_item(body["choices"], 0))["message"])["content"])
def _embedding_vector(body: dict[str, JsonValue]) -> JsonValue:
return object_value(_item(body["data"], 0))["embedding"]
def _responses_text(body: dict[str, JsonValue]) -> str:
return string_value(object_value(_item(object_value(_item(body["output"], 0))["content"], 0))["text"])
def _video_job(body: dict[str, JsonValue]) -> tuple[str, str]:
encoded_id: Final = string_value(body["id"]).removeprefix("video_")
decoded: Final = base64.b64decode(encoded_id).decode()
return decoded.rsplit("video_id:", 1)[-1], string_value(body["status"])
@dataclass(frozen=True, slots=True)
class _Surface:
route: str
upstream_model: str
body: Callable[[str], dict[str, JsonValue]]
observed: Callable[[dict[str, JsonValue]], JsonValue | tuple[str, str]]
expected: JsonValue | tuple[str, str]
_SURFACES: Final = (
pytest.param(
_Surface(
"/v1/chat/completions",
"openai/gpt-5.6",
lambda model: {"model": model, "messages": [{"role": "user", "content": "hook isolation"}]},
_chat_text,
"hi",
),
id="chat",
),
pytest.param(
_Surface(
"/v1/embeddings",
"openai/text-embedding-3-small",
lambda model: {"model": model, "input": "hook isolation"},
_embedding_vector,
[0.1, 0.2],
),
id="embeddings",
),
pytest.param(
_Surface(
"/v1/responses",
"openai/gpt-5.6",
lambda model: {"model": model, "input": "hook isolation"},
_responses_text,
"hi",
),
id="responses",
),
pytest.param(
_Surface(
"/v1/videos",
"openai/sora-2",
lambda model: {"model": model, "prompt": "a cat"},
_video_job,
(_VIDEO_JOB["id"], _VIDEO_JOB["status"]),
),
id="videos",
),
)
@pytest.mark.covers("other.observability.callbacks.raising_success_deployment_hook_keeps_response")
@pytest.mark.parametrize("surface", _SURFACES)
def test_response_survives_raising_success_deployment_hook(gateway: Gateway, tmp_path: Path, surface: _Surface) -> None:
def upstream(request: Request) -> Reply:
assert request.target == surface.route, request.target
assert b"hook isolation" in request.body or b"a cat" in request.body, request.body[:300]
return Reply(body=json.dumps(_UPSTREAM_REPLIES[surface.route]).encode())
(tmp_path / "raising_hook.py").write_text(_RAISING_HOOK)
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update({"callbacks": ["raising_hook.instance"]})
path: Final = tmp_path / "hook.yaml"
path.write_text(yaml.safe_dump(config))
with (
wire_server(upstream) as provider,
owned_proxy(gateway, tmp_path, {}, config=path) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(model=surface.upstream_model, api_base=provider.url + "/v1")
response: Final = candidate.request("POST", surface.route, surface.body(model))
assert response.status_code == 200, response.text
assert surface.observed(object_value(response.json())) == surface.expected, response.text

View file

@ -0,0 +1,189 @@
import json
import os
import signal
import threading
import uuid
from collections.abc import Callable, Iterator
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import psycopg
import pytest
import yaml
from integration._support.client import Gateway, delete_key_if_present, eventually, string_value
from integration._support.database import read_rows
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
REQUESTS_WHILE_BLOCKED: Final = 6
CANCEL_LOG_LINE: Final = "in-flight scheduled job(s) for shutdown"
BATCH_DRAINED_LOG_LINE: Final = f"flushed {REQUESTS_WHILE_BLOCKED} daily spend update items from in-memory queue"
MODEL_INSERT_ARRIVED_LOG_LINE: Final = "path=/model/new"
def _api_requests(table: str, column: str, identity: str) -> int:
rows: Final = read_rows(
f'SELECT coalesce(sum(api_requests), 0)::int AS total FROM "{table}" WHERE {column}=%s', (identity,)
)
total: Final = rows[0]["total"]
assert isinstance(total, int)
return total
def _waiting_on(table: str) -> int:
rows: Final = read_rows(
"SELECT count(*)::int AS waiting FROM pg_stat_activity WHERE wait_event_type='Lock' AND query LIKE %s",
(f'%"{table}"%',),
)
waiting: Final = rows[0]["waiting"]
assert isinstance(waiting, int)
return waiting
def _provider(request: Request) -> Reply:
if request.method != "POST":
return Reply(status=404, body=b'{"error":"not scripted"}')
assert request.target == "/v1/chat/completions"
return Reply(
body=json.dumps(
{
"id": "chatcmpl-" + uuid.uuid4().hex,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
}
).encode()
)
@dataclass(frozen=True, slots=True)
class _Shutdown:
owner: str
team: str
owned: OwnedProxy
key: str
model: str
def chat(self) -> None:
body: Final = {"model": self.model, "messages": [{"role": "user", "content": f"spend {uuid.uuid4().hex}"}]}
assert self.owned.gateway.request("POST", "/v1/chat/completions", body, key=self.key).status_code == 200
def daily_user_requests(self) -> int:
return _api_requests("LiteLLM_DailyUserSpend", "user_id", self.owner)
def logged(self, line: str, times: int = 1) -> bool:
return self.owned.log.read_text(errors="replace").count(line) >= times
def chat_while_spend_update_is_blocked(self, blocker: psycopg.Connection, table: str) -> None:
blocker.execute(f'LOCK TABLE "{table}" IN EXCLUSIVE MODE')
for _ in range(REQUESTS_WHILE_BLOCKED):
self.chat()
eventually(lambda: _waiting_on(table), lambda waiting: waiting == 1, seconds=30)
def start_blocked_model_insert(self) -> threading.Thread:
body: Final = {
"model_name": f"integration-blocked-{uuid.uuid4().hex}",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "integration-provider-key"},
"model_info": {},
}
def insert() -> None:
try:
self.owned.gateway.request("POST", "/model/new", body)
except httpx.TransportError:
pass
thread: Final = threading.Thread(target=insert, daemon=True)
thread.start()
return thread
def terminate_once(self, blocked: Callable[[], bool], release: Callable[[], None]) -> None:
eventually(blocked, lambda state: state, seconds=60)
self.owned.process.send_signal(signal.SIGTERM)
eventually(lambda: self.logged(CANCEL_LOG_LINE), lambda seen: seen, seconds=60)
release()
self.owned.process.wait(timeout=120)
def _config_with_pool_limit(tmp_path: Path, pool_limit: int) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["general_settings"]["database_connection_pool_limit"] = pool_limit
config["general_settings"]["database_connection_pool_timeout"] = 60
path: Final = tmp_path / f"pool-{pool_limit}.yaml"
path.write_text(yaml.safe_dump(config))
return path
@contextmanager
def _proxy_with_one_seeded_row(gateway: Gateway, tmp_path: Path, pool_limit: int) -> Iterator[_Shutdown]:
owner: Final = f"integration-owner-{uuid.uuid4().hex}"
with gateway.scenario() as scenario, wire_server(_provider) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
team: Final = scenario.team(models=[model])
with owned_proxy_process(
gateway,
tmp_path,
{
"LITELLM_LOG": "DEBUG",
"GRACEFUL_SHUTDOWN_TIMEOUT": "1",
"SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS": "1",
"SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS": "5",
},
config=_config_with_pool_limit(tmp_path, pool_limit),
) as owned:
key: Final = string_value(
owned.gateway.post("/key/generate", {"user_id": owner, "team_id": team, "models": [model]})["key"]
)
scenario.cleanups.callback(delete_key_if_present, gateway, key)
shutdown: Final = _Shutdown(owner, team, owned, key, model)
shutdown.chat()
eventually(shutdown.daily_user_requests, lambda total: total == 1, seconds=60)
yield shutdown
assert _api_requests("LiteLLM_DailyUserSpend", "user_id", owner) == 1 + REQUESTS_WHILE_BLOCKED
assert _api_requests("LiteLLM_DailyTeamSpend", "team_id", team) == 1 + REQUESTS_WHILE_BLOCKED
@pytest.mark.covers("quota_management.spend_tracking.shutdown_cancel_keeps_in_flight_daily_batch")
def test_daily_spend_batch_cancelled_while_waiting_for_a_pool_connection_is_written_by_the_final_flush(
gateway: Gateway, tmp_path: Path
) -> None:
with (
_proxy_with_one_seeded_row(gateway, tmp_path, pool_limit=2) as shutdown,
psycopg.connect(os.environ["DATABASE_URL"]) as models,
psycopg.connect(os.environ["DATABASE_URL"]) as memberships,
):
models.execute('LOCK TABLE "LiteLLM_ProxyModelTable" IN EXCLUSIVE MODE')
first: Final = shutdown.start_blocked_model_insert()
eventually(lambda: _waiting_on("LiteLLM_ProxyModelTable"), lambda waiting: waiting == 1, seconds=30)
shutdown.chat_while_spend_update_is_blocked(memberships, "LiteLLM_TeamMembership")
second: Final = shutdown.start_blocked_model_insert()
eventually(lambda: shutdown.logged(MODEL_INSERT_ARRIVED_LOG_LINE, times=2), lambda seen: seen, seconds=30)
memberships.rollback()
eventually(lambda: _waiting_on("LiteLLM_ProxyModelTable"), lambda waiting: waiting == 2, seconds=30)
shutdown.terminate_once(lambda: shutdown.logged(BATCH_DRAINED_LOG_LINE), models.rollback)
first.join(timeout=30)
second.join(timeout=30)
@pytest.mark.covers("quota_management.spend_tracking.shutdown_cancel_keeps_in_flight_daily_batch")
def test_daily_spend_batch_cancelled_while_waiting_for_a_row_lock_is_written_exactly_once(
gateway: Gateway, tmp_path: Path
) -> None:
with (
_proxy_with_one_seeded_row(gateway, tmp_path, pool_limit=10) as shutdown,
psycopg.connect(os.environ["DATABASE_URL"]) as holder,
psycopg.connect(os.environ["DATABASE_URL"]) as memberships,
):
holder.execute('SELECT 1 FROM "LiteLLM_DailyUserSpend" WHERE user_id=%s FOR UPDATE', (shutdown.owner,))
shutdown.chat_while_spend_update_is_blocked(memberships, "LiteLLM_TeamMembership")
memberships.rollback()
shutdown.terminate_once(
lambda: shutdown.logged(BATCH_DRAINED_LOG_LINE) and _waiting_on("LiteLLM_DailyUserSpend") == 1,
holder.rollback,
)

View file

@ -263,7 +263,7 @@ def test_trimming_should_not_change_original_messages():
assert messages == messages_copy
@pytest.mark.parametrize("model", ["gpt-4-0125-preview", "claude-sonnet-4-6"])
@pytest.mark.parametrize("model", ["gpt-5.4-mini", "claude-sonnet-4-6"])
def test_trimming_with_model_cost_max_input_tokens(model):
messages = [
{"role": "system", "content": "This is a normal system message"},

View file

@ -9,6 +9,12 @@ from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
fireworks = FireworksAIConfig()
VISION_MODEL = next(
key.removeprefix("fireworks_ai/")
for key, info in litellm.model_cost.items()
if key.startswith("fireworks_ai/accounts/fireworks/models/") and info.get("supports_vision") is True
)
def test_map_openai_params_tool_choice():
# Test case 1: tool_choice is "required"
@ -97,7 +103,7 @@ def test_document_inlining_example(disable_add_transform_inline_image_block):
with patch.object(client, "post") as mock_post:
try:
completion(
model="fireworks_ai/accounts/fireworks/models/minimax-m3",
model=f"fireworks_ai/{VISION_MODEL}",
messages=[
{
"role": "user",
@ -157,7 +163,7 @@ def test_transform_inline_no_longer_added(content, expected_url):
result = litellm.FireworksAIConfig()._transform_messages_helper(
messages=messages,
model="accounts/fireworks/models/minimax-m3",
model=VISION_MODEL,
litellm_params={},
)
result_image_block = result[0]["content"][0]
@ -182,7 +188,7 @@ def test_global_disable_flag_no_longer_adds_transform_inline(is_disabled):
]
result = litellm.FireworksAIConfig()._transform_messages_helper(
messages=messages,
model="accounts/fireworks/models/minimax-m3",
model=VISION_MODEL,
litellm_params={},
)
assert result[0]["content"][0]["image_url"] == url
@ -204,7 +210,7 @@ def test_global_disable_flag_with_transform_messages_helper(monkeypatch):
) as mock_post:
try:
completion(
model="fireworks_ai/accounts/fireworks/models/minimax-m3",
model=f"fireworks_ai/{VISION_MODEL}",
messages=[
{
"role": "user",

View file

@ -238,7 +238,7 @@ def test_gemini_image_generation_accumulates_multiple_image_prompt_token_details
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "gemini/gemini-3-pro-image-preview"
model = "gemini/gemini-3-pro-image"
config = GoogleImageGenConfig()
usage_metadata = {

View file

@ -32,7 +32,7 @@ class TestGroq(BaseLLMChatTest):
@pytest.mark.parametrize(
"model",
["groq/qwen/qwen3-32b", "groq/openai/gpt-oss-20b", "groq/openai/gpt-oss-120b"],
["groq/qwen/qwen3.8-27b", "groq/openai/gpt-oss-20b", "groq/openai/gpt-oss-120b"],
)
def test_reasoning_effort_in_supported_params(self, model):
"""Test that reasoning_effort is in the list of supported parameters for Groq"""

View file

@ -537,7 +537,7 @@ def test_dynamic_drop_params_e2e():
) as mock_response:
try:
response = litellm.completion(
model="command-r",
model="command-r-08-2024",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
response_format={"key": "value"},
drop_params=True,
@ -556,7 +556,7 @@ def test_dynamic_pass_additional_params():
) as mock_response:
try:
response = litellm.completion(
model="command-r",
model="command-r-08-2024",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
custom_param="test",
api_key="my-custom-key",
@ -606,7 +606,7 @@ def test_dynamic_drop_params_parallel_tool_calls():
) as mock_response:
try:
response = litellm.completion(
model="command-r",
model="command-r-08-2024",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
parallel_tool_calls=True,
drop_params=True,
@ -663,7 +663,7 @@ def test_dynamic_drop_additional_params_e2e():
) as mock_response:
try:
response = litellm.completion(
model="command-r",
model="command-r-08-2024",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
response_format={"key": "value"},
additional_drop_params=["response_format"],

View file

@ -164,7 +164,7 @@ def test_xai_message_name_filtering():
class TestXAIReasoningEffort(BaseReasoningLLMTests):
def get_base_completion_call_args(self):
return {
"model": "xai/grok-3-mini-beta",
"model": "xai/grok-4.7",
"messages": [{"role": "user", "content": "Hello"}],
}

View file

@ -2863,7 +2863,7 @@ def test_gemini_function_call_parameter_in_messages():
mock_client.return_value = mock_response
try:
completion(
model="vertex_ai/gemini-2.0-flash",
model="vertex_ai/gemini-2.5-flash-preview-09-2025",
messages=messages,
tools=tools,
tool_choice="auto",
@ -3263,7 +3263,7 @@ def test_vertex_anthropic_completion():
client, "post", side_effect=vertex_ai_anthropic_thinking_mock_response
):
response = completion(
model="vertex_ai/claude-3-7-sonnet@20250219",
model="vertex_ai/claude-sonnet-4-6@default",
messages=[{"role": "user", "content": "Hello, world!"}],
vertex_ai_location="us-east5",
vertex_ai_project="test-project",
@ -3271,7 +3271,7 @@ def test_vertex_anthropic_completion():
client=client,
)
print(response)
assert response.model == "claude-3-7-sonnet@20250219"
assert response.model == "claude-sonnet-4-6@default"
assert response._hidden_params["response_cost"] is not None
assert response._hidden_params["response_cost"] > 0

View file

@ -445,7 +445,7 @@ def test_groq_response_cost_tracking(is_streaming):
response_cost = litellm.response_cost_calculator(
response_object=response,
model="groq/llama-3.3-70b-versatile",
model="groq/openai/gpt-oss-120b",
custom_llm_provider="groq",
call_type=CallTypes.acompletion.value,
optional_params={},
@ -515,7 +515,7 @@ def test_gemini_completion_cost(provider):
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model_name = "gemini-2.0-flash"
model_name = "gemini-3.8-flash"
prompt_tokens = 128.0
output_tokens = 228.0
## GET MODEL FROM LITELLM.MODEL_INFO
@ -543,7 +543,7 @@ def test_vertex_ai_completion_cost():
prompt_tokens = 100
model_info = litellm.get_model_info(model="gemini-2.0-flash")
model_info = litellm.get_model_info(model="gemini-3.8-flash")
print("\nExpected model info:\n{}\n\n".format(model_info))
@ -551,7 +551,7 @@ def test_vertex_ai_completion_cost():
## CALCULATED COST
calculated_input_cost, calculated_output_cost = cost_per_token(
model="gemini-2.0-flash",
model="gemini-3.8-flash",
custom_llm_provider="vertex_ai",
prompt_tokens=prompt_tokens,
completion_tokens=0,
@ -676,7 +676,7 @@ async def test_completion_cost_hidden_params(sync_mode):
def test_vertex_ai_gemini_predict_cost():
model = "gemini-2.0-flash"
model = "gemini-3.8-flash"
messages = [{"role": "user", "content": "Hey, hows it going???"}]
predictive_cost = completion_cost(model=model, messages=messages)
@ -757,24 +757,24 @@ def test_completion_cost_tts(model):
def test_completion_cost_anthropic():
"""
model_name: claude-3-haiku-20240307
model_name: claude-haiku-4-5
litellm_params:
model: anthropic/claude-3-haiku-20240307
model: anthropic/claude-haiku-4-5
max_tokens: 4096
"""
router = litellm.Router(
model_list=[
{
"model_name": "claude-3-haiku-20240307",
"model_name": "claude-haiku-4-5",
"litellm_params": {
"model": "anthropic/claude-3-haiku-20240307",
"model": "anthropic/claude-haiku-4-5",
"max_tokens": 4096,
},
}
]
)
data = {
"model": "claude-3-haiku-20240307",
"model": "claude-haiku-4-5",
"prompt_tokens": 21,
"completion_tokens": 20,
"response_time_ms": 871.7040000000001,
@ -2068,14 +2068,14 @@ def test_completion_cost_params():
"""
litellm.set_verbose = True
resp1_prompt_cost, resp1_completion_cost = cost_per_token(
model="gemini-2.0-flash",
model="gemini-3.8-flash",
prompt_tokens=1000,
completion_tokens=1000,
custom_llm_provider="vertex_ai_beta",
)
resp2_prompt_cost, resp2_completion_cost = cost_per_token(
model="gemini-2.0-flash", prompt_tokens=1000, completion_tokens=1000
model="gemini-3.8-flash", prompt_tokens=1000, completion_tokens=1000
)
assert resp2_prompt_cost > 0
@ -2084,7 +2084,7 @@ def test_completion_cost_params():
assert resp1_completion_cost == resp2_completion_cost
resp3_prompt_cost, resp3_completion_cost = cost_per_token(
model="vertex_ai/gemini-2.0-flash", prompt_tokens=1000, completion_tokens=1000
model="vertex_ai/gemini-3.8-flash", prompt_tokens=1000, completion_tokens=1000
)
assert resp3_prompt_cost > 0
@ -2102,14 +2102,14 @@ def test_completion_cost_params_2():
prompt_tokens = 1000
completion_tokens = 1000
resp1_prompt_cost, resp1_completion_cost = cost_per_token(
model="gemini-2.0-flash",
model="gemini-3.8-flash",
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
)
print(resp1_prompt_cost, resp1_completion_cost)
model_info = litellm.get_model_info("gemini-2.0-flash")
model_info = litellm.get_model_info("gemini-3.8-flash")
input_cost_per_token = model_info["input_cost_per_token"]
output_cost_per_token = model_info["output_cost_per_token"]
@ -2148,7 +2148,7 @@ def test_completion_cost_params_gemini_3():
)
],
created=1728529259,
model="gemini-2.0-flash",
model="gemini-3.8-flash",
object="chat.completion",
system_fingerprint=None,
usage=usage,
@ -2172,7 +2172,7 @@ def test_completion_cost_params_gemini_3():
pc, cc = cost_per_character(
**{
"model": "gemini-2.0-flash",
"model": "gemini-3.8-flash",
"custom_llm_provider": "vertex_ai",
"prompt_characters": None,
"completion_characters": 3,
@ -2180,9 +2180,9 @@ def test_completion_cost_params_gemini_3():
}
)
model_info = litellm.get_model_info("gemini-2.0-flash")
model_info = litellm.get_model_info("gemini-3.8-flash")
# gemini-2.0-flash has no per-character pricing, so cost_per_character
# gemini-3.8-flash has no per-character pricing, so cost_per_character
# falls back to per-token pricing using usage.prompt_tokens / usage.completion_tokens
assert round(pc, 10) == round(3771 * model_info["input_cost_per_token"], 10)
assert round(cc, 10) == round(
@ -2239,16 +2239,16 @@ async def test_test_completion_cost_gpt4o_audio_output_from_model(stream):
)
],
created=1729282652,
model="gpt-4o-audio-preview",
model="gpt-audio-1.5",
object="chat.completion",
system_fingerprint="fp_4eafc16e9d",
usage=usage_object,
service_tier=None,
)
cost = completion_cost(completion, model="gpt-4o-audio-preview")
cost = completion_cost(completion, model="gpt-audio-1.5")
model_info = litellm.get_model_info("gpt-4o-audio-preview")
model_info = litellm.get_model_info("gpt-audio-1.5")
print(f"model_info: {model_info}")
## input cost
@ -2517,7 +2517,7 @@ def test_cost_calculator_with_base_model():
resp = litellm.completion(
model="bedrock/random-model",
messages=[{"role": "user", "content": "Hello, how are you?"}],
base_model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
base_model="bedrock/anthropic.claude-sonnet-5",
mock_response="Hello, how are you?",
)
assert resp.model == "random-model"
@ -2551,10 +2551,10 @@ def test_cost_calculator_with_base_model_with_router(base_model_arg):
if base_model_arg == "litellm_param":
model_item["litellm_params"][
"base_model"
] = "bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
] = "bedrock/anthropic.claude-sonnet-5"
elif base_model_arg == "model_info":
model_item["model_info"] = {
"base_model": "bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
"base_model": "bedrock/anthropic.claude-sonnet-5",
}
router = Router(model_list=[model_item])

View file

@ -1148,7 +1148,7 @@ def test_openai_gateway_timeout_error():
@pytest.mark.parametrize(
"provider, model, call_type",
[
("anthropic", "claude-3-haiku-20240307", "chat_completion"),
("anthropic", "claude-haiku-4-5-20251001", "chat_completion"),
],
)
@pytest.mark.asyncio

View file

@ -136,7 +136,7 @@ def trade(model_name: str) -> List[Trade]: # type: ignore
@pytest.mark.parametrize(
"model", ["claude-haiku-4-5-20251001", "anthropic.claude-3-haiku-20240307-v1:0"]
"model", ["claude-haiku-4-5-20251001", "us.anthropic.claude-haiku-4-5-20251001-v1:0"]
)
@pytest.mark.flaky(retries=6, delay=10)
def test_function_call_parsing(model):

View file

@ -67,7 +67,17 @@ def test_get_llm_provider_deepseek_custom_api_base():
os.environ.pop("DEEPSEEK_API_BASE")
def test_get_llm_provider_vertex_ai_image_models():
def test_get_llm_provider_vertex_ai_image_models(monkeypatch):
monkeypatch.setattr(litellm, "vertex_ai_image_models", set())
monkeypatch.setattr(litellm, "models_by_provider", dict(litellm.models_by_provider))
litellm.add_known_models(
model_cost_map={
"vertex_ai/imagegeneration@006": {
"litellm_provider": "vertex_ai-image-models",
"mode": "image_generation",
}
}
)
model, custom_llm_provider, dynamic_api_key, api_base = litellm.get_llm_provider(
model="imagegeneration@006", custom_llm_provider=None
)
@ -101,17 +111,17 @@ def test_get_llm_provider_ai21_chat_test2():
def test_get_llm_provider_cohere_chat_test2():
"""
if user prefix with cohere/ but calls command-r-plus then it should be cohere_chat provider
if user prefix with cohere/ but calls command-r-plus-08-2024 then it should be cohere_chat provider
"""
model, custom_llm_provider, dynamic_api_key, api_base = litellm.get_llm_provider(
model="cohere/command-r-plus",
model="cohere/command-r-plus-08-2024",
)
print("model=", model)
print("custom_llm_provider=", custom_llm_provider)
print("api_base=", api_base)
assert custom_llm_provider == "cohere_chat"
assert model == "command-r-plus"
assert model == "command-r-plus-08-2024"
def test_get_llm_provider_azure_o1():

View file

@ -16,7 +16,7 @@ def test_get_model_info_simple_model_name():
"""
tests if model name given, and model exists in model info - the object is returned
"""
model = "claude-3-opus-20240229"
model = "claude-opus-5-5"
litellm.get_model_info(model)
@ -24,7 +24,7 @@ def test_get_model_info_custom_llm_with_model_name():
"""
Tests if {custom_llm_provider}/{model_name} name given, and model exists in model info, the object is returned
"""
model = "anthropic/claude-3-opus-20240229"
model = "anthropic/claude-opus-5-5"
litellm.get_model_info(model)

View file

@ -28,7 +28,7 @@ async def test_get_available_deployments():
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {"model": "groq/llama-3.1-8b-instant"},
"litellm_params": {"model": "groq/openai/gpt-oss-20b"},
"model_info": {"id": "groq-llama"},
},
]

View file

@ -31,7 +31,7 @@ async def test_openai_moderation_error_raising(monkeypatch):
from unittest.mock import AsyncMock, MagicMock
from litellm.types.llms.openai import OpenAIModerationResponse
litellm.openai_moderations_model_name = "text-moderation-latest"
litellm.openai_moderations_model_name = "omni-moderation-latest"
openai_mod = _ENTERPRISE_OpenAI_Moderation()
_api_key = "sk-12345"
_api_key = hash_token("sk-12345")
@ -41,9 +41,9 @@ async def test_openai_moderation_error_raising(monkeypatch):
llm_router = litellm.Router(
model_list=[
{
"model_name": "text-moderation-latest",
"model_name": "omni-moderation-latest",
"litellm_params": {
"model": "text-moderation-latest",
"model": "omni-moderation-latest",
"api_key": os.environ.get("OPENAI_API_KEY", "fake-key"),
},
}

View file

@ -188,7 +188,7 @@ def test_router_get_model_info_wildcard_routes():
]
)
model_info = router.get_router_model_info(
deployment=None, received_model_name="gemini/gemini-1.5-flash", id="1"
deployment=None, received_model_name="gemini/gemini-2.5-flash", id="1"
)
print(model_info)
assert model_info is not None
@ -212,7 +212,7 @@ async def test_router_get_model_group_usage_wildcard_routes():
)
resp = await router.acompletion(
model="gemini/gemini-1.5-flash",
model="gemini/gemini-2.5-flash",
messages=[{"role": "user", "content": "Hello, how are you?"}],
mock_response="Hello, I'm good.",
)
@ -220,7 +220,7 @@ async def test_router_get_model_group_usage_wildcard_routes():
await asyncio.sleep(2)
tpm, rpm = await router.get_model_group_usage(model_group="gemini/gemini-1.5-flash")
tpm, rpm = await router.get_model_group_usage(model_group="gemini/gemini-2.5-flash")
assert tpm is not None, "tpm is None"
assert rpm is not None, "rpm is None"
@ -242,7 +242,7 @@ async def test_call_router_callbacks_on_success():
router.cache, "async_increment_cache_pipeline", new=AsyncMock()
) as mock_callback:
await router.acompletion(
model="gemini/gemini-1.5-flash",
model="gemini/gemini-2.5-flash",
messages=[{"role": "user", "content": "Hello, how are you?"}],
mock_response="Hello, I'm good.",
)
@ -255,12 +255,12 @@ async def test_call_router_callbacks_on_success():
for increment in increment_list:
if "tpm" in increment["key"]:
assert increment["key"].startswith(
"global_router:1:gemini/gemini-1.5-flash:tpm"
"global_router:1:gemini/gemini-2.5-flash:tpm"
)
assert increment["increment_value"] == 30
elif "rpm" in increment["key"]:
assert increment["key"].startswith(
"global_router:1:gemini/gemini-1.5-flash:rpm"
"global_router:1:gemini/gemini-2.5-flash:rpm"
)
assert increment["increment_value"] == 1
@ -283,7 +283,7 @@ async def test_call_router_callbacks_on_failure():
) as mock_callback:
with pytest.raises(litellm.RateLimitError):
await router.acompletion(
model="gemini/gemini-1.5-flash",
model="gemini/gemini-2.5-flash",
messages=[{"role": "user", "content": "Hello, how are you?"}],
mock_response="litellm.RateLimitError",
num_retries=0,
@ -295,7 +295,7 @@ async def test_call_router_callbacks_on_failure():
assert (
mock_callback.call_args_list[0]
.kwargs["key"]
.startswith("global_router:1:gemini/gemini-1.5-flash:rpm")
.startswith("global_router:1:gemini/gemini-2.5-flash:rpm")
)
@ -317,7 +317,7 @@ async def test_router_model_group_headers():
for _ in range(2):
resp = await router.acompletion(
model="gemini/gemini-1.5-flash",
model="gemini/gemini-2.5-flash",
messages=[{"role": "user", "content": "Hello, how are you?"}],
mock_response="Hello, I'm good.",
)
@ -325,7 +325,7 @@ async def test_router_model_group_headers():
assert (
resp._hidden_params["additional_headers"]["x-litellm-model-group"]
== "gemini/gemini-1.5-flash"
== "gemini/gemini-2.5-flash"
)
assert "x-ratelimit-remaining-requests" in resp._hidden_params["additional_headers"]
@ -349,7 +349,7 @@ async def test_get_remaining_model_group_usage():
)
for _ in range(2):
resp = await router.acompletion(
model="gemini/gemini-1.5-flash",
model="gemini/gemini-2.5-flash",
messages=[{"role": "user", "content": "Hello, how are you?"}],
mock_response="Hello, I'm good.",
)
@ -363,7 +363,7 @@ async def test_get_remaining_model_group_usage():
await asyncio.sleep(1)
remaining_usage = await router.get_remaining_model_group_usage(
model_group="gemini/gemini-1.5-flash"
model_group="gemini/gemini-2.5-flash"
)
assert remaining_usage is not None
assert "x-ratelimit-remaining-requests" in remaining_usage

View file

@ -38,7 +38,7 @@ async def test_spend_calc_model_on_router_messages():
{
"model_name": "special-llama-model",
"litellm_params": {
"model": "groq/llama-3.1-8b-instant",
"model": "groq/openai/gpt-oss-20b",
},
}
]
@ -81,7 +81,7 @@ async def test_spend_calc_using_response():
}
],
"created": "1677652288",
"model": "groq/llama-3.1-8b-instant",
"model": "groq/openai/gpt-oss-20b",
"object": "chat.completion",
"system_fingerprint": "fp_873a560973",
"usage": {

View file

@ -31,14 +31,14 @@
"model_id": null,
"cache_key": null,
"api_base": null,
"response_cost": 7.5e-06,
"response_cost": 3.5e-05,
"additional_headers": {},
"litellm_overhead_time_ms": null,
"batch_models": null,
"litellm_model_name": "vertex_ai/gemini-2.0-flash-001",
"litellm_model_name": "vertex_ai/gemini-3-flash-preview",
"usage_object": null
},
"litellm_response_cost": 7.5e-06,
"litellm_response_cost": 3.5e-05,
"cache_hit": false,
"requester_metadata": {}
},
@ -54,13 +54,13 @@
"id": "time-14-15-40-349639_chatcmpl-59a988d0-7ef1-4dc4-bc18-d2e78961817f",
"endTime": "2025-05-26T14:15:40.607266-07:00",
"completionStartTime": "2025-05-26T14:15:40.607266-07:00",
"model": "gemini-2.0-flash-001",
"model": "gemini-3-flash-preview",
"modelParameters": {},
"usage": {
"input": 10,
"output": 10,
"unit": "TOKENS",
"totalCost": 7.5e-06
"totalCost": 3.5e-05
},
"usageDetails": {
"input": 10,

View file

@ -582,7 +582,7 @@ async def test_webhook_alerting(alerting_type):
None,
None,
),
("gemini-2.0-flash", None, "vertex_ai", "hardy-device-38811", "us-central1"),
("gemini-3.8-flash", None, "vertex_ai", "hardy-device-38811", "us-central1"),
],
)
@pytest.mark.parametrize("error_code", [500, 408, 400])
@ -688,7 +688,7 @@ async def test_outage_alerting_called(
None,
None,
),
("gemini-2.0-flash", None, "vertex_ai", "hardy-device-38811", "us-central1"),
("gemini-3.8-flash", None, "vertex_ai", "hardy-device-38811", "us-central1"),
],
)
@pytest.mark.parametrize("error_code", [500, 408, 400])
@ -775,7 +775,7 @@ async def test_region_outage_alerting_called(
await slack_alerting.region_outage_alerts(
exception=error_to_raise, deployment_id=deployment_id # type: ignore
)
if model == "gemini-2.0-flash" and (error_code == 500 or error_code == 408):
if model == "gemini-3.8-flash" and (error_code == 500 or error_code == 408):
mock_send_alert.assert_called_once()
else:
mock_send_alert.assert_not_called()

View file

@ -481,12 +481,12 @@ class TestLangfuseLogging:
completion_tokens=10,
total_tokens=20,
),
model="vertex/gemini-2.0-flash-001",
model="vertex/gemini-3-flash-preview",
object="chat.completion",
created=1723081200,
).model_dump()
await litellm.acompletion(
model="vertex_ai/gemini-2.0-flash-001",
model="vertex_ai/gemini-3-flash-preview",
messages=[{"role": "user", "content": "Hello!"}],
mock_response=mock_response,
metadata={"trace_id": setup["trace_id"]},

View file

@ -100,6 +100,7 @@ async def test_daily_tag_spend_retries_then_succeeds():
1,
]
)
prisma_client.db.tx.return_value.__aenter__.return_value.execute_raw = prisma_client.db.execute_raw
daily_spend_transactions: Dict[str, DailyTagSpendTransaction] = {
"k": {

View file

@ -1273,9 +1273,7 @@ async def test_combined_prefix_reflects_in_s3_object_key():
assert "myteam/apikey/" in key, f"Expected both prefixes in key: {key}"
def test_s3_object_key_sanitizes_slashes_in_file_name():
"""Response ids containing slashes (e.g. bedrock batch job ARNs) must not
create nested S3 folders; only path/prefix/date slashes are separators."""
def test_s3_object_key_sanitizes_slashes_and_colons_in_file_name():
from litellm.integrations.s3 import get_s3_object_key
start_time = datetime(2026, 2, 11, 0, 35, 18, 391582)
@ -1290,10 +1288,32 @@ def test_s3_object_key_sanitizes_slashes_in_file_name():
assert key == (
"LiteLLMAPPLogs/myteam/2026-02-11/"
"time-00-35-18-391582_arn:aws:bedrock:us-east-1:123456789012:model-invocation-job_gl18r6skk9yy.json"
"time-00-35-18-391582_arn_aws_bedrock_us-east-1_123456789012_model-invocation-job_gl18r6skk9yy.json"
)
@pytest.mark.parametrize(
"response_id",
[
"s3://example-batch-bucket/litellm-bedrock-files/input.jsonl",
"gs://example-batch-bucket/litellm-vertex-files/input.jsonl",
],
)
def test_s3_object_key_has_no_colon_for_cloud_uri_file_ids(response_id: str):
from litellm.integrations.s3 import get_s3_object_key
key = get_s3_object_key(
s3_path="",
prefix="",
start_time=datetime(2026, 9, 7, 4, 51, 6, 685889),
s3_file_name=f"time-04-51-06-685889_{response_id}",
)
filename = key.rsplit("/", 1)[-1]
assert ":" not in filename
assert filename.endswith("_input.jsonl.json")
def test_create_s3_batch_logging_element_flat_key_for_arn_response_id():
"""End-to-end through the s3_v2 element builder: an ARN response id must
yield a flat file directly under the date segment."""

View file

@ -0,0 +1,93 @@
import json
import pytest
import litellm
from litellm.litellm_core_utils.llm_response_utils import get_api_base as get_api_base_module
from litellm.llms.chatgpt.common_utils import CHATGPT_API_BASE
from litellm.llms.github_copilot.common_utils import DEFAULT_GITHUB_COPILOT_API_BASE
@pytest.fixture
def isolated_token_dirs(tmp_path, monkeypatch):
monkeypatch.setenv("GITHUB_COPILOT_TOKEN_DIR", str(tmp_path / "github_copilot"))
monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path / "chatgpt"))
monkeypatch.delenv("GITHUB_COPILOT_API_BASE", raising=False)
monkeypatch.delenv("CHATGPT_API_BASE", raising=False)
monkeypatch.delenv("OPENAI_CHATGPT_API_BASE", raising=False)
return tmp_path
@pytest.fixture
def resolution_lookups(monkeypatch):
lookups: list = []
def _record(*args, **kwargs):
lookups.append((args, kwargs))
raise RuntimeError("provider resolution must not run for an authenticating provider")
monkeypatch.setattr(get_api_base_module, "get_llm_provider", _record)
return lookups
class TestDeclaredAuthenticatingProvider:
"""get_llm_provider runs the OAuth device flow for github_copilot and chatgpt, and get_api_base
runs on every response's hidden params and on every mapped exception, so it must answer from
the declaration without resolving. The recorder appends before raising, and get_api_base
swallows resolver errors, so an empty list proves the lookup never ran."""
@pytest.mark.parametrize(
"model, custom_llm_provider, expected",
[
("github_copilot/gpt-4o", None, DEFAULT_GITHUB_COPILOT_API_BASE),
("gpt-4o", "github_copilot", DEFAULT_GITHUB_COPILOT_API_BASE),
("chatgpt/gpt-5", None, CHATGPT_API_BASE),
("gpt-5", "chatgpt", CHATGPT_API_BASE),
],
)
def test_answers_without_resolving(
self, model, custom_llm_provider, expected, isolated_token_dirs, resolution_lookups
):
api_base = litellm.get_api_base(model=model, optional_params={"custom_llm_provider": custom_llm_provider})
assert resolution_lookups == []
assert api_base == expected
def test_copilot_keeps_the_enterprise_endpoint_from_disk(self, isolated_token_dirs, resolution_lookups):
token_dir = isolated_token_dirs / "github_copilot"
token_dir.mkdir()
(token_dir / "api-key.json").write_text(
json.dumps({"endpoints": {"api": "https://api.enterprise.githubcopilot.com"}})
)
api_base = litellm.get_api_base(model="github_copilot/gpt-4o", optional_params={})
assert resolution_lookups == []
assert api_base == "https://api.enterprise.githubcopilot.com"
def test_explicit_api_base_still_wins(self, isolated_token_dirs, resolution_lookups):
api_base = litellm.get_api_base(
model="github_copilot/gpt-4o", optional_params={"api_base": "https://copilot.example/v1"}
)
assert resolution_lookups == []
assert api_base == "https://copilot.example/v1"
def test_other_providers_still_resolve(self, isolated_token_dirs, resolution_lookups):
litellm.get_api_base(model="openai/gpt-4o", optional_params={})
assert len(resolution_lookups) == 1
@pytest.mark.parametrize(
"model, expected",
[
("gemini/gemini-2.5-pro", "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent"),
("openai/gpt-4o", "https://api.openai.com"),
],
)
def test_providers_with_a_fixed_base_still_get_it(model, expected, monkeypatch):
for env in ("GEMINI_API_BASE", "OPENAI_API_BASE", "OPENAI_BASE_URL"):
monkeypatch.delenv(env, raising=False)
assert litellm.get_api_base(model=model, optional_params={}) == expected

View file

@ -21,7 +21,9 @@ class _RecordingClientWebSocket:
@pytest.mark.asyncio
async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_policy_close():
import websockets
from websockets.datastructures import Headers
from websockets.exceptions import InvalidStatus
from websockets.http11 import Response
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
from litellm.types.realtime import RealtimeErrorEvent
@ -32,9 +34,7 @@ async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_
dummy_websocket = _RecordingClientWebSocket()
dummy_logging_obj = MagicMock()
refused = websockets.exceptions.InvalidStatus(
websockets.http11.Response(401, "Unauthorized", websockets.datastructures.Headers())
)
refused = InvalidStatus(Response(401, "Unauthorized", Headers()))
with patch("websockets.connect", side_effect=refused):
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # handler's websocket param is a Protocol here but the mock connect type is incomplete

View file

@ -19,9 +19,8 @@ from unittest.mock import MagicMock, patch
import httpx
import pytest
from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
from litellm.types.utils import LiteLLMBatch, LlmProviders
from litellm.types.utils import LlmProviders
# AWS JobStatus -> OpenAI BatchJobStatus, exactly as encoded in transformation.py
# (both transform_create_batch_response and transform_retrieve_batch_response).
@ -270,6 +269,44 @@ def test_create_request_keeps_kms_key_alongside_s3_bucket_owner(config, monkeypa
}
def test_create_request_omits_kms_key_when_env_var_is_blank(config, monkeypatch):
monkeypatch.setenv("AWS_S3_ENCRYPTION_KEY_ID", "")
monkeypatch.delenv("AWS_S3_BUCKET_OWNER", raising=False)
bedrock_request = _signed_batch_request(config, {}, {})
assert bedrock_request["outputDataConfig"] == {
"s3OutputDataConfig": {"s3Uri": "s3://in-bucket/litellm-batch-outputs/litellm-batch-1/"}
}
def test_create_request_omits_s3_bucket_owner_when_env_var_is_blank(config, monkeypatch):
monkeypatch.setenv("AWS_S3_BUCKET_OWNER", "")
monkeypatch.delenv("AWS_S3_ENCRYPTION_KEY_ID", raising=False)
bedrock_request = _signed_batch_request(config, {}, {})
assert bedrock_request["inputDataConfig"] == {"s3InputDataConfig": {"s3Uri": "s3://in-bucket/in.jsonl"}}
assert bedrock_request["outputDataConfig"] == {
"s3OutputDataConfig": {"s3Uri": "s3://in-bucket/litellm-batch-outputs/litellm-batch-1/"}
}
def test_create_request_emits_real_values_alongside_blank_sibling_env_var(config, monkeypatch):
monkeypatch.setenv("AWS_S3_ENCRYPTION_KEY_ID", "kms-key-123")
monkeypatch.setenv("AWS_S3_BUCKET_OWNER", "")
bedrock_request = _signed_batch_request(config, {}, {})
assert bedrock_request["inputDataConfig"] == {"s3InputDataConfig": {"s3Uri": "s3://in-bucket/in.jsonl"}}
assert bedrock_request["outputDataConfig"] == {
"s3OutputDataConfig": {
"s3Uri": "s3://in-bucket/litellm-batch-outputs/litellm-batch-1/",
"s3EncryptionKeyId": "kms-key-123",
}
}
def test_create_request_missing_input_file_id_raises(config):
with pytest.raises(ValueError, match="input_file_id is required"):
config.transform_create_batch_request(

View file

@ -1,4 +1,7 @@
import base64
import io
import json
import sys
from litellm._uuid import uuid
from unittest.mock import MagicMock, patch
@ -544,3 +547,74 @@ async def test_ollama_async_completion_inlines_remote_images_off_the_event_loop(
assert response.choices[0].message.content == "Green"
assert async_only_image_fetch.fetched == [image_url]
assert captured["body"]["images"] == [async_only_image_fetch.base64_png]
def _image_base64(image_format: str) -> str:
from PIL import Image
buffer = io.BytesIO()
Image.new("RGB", (4, 4), "green").save(buffer, image_format)
return base64.b64encode(buffer.getvalue()).decode("utf-8")
def _transform_image_request(image_base64: str, mime_subtype: str) -> dict:
return OllamaConfig().transform_request(
model="llava",
messages=[
{
"role": "user",
"content": [
{"type": "text", "text": "What colour is this?"},
{
"type": "image_url",
"image_url": {"url": f"data:image/{mime_subtype};base64,{image_base64}"},
},
],
}
],
optional_params={},
litellm_params={},
headers={},
)
@pytest.mark.parametrize("image_format", ["PNG", "JPEG"])
def test_transform_request_sends_png_and_jpeg_images_without_pillow(
image_format: str, monkeypatch: pytest.MonkeyPatch
) -> None:
image_base64 = _image_base64(image_format)
monkeypatch.setitem(sys.modules, "PIL", None)
data = _transform_image_request(image_base64, image_format.lower())
assert data["images"] == [image_base64]
def test_transform_request_without_pillow_says_how_to_convert_other_image_formats(
monkeypatch: pytest.MonkeyPatch,
) -> None:
gif_base64 = _image_base64("GIF")
monkeypatch.setitem(sys.modules, "PIL", None)
with pytest.raises(Exception, match="pip install Pillow"):
_transform_image_request(gif_base64, "gif")
def test_transform_request_reencodes_other_image_formats_as_jpeg() -> None:
from PIL import Image
data = _transform_image_request(_image_base64("GIF"), "gif")
(encoded,) = data["images"]
assert Image.open(io.BytesIO(base64.b64decode(encoded))).format == "JPEG"
@pytest.mark.parametrize(
"payload",
[base64.b64encode(b"not an image").decode("utf-8"), "abc"],
ids=["decodable_but_not_an_image", "invalid_base64"],
)
def test_transform_request_leaves_unreadable_images_untouched(payload: str) -> None:
data = _transform_image_request(payload, "png")
assert data["images"] == [payload]

View file

@ -422,7 +422,9 @@ async def test_async_realtime_ws_url_has_no_ssl():
async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_policy_close():
from typing import cast
import websockets
from websockets.datastructures import Headers
from websockets.exceptions import InvalidStatus
from websockets.http11 import Response
from litellm.llms.openai.realtime.handler import OpenAIRealtime
from litellm.types.realtime import RealtimeErrorEvent
@ -445,9 +447,7 @@ async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_
dummy_websocket = RecordingClientWebSocket()
dummy_logging_obj = MagicMock()
refused = websockets.exceptions.InvalidStatus(
websockets.http11.Response(401, "Unauthorized", websockets.datastructures.Headers())
)
refused = InvalidStatus(Response(401, "Unauthorized", Headers()))
with patch("websockets.connect", side_effect=refused):
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # handler's websocket param is Any

View file

@ -6529,6 +6529,53 @@ class TestMCPDcrBridgeDelegateAdmission:
assert exc_info.value.status_code == 503
async def test_reload_admitted_key_returns_admin_for_master_key_hash(self):
"""An envelope sealed under the master key has no DB row to reload; the reload resolves it
to the PROXY_ADMIN auth context (api_key is the alias, never the hash) rather than failing.
A hash that is NOT the master key's still reaches the prisma gate and fails the same as
before (500 with no database connection)."""
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy._types import LitellmUserRoles, hash_token
with (
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
admitted = await MCPRequestHandler._reload_admitted_key(hash_token(self._MASTER_KEY))
assert admitted.user_role == LitellmUserRoles.PROXY_ADMIN
assert admitted.api_key == LITELLM_PROXY_MASTER_KEY_ALIAS
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler._reload_admitted_key("not-the-master-hash")
assert exc_info.value.status_code == 500
@pytest.mark.parametrize(
"flag_enabled, scope, expected",
[(True, "scoped", []), (False, "scoped", ["public"]), (True, "unscoped", ["public"])],
)
async def test_master_envelope_respects_allow_all_scope(self, flag_enabled, scope, expected):
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.proxy._types import hash_token
manager = MCPServerManager()
with patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY):
admitted = await MCPRequestHandler._reload_admitted_key(hash_token(self._MASTER_KEY))
with (
patch.object(manager, "get_allow_all_keys_server_ids", return_value=["public"]),
patch.object(manager, "_get_active_submitted_mcp_server_ids_for_user", new=AsyncMock(return_value=[])),
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_user",
new=AsyncMock(return_value=["granted"] if scope == "scoped" else []),
),
):
servers = await manager.get_allowed_mcp_servers(
admitted,
access=MCPServerAccess(server_ids=(), scope=scope),
general_settings={"mcp_allow_all_keys_respects_mcp_scope": flag_enabled},
)
assert servers == expected
async def test_envelope_for_key_barred_from_mcp_routes_is_rejected_403(self):
"""A key whose allowed_routes exclude MCP must not reach tools via an envelope: the arm runs
RouteChecks.should_call_route before admitting, exactly as the standard pipeline does between
@ -6959,10 +7006,12 @@ class TestMCPDcrBridgeDelegateAdmission:
assert exc_info.value.status_code == 403
assert not exc_info.value.headers
async def test_explicit_litellm_key_wins_over_envelope_arm(self):
"""An explicit x-litellm-api-key is always a LiteLLM credential and its arm precedes the
envelope arm: user_api_key_auth validates the key and NO inner token is injected, even
though the Authorization header carries a valid envelope."""
async def test_explicit_litellm_key_matching_envelope_admits_under_explicit_key(self):
"""The dual-credential arm: an explicit x-litellm-api-key paired with an envelope sealing the
SAME key hash admits under the explicit key's auth context AND injects the sealed upstream
token for egress. When the envelope seals a different principal the request is a 403 instead
(covered by the mismatch tests), and the explicit key never silently drops the envelope the
way the pre-fix ordering did."""
envelope = self._mint_bridge_envelope()
scope = {
"type": "http",
@ -6975,7 +7024,7 @@ class TestMCPDcrBridgeDelegateAdmission:
}
async def mock_user_api_key_auth(api_key, request):
return UserAPIKeyAuth(api_key=api_key, user_id="litellm-key-user")
return UserAPIKeyAuth(api_key=self._KEY_HASH, user_id="litellm-key-user")
with (
patch(
@ -6997,9 +7046,10 @@ class TestMCPDcrBridgeDelegateAdmission:
mock_auth.assert_called_once()
assert mock_auth.call_args.kwargs["api_key"] == "Bearer sk-explicit-litellm-key"
# The explicit-key arm admitted; the envelope arm never ran, so no inner token is injected.
assert auth_result.user_id == "litellm-key-user"
assert mcp_server_auth_headers == {}
assert mcp_server_auth_headers == {
"bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"}
}
async def test_non_bridge_oauth_delegate_server_does_not_take_envelope_arm(self):
"""An oauth_delegate server that is NOT a DCR bridge (``dcr_bridge`` unset) must not take the
@ -7139,6 +7189,199 @@ class TestMCPDcrBridgeDelegateAdmission:
assert exc_info.value.status_code == 500
@pytest.mark.asyncio
class TestMCPDcrBridgeDualCredential:
"""Dual-credential arm: ``x-litellm-api-key`` alongside an ``llm_env_`` bearer on a
DCR-bridge ``oauth_delegate`` route (issue #38208).
Real MCP clients send their litellm key on every request, so the envelope minted at
``/{server}/token`` arrives paired with the key rather than alone. The explicit credential
is the admission context and the envelope supplies the upstream token, but only when both
name the same principal; a mismatch is a 403, an invalid envelope is the scope's
``invalid_token`` challenge, and the envelope itself never reaches egress.
"""
_DELEGATE = TestMCPDcrBridgeDelegateAdmission
_MASTER_KEY = TestMCPDcrBridgeDelegateAdmission._MASTER_KEY
_KEY_HASH = TestMCPDcrBridgeDelegateAdmission._KEY_HASH
@staticmethod
def _dual_scope(envelope: str, explicit_key: str):
return {
"type": "http",
"method": "POST",
"path": "/mcp/bridge_delegate_server",
"headers": [
(b"authorization", f"Bearer {envelope}".encode("latin-1")),
(b"x-litellm-api-key", explicit_key.encode("latin-1")),
],
}
@pytest.mark.parametrize("dual_credential", [False, True])
async def test_admission_rejects_server_without_routable_name(self, dual_credential):
envelope = self._DELEGATE._mint_bridge_envelope()
server = self._DELEGATE._bridge_delegate_server(server_name=None)
admission = (
MCPRequestHandler._admit_dcr_bridge_dual_credential(
server=server,
requested_name="bridge_delegate_server",
authorization_value=f"Bearer {envelope}",
litellm_api_key="sk-explicit-key",
mcp_server_auth_headers=None,
request=self._DELEGATE._mcp_request(),
route="/mcp/bridge_delegate_server",
)
if dual_credential
else MCPRequestHandler._admit_dcr_bridge_delegate(
server=server,
requested_name="bridge_delegate_server",
authorization_value=f"Bearer {envelope}",
mcp_server_auth_headers=None,
request=self._DELEGATE._mcp_request(),
route="/mcp/bridge_delegate_server",
)
)
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
new=AsyncMock(return_value=self._DELEGATE._reloaded_key()),
),
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
self._DELEGATE._patch_key_reload() as reload_key,
pytest.raises(HTTPException) as exc_info,
):
await admission
assert exc_info.value.status_code == 500
assert exc_info.value.detail == "Server misconfigured: MCP server has no routable name"
reload_key.assert_not_awaited()
@pytest.mark.parametrize("mapped_jwt", [False, True])
async def test_dual_credential_matching_key_admits_under_explicit_key_and_forwards_upstream_token(self, mapped_jwt):
"""The reported bug: before the fix this request validated the key and dropped the
envelope, so egress forwarded no upstream credential and the upstream 401 yielded
``tools: []``. Now the explicit key's auth context wins admission AND the sealed
upstream token is injected per-server, while the envelope bearer is scrubbed from
every egress header context."""
envelope = self._DELEGATE._mint_bridge_envelope(key_hash=self._KEY_HASH)
explicit_auth = self._DELEGATE._reloaded_key(
api_key=None if mapped_jwt else self._KEY_HASH,
token=self._KEY_HASH,
user_id=None if mapped_jwt else "explicit-key-user",
)
presented_token = "aaa.bbb.ccc" if mapped_jwt else "sk-explicit-key"
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
new_callable=AsyncMock,
return_value=explicit_auth,
) as mock_auth,
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
self._DELEGATE._patch_key_reload() as get_key_object,
):
mock_mgr.get_mcp_server_by_name.return_value = self._DELEGATE._bridge_delegate_server()
(
auth_result,
_mcp_auth_header,
_mcp_servers,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = await MCPRequestHandler.process_mcp_request(self._dual_scope(envelope, presented_token))
mock_auth.assert_awaited_once()
assert mock_auth.await_args.kwargs["api_key"] == f"Bearer {presented_token}"
assert auth_result is explicit_auth
get_key_object.assert_not_awaited()
assert mcp_server_auth_headers == {
"bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"}
}
assert oauth2_headers is None
assert all("llm_env_" not in str(v) for v in raw_headers.values())
async def test_dual_credential_principal_mismatch_is_403(self):
"""An envelope minted under one key presented alongside a different key must not admit:
the request names two different principals, so it fails closed with
``oauth_principal_mismatch`` rather than falling back onto either credential."""
envelope = self._DELEGATE._mint_bridge_envelope(key_hash=self._KEY_HASH)
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
new_callable=AsyncMock,
return_value=self._DELEGATE._reloaded_key(api_key="a-different-key-hash"),
),
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
self._DELEGATE._patch_key_reload(),
):
mock_mgr.get_mcp_server_by_name.return_value = self._DELEGATE._bridge_delegate_server()
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(self._dual_scope(envelope, "sk-other-key"))
assert exc_info.value.status_code == 403
assert exc_info.value.detail == {"error": "oauth_principal_mismatch"}
async def test_dual_credential_user_subject_envelope_matches_on_user_id(self):
"""An interactive (user_id) envelope pairs with an explicit credential whose resolved
user_id is the same user; a different user is a 403, never a silent admit."""
for presented_user, expected_status in (("sso-user-7", None), ("sso-user-9", 403)):
envelope = self._DELEGATE._mint_bridge_envelope(user_id="sso-user-7")
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
new_callable=AsyncMock,
return_value=UserAPIKeyAuth(user_id=presented_user, api_key="any-hash"),
),
patch(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
) as mock_mgr,
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
):
mock_mgr.get_mcp_server_by_name.return_value = self._DELEGATE._bridge_delegate_server()
if expected_status is not None:
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(self._dual_scope(envelope, "sk-key"))
assert exc_info.value.status_code == expected_status
assert exc_info.value.detail == {"error": "oauth_principal_mismatch"}
else:
(
auth_result,
_h,
_s,
mcp_server_auth_headers,
_o,
_r,
) = await MCPRequestHandler.process_mcp_request(self._dual_scope(envelope, "sk-key"))
assert auth_result.user_id == "sso-user-7"
assert mcp_server_auth_headers == {
"bridge_delegate_server": {"Authorization": "Bearer inner-upstream-access-token"}
}
async def test_dual_credential_invalid_envelope_is_401_challenge_not_silent_admit(self):
"""A tampered envelope next to a perfectly valid key must still fail closed with the
scope's ``invalid_token`` challenge; the explicit key alone never unlocks a bridge
server's upstream token."""
envelope = self._DELEGATE._mint_bridge_envelope(key_hash=self._KEY_HASH)
tampered = envelope[:-4] + "AAAA"
with (
patch(
"litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth",
new_callable=AsyncMock,
return_value=self._DELEGATE._reloaded_key(),
),
patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr,
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
self._DELEGATE._patch_key_reload(),
):
mock_mgr.get_mcp_server_by_name.return_value = self._DELEGATE._bridge_delegate_server()
with pytest.raises(HTTPException) as exc_info:
await MCPRequestHandler.process_mcp_request(self._dual_scope(tampered, "sk-explicit-key"))
assert exc_info.value.status_code == 401
assert "invalid_token" in str(exc_info.value.headers)
@pytest.mark.asyncio
class TestAggregateGatewayDcrChallenge:
"""The mcp_gateway_dcr front door: a 401 on the aggregate /mcp scope must

View file

@ -6732,6 +6732,270 @@ async def test_bridge_mint_unresolvable_identity_is_500_before_upstream():
post.assert_not_called()
async def _exchange_for_bridge_server_with_jwt(jwt_auth_result, upstream_body=None):
"""Drive exchange_token_with_server for a bridge oauth_delegate authorization_code request whose
presented credential is JWT-shaped, with _resolve_jwt_auth stubbed to a given result. Returns
(response, post_mock) so a test can assert the minted envelope's sealed identity or the mapped
error status."""
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import exchange_token_with_server
from litellm.types.mcp import MCPAuth
request = _bridge_mock_request()
request.headers = {"x-litellm-api-key": "aaa.bbb.ccc"}
fake_http_response = MagicMock()
fake_http_response.json.return_value = upstream_body or {
"access_token": "UP",
"token_type": "Bearer",
"expires_in": 3600,
}
fake_http_response.raise_for_status = MagicMock()
fake_http_client = MagicMock()
fake_http_client.post = AsyncMock(return_value=fake_http_response)
server = _bridge_server(auth_type=MCPAuth.oauth_delegate)
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=fake_http_client,
),
patch(
"litellm.proxy._experimental.mcp_server.bridge_token_flow._resolve_jwt_auth",
new=AsyncMock(return_value=jwt_auth_result),
),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
):
response = await exchange_token_with_server(
request=request,
mcp_server=server,
grant_type="authorization_code",
code="auth-code",
redirect_uri="https://claude.ai/api/mcp/auth_callback",
client_id="dcr-client-123",
client_secret=None,
code_verifier="verifier",
)
return response, fake_http_client.post
@pytest.mark.asyncio
async def test_bridge_mint_unmapped_jwt_is_rejected_before_upstream():
from litellm.proxy.auth.handle_jwt import JWTIdentity
response, post = await _exchange_for_bridge_server_with_jwt(
JWTIdentity(user_id="jwt-user-5", user_object=None, agent_id=None)
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_request"
post.assert_not_called()
@pytest.mark.asyncio
async def test_bridge_mint_jwt_mapped_to_virtual_key_seals_key_hash_subject():
from datetime import datetime, timezone
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import (
BridgeEnvelopeAdmitted,
envelope_keys_from_master_key,
resolve_bridge_envelope,
)
from litellm.proxy._types import UserAPIKeyAuth
response, _post = await _exchange_for_bridge_server_with_jwt(
UserAPIKeyAuth(token="mapped-key-hash-99", user_id="mapped-user")
)
assert response.status_code == 200
token = json.loads(response.body)["access_token"]
keys = envelope_keys_from_master_key(_BRIDGE_MASTER_KEY)
opened = resolve_bridge_envelope(token, keys, datetime.now(timezone.utc), "bridge_srv")
assert isinstance(opened, BridgeEnvelopeAdmitted)
assert opened.identity.subject_type == "key_hash"
assert opened.identity.subject == "mapped-key-hash-99"
@pytest.mark.asyncio
@pytest.mark.parametrize("jwt_claims", [{"client_id": "allowed"}, {"client_id": "denied"}, {}])
async def test_bridge_mint_jwt_cannot_drop_signed_client_policy(jwt_claims):
from litellm.proxy._types import UserAPIKeyAuth
settings = {
"mcp_allowed_clients": [{"alias": "Allowed", "value": "allowed"}],
"mcp_client_id_header": "x-client-id",
"litellm_jwtauth": {"mcp_client_id_jwt_field": "client_id"},
}
with patch("litellm.proxy.proxy_server.general_settings", settings):
response, post = await _exchange_for_bridge_server_with_jwt(
UserAPIKeyAuth(token="mapped-key-hash-99", jwt_claims=jwt_claims)
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_request"
assert "signed client identity" in json.loads(response.body)["error_description"]
post.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"settings",
[
{},
{"litellm_jwtauth": {"mcp_client_id_jwt_field": "client_id"}},
{
"mcp_allowed_clients": [{"alias": "Allowed", "value": "allowed"}],
"mcp_client_id_header": "x-client-id",
},
],
)
async def test_bridge_mint_mapped_jwt_without_signed_client_policy(settings):
from litellm.proxy._types import UserAPIKeyAuth
with patch("litellm.proxy.proxy_server.general_settings", settings):
response, post = await _exchange_for_bridge_server_with_jwt(
UserAPIKeyAuth(token="mapped-key-hash-99", jwt_claims={"client_id": "allowed"})
)
assert response.status_code == 200
assert json.loads(response.body)["access_token"].startswith("llm_env_")
post.assert_awaited_once()
@pytest.mark.asyncio
async def test_bridge_mint_jwt_with_no_resolved_identity_is_400_before_upstream():
"""A JWT that resolves to nothing (or to an identity with no user_id) cannot back an envelope:
the mint returns 400 invalid_request WITHOUT consuming the single-use code upstream, matching
the no-credential path rather than hashing the raw JWT string."""
response, post = await _exchange_for_bridge_server_with_jwt(None)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_request"
post.assert_not_called()
from litellm.proxy.auth.handle_jwt import JWTIdentity
response, post = await _exchange_for_bridge_server_with_jwt(
JWTIdentity(user_id=None, user_object=None, agent_id=None)
)
assert response.status_code == 400
assert json.loads(response.body)["error"] == "invalid_request"
post.assert_not_called()
def _jwt_auth_patches(mapped_key):
from contextlib import ExitStack
from litellm.proxy._types import LiteLLM_JWTAuth
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
cache = UserApiKeyCache()
cache.set_cache(key=jwt_key_mapping_cache_key("sub", "mapped-client"), value=mapped_key.token)
cache.set_cache(key=mapped_key.token, value=mapped_key)
handler = MagicMock()
handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_claim_field="sub")
handler.auth_jwt = AsyncMock(return_value={"sub": "mapped-client"})
stack = ExitStack()
stack.enter_context(patch("litellm.proxy.proxy_server.general_settings", {"enable_jwt_auth": True}))
stack.enter_context(patch("litellm.proxy.proxy_server.premium_user", True))
stack.enter_context(patch("litellm.proxy.proxy_server.prisma_client", object()))
stack.enter_context(patch("litellm.proxy.proxy_server.jwt_handler", handler))
stack.enter_context(patch("litellm.proxy.proxy_server.user_api_key_cache", cache))
stack.enter_context(
patch(
"litellm.proxy._experimental.mcp_server.bridge_token_flow._key_owner_scim_deactivated",
new=AsyncMock(return_value=False),
)
)
return stack
@pytest.mark.asyncio
async def test_jwt_mapped_to_service_account_key_without_user_id_resolves():
"""A JWT mapped to a team or service-account virtual key (no user_id) is still an active
credential: _resolve_jwt_auth returns the mapped key, and the mint seals a key_hash-subject
envelope rather than 400ing with no_identity."""
from litellm.proxy._experimental.mcp_server import bridge_token_flow
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp import MCPAuth
mapped_key = UserAPIKeyAuth(user_id=None, token="svc-key-hash-1")
request = _bridge_mock_request()
request.headers = {"x-litellm-api-key": "aaa.bbb.ccc"}
with (
_jwt_auth_patches(mapped_key),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
):
resolved = await bridge_token_flow._resolve_jwt_auth(request, "aaa.bbb.ccc", None)
assert isinstance(resolved, UserAPIKeyAuth)
assert resolved.token == mapped_key.token
assert resolved.api_key is None
assert resolved.user_id is None
mint = await bridge_token_flow._prepare_bridge_mint(
request=request,
mcp_server=_bridge_server(auth_type=MCPAuth.oauth_delegate),
)
assert isinstance(mint, bridge_token_flow._BridgeMintReady)
assert mint.identity.subject_type == "key_hash"
assert mint.identity.subject == "svc-key-hash-1"
@pytest.mark.asyncio
async def test_jwt_mapped_to_blocked_key_is_rejected():
"""The relaxed gate is still active-state gated: a JWT mapped to a blocked virtual key resolves
to None, so the mint cannot seal an envelope under it."""
from litellm.proxy._experimental.mcp_server import bridge_token_flow
from litellm.proxy._types import UserAPIKeyAuth
mapped_key = UserAPIKeyAuth(user_id=None, token="blocked-key-hash-1", blocked=True)
request = _bridge_mock_request()
request.headers = {"x-litellm-api-key": "aaa.bbb.ccc"}
with (
_jwt_auth_patches(mapped_key),
patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY),
):
resolved = await bridge_token_flow._resolve_jwt_auth(request, "aaa.bbb.ccc", None)
assert resolved is None
mint = await bridge_token_flow._prepare_bridge_mint(
request=request,
mcp_server=_bridge_server(auth_type=MCPAuth.oauth_delegate),
)
assert mint == "no_identity"
@pytest.mark.asyncio
async def test_master_key_at_token_endpoint_mints_key_hash_envelope():
"""The master key has no row in LiteLLM_VerificationTokenTable, but it is the proxy's root
credential: presented at the bridge /token endpoint it must mint a key_hash-subject envelope
(sealed under hash_token(master_key)) even with no database connection at all. A presented key
that is NOT the master key still hits the unresolvable gate when prisma is down, unchanged."""
from litellm.proxy._experimental.mcp_server import bridge_token_flow
from litellm.proxy._types import hash_token
from litellm.types.mcp import MCPAuth
master = "sk-test-master-key-mint-0000"
request = _bridge_mock_request()
request.headers = {"x-litellm-api-key": master}
with (
patch("litellm.proxy.proxy_server.master_key", master),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
mint = await bridge_token_flow._prepare_bridge_mint(
request=request,
mcp_server=_bridge_server(auth_type=MCPAuth.oauth_delegate),
)
assert isinstance(mint, bridge_token_flow._BridgeMintReady)
assert mint.identity.subject_type == "key_hash"
assert mint.identity.subject == hash_token(master)
other = _bridge_mock_request()
other.headers = {"x-litellm-api-key": "sk-not-the-master-key"}
with (
patch("litellm.proxy.proxy_server.master_key", master),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
mint = await bridge_token_flow._prepare_bridge_mint(
request=other,
mcp_server=_bridge_server(auth_type=MCPAuth.oauth_delegate),
)
assert mint == "identity_unresolvable"
@pytest.mark.asyncio
async def test_bridge_mint_upstream_expired_lifetime_is_502():
"""An upstream token response reporting an already-elapsed lifetime (a parseable non-positive

View file

@ -1,6 +1,8 @@
"""Tests for the single-statement daily spend upsert (LIT-5291)."""
import re
from collections.abc import AsyncIterator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
import pytest
@ -149,6 +151,13 @@ class _RecordingDb:
self.statements.append((query, args))
return len(args)
@asynccontextmanager
async def _tx(self) -> AsyncIterator["_RecordingDb"]:
yield self
def tx(self, timeout: object = None) -> AbstractAsyncContextManager["_RecordingDb"]:
return self._tx()
class _RecordingPrismaClient:
def __init__(self) -> None:

View file

@ -5,11 +5,11 @@ import logging
import re
from collections.abc import Callable
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from collections.abc import AsyncIterator, Callable
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
from typing import Final, cast
from unittest.mock import AsyncMock, MagicMock, call, patch
import httpx
@ -20,7 +20,7 @@ from redis.exceptions import DataError
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem
from litellm.proxy._types import DailyTagSpendTransaction, Litellm_EntityType, SpendUpdateQueueItem
from litellm.proxy.db.db_spend_update_writer import (
_TEAM_ADVISORY_LOCK_SQL,
_TEAM_MEMBER_SPEND_SQL,
@ -28,6 +28,8 @@ from litellm.proxy.db.db_spend_update_writer import (
_SpendTableName,
_spend_tables_left_to_send,
)
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import DailySpendUpdateQueue
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
build_window_spend_transaction,
@ -303,6 +305,13 @@ class _RecordingDb:
return self._execute_raw()
return len(args)
@asynccontextmanager
async def _tx(self) -> AsyncIterator["_RecordingDb"]:
yield self
def tx(self, timeout: timedelta | None = None) -> AbstractAsyncContextManager["_RecordingDb"]:
return self._tx()
class _RecordingPrisma:
def __init__(self, execute_raw: Callable[[], int] | None = None) -> None:
@ -3770,8 +3779,15 @@ async def test_commit_spend_updates_does_not_retry_non_deadlock_data_error(monke
@pytest.mark.asyncio
async def test_update_daily_spend_retries_deadlock(monkeypatch):
"""The daily-spend upsert path retries a deadlock on the bulk upsert and then drains successfully."""
mock_prisma_client = MagicMock()
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[_deadlock_error(), None])
outcomes = iter([_deadlock_error(), None])
def first_attempt_deadlocks():
outcome = next(outcomes)
if outcome is not None:
raise outcome
return 1
mock_prisma_client = _RecordingPrisma(execute_raw=first_attempt_deadlocks)
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
@ -3786,7 +3802,7 @@ async def test_update_daily_spend_retries_deadlock(monkeypatch):
entity_id_field="user_id",
)
assert mock_prisma_client.db.execute_raw.call_count == 2
assert len(mock_prisma_client.db.statements) == 2
assert daily_spend_transactions == {}
proxy_logging.failure_handler.assert_not_called()
@ -4247,3 +4263,241 @@ async def test_daily_transaction_attributes_caching_savings_only_with_an_injecti
assert transaction["cache_creation_input_tokens"] == 1111
assert transaction["prompt_caching_savings_spend"] != 0.0
assert transaction["gateway_injected_caching_savings_spend"] == 0.0
class _StallingDailySpendFakeDB(_DailySpendFakeDB):
"""Holds the daily upsert aimed at one table until it is cancelled, like a starved pool does.
The rollback of that transaction waits for ``rollback_release``: the query engine only
rolls back once the statement it is running has returned, which behind a lock takes
as long as the lock is held."""
def __init__(self, stalled_table: str) -> None:
super().__init__(failing_table=None)
self.stalled_table = stalled_table
self.stalled = asyncio.Event()
self.rollback_release = asyncio.Event()
self.rolled_back = asyncio.Event()
self.transaction_outcomes: list[str] = []
async def execute_raw(self, query: str, *args: object) -> int:
if self.stalled_table in query:
self.stalled.set()
await asyncio.Event().wait()
return await super().execute_raw(query, *args)
@asynccontextmanager
async def _tx(self) -> AsyncIterator["_StallingDailySpendFakeDB"]:
try:
yield self
except BaseException:
await self.rollback_release.wait()
self.transaction_outcomes.append("rollback")
self.rolled_back.set()
raise
self.transaction_outcomes.append("commit")
def _daily_entity_txn(entity_id_field: str) -> dict:
return {key: value for key, value in _daily_txn().items() if key != "user_id"} | {entity_id_field: "entity-1"}
_DAILY_SPEND_ENTITIES: Final = [
pytest.param("daily_spend_update_queue", "user", "user_id", "LiteLLM_DailyUserSpend", id="user"),
pytest.param("daily_team_spend_update_queue", "team", "team_id", "LiteLLM_DailyTeamSpend", id="team"),
pytest.param("daily_org_spend_update_queue", "org", "organization_id", "LiteLLM_DailyOrganizationSpend", id="org"),
pytest.param("daily_tag_spend_update_queue", "tag", "tag", "LiteLLM_DailyTagSpend", id="tag"),
pytest.param(
"daily_end_user_spend_update_queue", "end_user", "end_user_id", "LiteLLM_DailyEndUserSpend", id="end_user"
),
pytest.param("daily_agent_spend_update_queue", "agent", "agent_id", "LiteLLM_DailyAgentSpend", id="agent"),
]
_DAILY_SPEND_QUEUES: Final[dict[str, Callable[[DBSpendUpdateWriter], DailySpendUpdateQueue]]] = {
"daily_spend_update_queue": lambda writer: writer.daily_spend_update_queue,
"daily_team_spend_update_queue": lambda writer: writer.daily_team_spend_update_queue,
"daily_org_spend_update_queue": lambda writer: writer.daily_org_spend_update_queue,
"daily_tag_spend_update_queue": lambda writer: writer.daily_tag_spend_update_queue,
"daily_end_user_spend_update_queue": lambda writer: writer.daily_end_user_spend_update_queue,
"daily_agent_spend_update_queue": lambda writer: writer.daily_agent_spend_update_queue,
}
_DAILY_SPEND_COMMITS: Final = {
"user": DBSpendUpdateWriter.update_daily_user_spend,
"team": DBSpendUpdateWriter.update_daily_team_spend,
"org": DBSpendUpdateWriter.update_daily_org_spend,
"tag": DBSpendUpdateWriter.update_daily_tag_spend,
"end_user": DBSpendUpdateWriter.update_daily_end_user_spend,
"agent": DBSpendUpdateWriter.update_daily_agent_spend,
}
@pytest.mark.parametrize(("queue_name", "entity_type", "entity_id_field", "table"), _DAILY_SPEND_ENTITIES)
@pytest.mark.asyncio
async def test_daily_spend_batch_cancelled_mid_flight_is_rolled_back_requeued_and_written_once_by_the_next_flush(
queue_name: str, entity_type: str, entity_id_field: str, table: str
):
"""Shutdown cancels the scheduler tick while a drained batch waits on the database. The
batch has left the queue, so unless the cancellation puts it back, the final flush finds
nothing and the spend is gone (F2). The upsert runs in an interactive transaction so a
statement that did reach Postgres is rolled back with the cancel and the requeued rows
land exactly once."""
db_writer = DBSpendUpdateWriter()
queue = _DAILY_SPEND_QUEUES[queue_name](db_writer)
await queue.add_update({"key-a": _daily_entity_txn(entity_id_field)})
await queue.add_update({"key-a": _daily_entity_txn(entity_id_field)})
db = _StallingDailySpendFakeDB(stalled_table=table)
proxy_logging_obj = MagicMock()
proxy_logging_obj.failure_handler = AsyncMock()
def flush(prisma_db: _DailySpendFakeDB):
return db_writer._flush_daily_spend_queue(
queue=queue,
entity_type=entity_type,
commit=_DAILY_SPEND_COMMITS[entity_type],
n_retry_times=0,
prisma_client=_WindowSpendFakePrisma(prisma_db),
proxy_logging_obj=proxy_logging_obj,
)
tick = asyncio.ensure_future(flush(db))
await asyncio.wait_for(db.stalled.wait(), timeout=5)
tick.cancel()
finished, _ = await asyncio.wait({tick}, timeout=1)
assert finished == {tick}, "the cancelled tick must return before the rolled-back statement unwinds"
with pytest.raises(asyncio.CancelledError):
tick.result()
assert not queue.update_queue.empty(), "the cancelled batch must go back on the queue before the rollback lands"
assert db.transaction_outcomes == []
db.rollback_release.set()
await asyncio.wait_for(db.rolled_back.wait(), timeout=5)
assert db.transaction_outcomes == ["rollback"]
assert _daily_upserts(db, table) == []
final_db = _DailySpendFakeDB(failing_table=None)
await flush(final_db)
(upsert,) = _daily_upserts(final_db, table)
assert _row_values(upsert, entity_id_field) == ["entity-1"]
assert _row_values(upsert, "spend") == [pytest.approx(0.2)]
assert _row_values(upsert, "api_requests") == [2]
assert queue.update_queue.empty()
@pytest.mark.asyncio
async def test_cancelled_flush_of_an_empty_daily_queue_requeues_nothing():
"""A cancel that lands with nothing drained must not push an empty batch onto the queue."""
db_writer = DBSpendUpdateWriter()
db = _StallingDailySpendFakeDB(stalled_table="LiteLLM_DailyUserSpend")
class _CancellingQueue(type(db_writer.daily_spend_update_queue)):
async def flush_and_get_aggregated_daily_spend_update_transactions(self):
drained = await super().flush_and_get_aggregated_daily_spend_update_transactions()
asyncio.current_task().cancel()
await asyncio.sleep(0)
return drained
queue = _CancellingQueue()
with pytest.raises(asyncio.CancelledError):
await db_writer._flush_daily_spend_queue(
queue=queue,
entity_type="user",
commit=DBSpendUpdateWriter.update_daily_user_spend,
n_retry_times=0,
prisma_client=_WindowSpendFakePrisma(db),
proxy_logging_obj=MagicMock(),
)
assert queue.update_queue.empty()
class _AnnouncingDailySpendFakeDB(_DailySpendFakeDB):
"""Signals ``written`` the moment the daily upsert has been committed."""
def __init__(self) -> None:
super().__init__(failing_table=None)
self.written = asyncio.Event()
async def execute_raw(self, query: str, *args: object) -> int:
rows = await super().execute_raw(query, *args)
self.written.set()
return rows
@pytest.mark.asyncio
async def test_cancel_that_lands_after_the_daily_batch_committed_does_not_requeue_it():
"""The commit has returned but the tick has not resumed yet when the cancel arrives.
Putting the batch back now would write the same spend twice on the final flush."""
db_writer = DBSpendUpdateWriter()
queue = db_writer.daily_spend_update_queue
await queue.add_update({"key-a": _daily_txn()})
db = _AnnouncingDailySpendFakeDB()
tick = asyncio.ensure_future(
db_writer._flush_daily_spend_queue(
queue=queue,
entity_type="user",
commit=DBSpendUpdateWriter.update_daily_user_spend,
n_retry_times=0,
prisma_client=_WindowSpendFakePrisma(db),
proxy_logging_obj=MagicMock(),
)
)
await db.written.wait()
tick.cancel()
with pytest.raises(asyncio.CancelledError):
await tick
assert len(_daily_upserts(db, "LiteLLM_DailyUserSpend")) == 1
assert queue.update_queue.empty(), "a batch that already committed must not be requeued"
class _DrainedTagRedisBuffer:
"""Hands out one drained tag batch and records whatever is restored."""
def __init__(self, drained: dict[str, DailyTagSpendTransaction]) -> None:
self.drained = drained
self.restored: list[dict[str, DailyTagSpendTransaction]] = []
async def get_all_daily_tag_spend_update_transactions_from_redis_buffer(
self,
) -> dict[str, DailyTagSpendTransaction]:
return self.drained
async def restore_transactions_to_redis(
self, daily_tag_spend_update_transactions: dict[str, DailyTagSpendTransaction]
) -> None:
self.restored.append(daily_tag_spend_update_transactions)
@pytest.mark.asyncio
async def test_tag_batch_drained_from_redis_and_cancelled_mid_flight_is_restored_before_its_rollback_returns():
"""The Redis tag drain is destructive. A shutdown cancel used to leave the batch nowhere:
Redis no longer had it and the interactive transaction rolled the statement back."""
db_writer = DBSpendUpdateWriter()
drained = {"key-a": cast(DailyTagSpendTransaction, _daily_entity_txn("tag"))}
redis_buffer = _DrainedTagRedisBuffer(drained)
db_writer.redis_update_buffer = cast(RedisUpdateBuffer, redis_buffer)
db = _StallingDailySpendFakeDB(stalled_table="LiteLLM_DailyTagSpend")
tick = asyncio.ensure_future(
db_writer._drain_and_commit_daily_tag_spend_from_redis(
prisma_client=_WindowSpendFakePrisma(db),
n_retry_times=0,
proxy_logging_obj=MagicMock(),
)
)
await asyncio.wait_for(db.stalled.wait(), timeout=5)
tick.cancel()
finished, _ = await asyncio.wait({tick}, timeout=1)
assert finished == {tick}, "the cancelled drain must return before the rolled-back statement unwinds"
with pytest.raises(asyncio.CancelledError):
tick.result()
assert redis_buffer.restored == [drained], "the drained tag batch must be back in Redis before the rollback lands"
assert db.transaction_outcomes == []
db.rollback_release.set()
await asyncio.wait_for(db.rolled_back.wait(), timeout=5)
assert db.transaction_outcomes == ["rollback"]
assert _daily_upserts(db, "LiteLLM_DailyTagSpend") == []

View file

@ -723,11 +723,13 @@ class TestAutoRouterBenchmarks:
def test_savings_compare_only_the_current_estimated_cohort(self, estimated_turns: int) -> None:
from litellm.proxy.management_endpoints.auto_router_endpoints import _benchmark_totals
row: Final = self.ROW.model_copy(update={
"savings_estimated_turns": estimated_turns,
"savings_estimated_actual_spend": 2.0 if estimated_turns else 0.0,
"savings_estimated_saved_spend": -0.5 if estimated_turns else 0.0,
})
row: Final = self.ROW.model_copy(
update={
"savings_estimated_turns": estimated_turns,
"savings_estimated_actual_spend": 2.0 if estimated_turns else 0.0,
"savings_estimated_saved_spend": -0.5 if estimated_turns else 0.0,
}
)
totals: Final = _benchmark_totals(row)
assert totals.spend == 10.0
assert totals.savings_estimated_turns == estimated_turns
@ -755,10 +757,16 @@ class TestAutoRouterBenchmarks:
_summed_agg_row,
)
other = self.ROW.model_copy(update={
"router_name": "auto-2", "sessions": 1, "turns": 10, "spend": 0.0,
"savings_estimated_turns": 10, "savings_estimated_actual_spend": 0.0,
})
other = self.ROW.model_copy(
update={
"router_name": "auto-2",
"sessions": 1,
"turns": 10,
"spend": 0.0,
"savings_estimated_turns": 10,
"savings_estimated_actual_spend": 0.0,
}
)
summed = _summed_agg_row([self.ROW, other])
totals = _benchmark_totals(summed)
assert summed.sessions == 5
@ -1091,18 +1099,27 @@ class TestAutoRouterSession:
return lookups
@pytest.mark.asyncio
@pytest.mark.parametrize("turns, estimated", [(3, True), (10, True), (10, False)], ids=["full", "partial", "legacy"])
@pytest.mark.parametrize(
"turns, estimated", [(3, True), (10, True), (10, False)], ids=["full", "partial", "legacy"]
)
async def test_a_key_reads_its_own_session_with_the_baseline_its_turns_were_priced_against(
self, monkeypatch: pytest.MonkeyPatch, turns: int, estimated: bool,
self,
monkeypatch: pytest.MonkeyPatch,
turns: int,
estimated: bool,
) -> None:
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session
caller = UserAPIKeyAuth(api_key="sk-caller")
row: Final = {key: value for key, value in self.ROW.items() if estimated or not key.startswith("savings_estimated_")}
row: Final = {
key: value for key, value in self.ROW.items() if estimated or not key.startswith("savings_estimated_")
}
spend: Final = 0.14 if turns == 3 else 10.0
if estimated and turns != 3:
row["savings_estimated_saved_spend"] = -0.04
self._rig(monkeypatch, [{**row, "api_key": caller.api_key, "session_id": "sess-1", "turns": turns, "spend": spend}])
self._rig(
monkeypatch, [{**row, "api_key": caller.api_key, "session_id": "sess-1", "turns": turns, "spend": spend}]
)
response = await get_auto_router_session(user_api_key_dict=caller, session_id="sess-1")
assert response.model_dump() == {
"session_id": "sess-1",
@ -1159,10 +1176,18 @@ class TestAutoRouterSession:
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session
priced = {"anthropic/claude-opus-5": 2, "anthropic/claude-sonnet-5": 1}
self._rig(monkeypatch, [{
**self.ROW, "api_key": ADMIN.api_key, "session_id": "s",
"baseline_models": {"old-baseline": 100}, "savings_estimated_baseline_models": priced,
}])
self._rig(
monkeypatch,
[
{
**self.ROW,
"api_key": ADMIN.api_key,
"session_id": "s",
"baseline_models": {"old-baseline": 100},
"savings_estimated_baseline_models": priced,
}
],
)
response = await get_auto_router_session(user_api_key_dict=ADMIN, session_id="s")
assert response.baseline_model == "anthropic/claude-opus-5"
assert response.baseline_models == priced
@ -3562,3 +3587,97 @@ async def test_start_shadow_eval_seeds_a_zero_funnel_row_per_leg(monkeypatch: py
if "group_id" in call.kwargs.get("where", {})
]
assert group_reads == []
@pytest.mark.asyncio
async def test_availability_counts_db_and_yaml_without_disclosing_router_names(monkeypatch):
from litellm.models.model import LiteLLM_ProxyModelTable
from litellm.proxy.management_helpers.auto_router_availability import build_auto_router_catalog
from litellm.types.management_endpoints.auto_router_endpoints import AutoRouterAvailabilityRequest
row = LiteLLM_ProxyModelTable(
model_id="db-router",
model_name="private-team-router",
created_by="someone-else",
litellm_params={
"model": "auto_router/complexity_router",
"complexity_router_config": {"classifier_type": "heuristic_v2"},
},
)
yaml_row = {
"model_name": "private-yaml-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"classifier_type": "capability"},
},
}
find_many = AsyncMock(side_effect=AssertionError("Availability must not query the model table"))
monkeypatch.setattr(
proxy_server,
"prisma_client",
SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))),
)
monkeypatch.setattr(proxy_server.proxy_config, "auto_router_db_catalog", build_auto_router_catalog((row,)))
monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(config_deployments=lambda: (yaml_row,)))
monkeypatch.setattr(proxy_server, "_license_check", SimpleNamespace(auto_router_capability_limit=lambda: 1))
monkeypatch.setattr(proxy_server, "heuristic_v1_tuning_baselines", {})
result = await auto_router_endpoints.get_auto_router_availability(AutoRouterAvailabilityRequest(), ADMIN)
assert {slot.key: slot.remaining for slot in result.allowances} == {
"heuristic_v2": 0,
"capability": 0,
"llm_v2": 1,
"tier_or_classifier_prompt": 1,
"heuristic_tuning": 1,
}
assert "private" not in result.model_dump_json()
edit = await auto_router_endpoints.get_auto_router_availability(
AutoRouterAvailabilityRequest(
saved_model_id="db-router", complexity_router_config={"classifier_type": "heuristic_v2"}
),
ADMIN,
)
assert edit.allowances[0].used_by_this_router
assert edit.allowances[0].remaining == 1
assert edit.error is None
find_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_availability_denies_another_teams_edit_exemption(monkeypatch):
from litellm.models.model import LiteLLM_ProxyModelTable
from litellm.proxy.management_helpers.auto_router_availability import build_auto_router_catalog
from litellm.types.management_endpoints.auto_router_endpoints import AutoRouterAvailabilityRequest
row = LiteLLM_ProxyModelTable(
model_id="other-router",
model_name="other",
created_by="other",
model_info={"team_id": "other-team"},
litellm_params={"model": "auto_router/complexity_router"},
)
find_many = AsyncMock(side_effect=AssertionError("Availability must not query the model table"))
monkeypatch.setattr(
proxy_server,
"prisma_client",
SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))),
)
monkeypatch.setattr(proxy_server.proxy_config, "auto_router_db_catalog", build_auto_router_catalog((row,)))
monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(config_deployments=lambda: ()))
monkeypatch.setattr(auto_router_endpoints, "_authorize_router_dry_run", AsyncMock(return_value=None))
with pytest.raises(HTTPException) as error:
await auto_router_endpoints.get_auto_router_availability(
AutoRouterAvailabilityRequest(team_id="own-team", saved_model_id="other-router"),
UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="owner"),
)
assert error.value.status_code == 403
@pytest.mark.asyncio
async def test_availability_waits_for_the_first_complete_catalog(monkeypatch):
from litellm.types.management_endpoints.auto_router_endpoints import AutoRouterAvailabilityRequest
monkeypatch.setattr(proxy_server.proxy_config, "auto_router_db_catalog", None)
monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(config_deployments=lambda: ()))
with pytest.raises(HTTPException) as error:
await auto_router_endpoints.get_auto_router_availability(AutoRouterAvailabilityRequest(), ADMIN)
assert error.value.status_code == 503

View file

@ -3,6 +3,7 @@ import asyncio
import contextlib
import json
from collections.abc import Iterator, Mapping
from types import SimpleNamespace
from typing import Dict, Final, Optional
from unittest.mock import AsyncMock, MagicMock, patch
@ -1222,6 +1223,117 @@ class TestDeleteModelClearsRouterRegistry:
assert mock_router.complexity_routers.get("shared-name") is config_router
@pytest.fixture
def deleted_auto_router_catalog(monkeypatch):
from litellm.proxy import proxy_server
from litellm.proxy.management_helpers.auto_router_availability import build_auto_router_catalog
rows = tuple(
LiteLLM_ProxyModelTable(
model_id=model_id,
model_name=f"model_name_{team_id}_{model_id}",
litellm_params={
"model": "auto_router/complexity_router",
"complexity_router_config": {"classifier_type": classifier},
},
model_info={"id": model_id, "team_id": team_id},
created_by="admin",
updated_by="admin",
blocked=True,
)
for model_id, team_id, classifier in (
("deleted-router", "deleted-team", "heuristic_v2"),
("surviving-router", "surviving-team", "llm_v2"),
)
)
config = proxy_server.ProxyConfig()
config.auto_router_db_catalog = build_auto_router_catalog(rows)
monkeypatch.setattr(proxy_server, "proxy_config", config)
monkeypatch.setattr(proxy_server, "MODEL_RECONCILE_LOCK", asyncio.Lock())
monkeypatch.setattr(proxy_server, "llm_router", Router(model_list=[]))
monkeypatch.setattr(proxy_server, "_license_check", SimpleNamespace(auto_router_capability_limit=lambda: 1))
monkeypatch.setattr(proxy_server, "heuristic_v1_tuning_baselines", {})
return config, rows
class TestDeletedAutoRouterAvailability:
@pytest.mark.asyncio
@pytest.mark.parametrize("delete_succeeds,has_router", ((True, True), (True, False), (False, True)))
async def test_single_delete_releases_allowance_only_after_success(
self, monkeypatch, deleted_auto_router_catalog, delete_succeeds, has_router
):
from litellm.proxy import proxy_server
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_availability
from litellm.proxy.management_endpoints.model_management_endpoints import ModelInfoDelete, delete_model
from litellm.types.management_endpoints.auto_router_endpoints import AutoRouterAvailabilityRequest
config, rows = deleted_auto_router_catalog
original = config.auto_router_db_catalog
row = rows[0].model_copy(update={"model_info": {"id": rows[0].model_id}})
table = SimpleNamespace(
find_unique=AsyncMock(return_value=row),
delete=AsyncMock(return_value=row, side_effect=None if delete_succeeds else RuntimeError("delete failed")),
)
prisma = SimpleNamespace(
db=SimpleNamespace(litellm_proxymodeltable=table, query_raw=AsyncMock(return_value=[]))
)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
request = AutoRouterAvailabilityRequest(complexity_router_config={"classifier_type": "heuristic_v2"})
before = await get_auto_router_availability(request, admin)
assert before.error is not None
if not has_router:
monkeypatch.setattr(proxy_server, "llm_router", None)
if not delete_succeeds:
with pytest.raises(ProxyException, match="delete failed"):
await delete_model(ModelInfoDelete(id=row.model_id), admin)
assert config.auto_router_db_catalog == original
return
await delete_model(ModelInfoDelete(id=row.model_id), admin)
monkeypatch.setattr(proxy_server, "llm_router", Router(model_list=[]))
after = await get_auto_router_availability(request, admin)
assert after.error is None
assert {slot.key: slot.remaining for slot in after.allowances} == {
"heuristic_v2": 1,
"capability": 1,
"llm_v2": 0,
"tier_or_classifier_prompt": 1,
"heuristic_tuning": 1,
}
@pytest.mark.asyncio
@pytest.mark.parametrize("has_router", (True, False))
async def test_team_delete_releases_only_its_routers_allowance(
self, monkeypatch, deleted_auto_router_catalog, has_router
):
from litellm.proxy import proxy_server
from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_availability
from litellm.types.management_endpoints.auto_router_endpoints import AutoRouterAvailabilityRequest
_, rows = deleted_auto_router_catalog
prisma = _TxPrismaClient(rows)
deleted = await delete_team_models(
team_ids=["deleted-team"], prisma_client=prisma, llm_router=proxy_server.llm_router if has_router else None
)
assert deleted == ["deleted-router"]
after = await get_auto_router_availability(
AutoRouterAvailabilityRequest(complexity_router_config={"classifier_type": "heuristic_v2"}),
UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert after.error is None
assert {slot.key: slot.remaining for slot in after.allowances} == {
"heuristic_v2": 1,
"capability": 1,
"llm_v2": 0,
"tier_or_classifier_prompt": 1,
"heuristic_tuning": 1,
}
class TestUpdateModel:
"""
Tests for the update_model (POST /model/update) handler.
@ -5274,7 +5386,7 @@ class TestDeleteEvictionsHoldTheReconcileLock:
"""
@staticmethod
async def _assert_evicts_under_lock(monkeypatch, call_endpoint, model_id: str) -> None:
async def _assert_evicts_under_lock(monkeypatch, call_endpoint, model_id: str, config) -> None:
"""Run ``call_endpoint`` with the lock already held and assert it blocks.
Holding MODEL_RECONCILE_LOCK stands in for a reconcile that is mid-flight. If
@ -5291,6 +5403,7 @@ class TestDeleteEvictionsHoldTheReconcileLock:
"""
lock = asyncio.Lock()
monkeypatch.setattr("litellm.proxy.proxy_server.MODEL_RECONCILE_LOCK", lock)
stale_catalog = config.auto_router_db_catalog
async with lock:
task = asyncio.create_task(call_endpoint())
@ -5301,16 +5414,19 @@ class TestDeleteEvictionsHoldTheReconcileLock:
f"deleting {model_id} did not wait for MODEL_RECONCILE_LOCK -- an "
f"in-flight reconcile can resurrect the deployment it just evicted"
)
config.auto_router_db_catalog = stale_catalog
await asyncio.wait_for(task, timeout=5)
assert tuple(row.model_id for row in config.auto_router_db_catalog) == ("surviving-router",)
@pytest.mark.asyncio
async def test_delete_model_waits_for_an_in_flight_reconcile(self, monkeypatch):
async def test_delete_model_waits_for_an_in_flight_reconcile(self, monkeypatch, deleted_auto_router_catalog):
from litellm.proxy.management_endpoints.model_management_endpoints import (
ModelInfoDelete,
delete_model,
)
model_id = "m-doomed"
config, rows = deleted_auto_router_catalog
model_id = rows[0].model_id
row = MagicMock()
row.model_dump.return_value = {
"model_name": "gpt-4o",
@ -5347,16 +5463,17 @@ class TestDeleteEvictionsHoldTheReconcileLock:
),
)
await self._assert_evicts_under_lock(monkeypatch, call, model_id)
await self._assert_evicts_under_lock(monkeypatch, call, model_id, config)
router.delete_deployment.assert_called_once_with(id=model_id)
@pytest.mark.asyncio
async def test_delete_team_models_waits_for_an_in_flight_reconcile(self, monkeypatch):
async def test_delete_team_models_waits_for_an_in_flight_reconcile(self, monkeypatch, deleted_auto_router_catalog):
from litellm.proxy.management_endpoints.model_management_endpoints import (
delete_team_models,
)
model_id = "m-team-doomed"
config, rows = deleted_auto_router_catalog
model_id = rows[0].model_id
router = MagicMock()
router.delete_deployment = MagicMock(return_value=True)
@ -5392,7 +5509,7 @@ class TestDeleteEvictionsHoldTheReconcileLock:
team_ids=["team-1"], prisma_client=prisma, llm_router=router
)
await self._assert_evicts_under_lock(monkeypatch, call, model_id)
await self._assert_evicts_under_lock(monkeypatch, call, model_id, config)
router.delete_deployment.assert_called_once_with(id=model_id)
@ -6092,7 +6209,8 @@ class TestStrategyRouterWriteValidation:
_TUNED_A = {"classifier_type": "heuristic", "tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"}}
_TUNED_A_EDITED = {**_TUNED_A, "dimension_weights": {"codePresence": 0.9}}
_TUNED_B = {"classifier_type": "heuristic", "tiers": {"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4.1"}}
_TUNED_B_EDITED = {**_TUNED_B, "tiers": {"SIMPLE": "gpt-4o", "MEDIUM": "gpt-4.1"}}
_TUNED_B_EDITED = {**_TUNED_B, "code_keywords": ["internal-api"]}
_MODELS_ONLY_B = {**_TUNED_B, "tiers": {"SIMPLE": "fast-model", "MEDIUM": "capable-model"}}
@staticmethod
def _db_router_row(model_id: str, config: Mapping[str, object]) -> dict[str, object]:
@ -6110,7 +6228,9 @@ class TestStrategyRouterWriteValidation:
(1, ["a", "b"], {"a": "_TUNED_A", "b": "_TUNED_B"}, "a", "_TUNED_A_EDITED", "allowed"),
(1, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "a", "_TUNED_A_EDITED", "allowed"),
(1, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "b", "_TUNED_B_EDITED", "refused"),
(1, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "c", "_TUNED_B", "refused"),
(1, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "c", "_TUNED_B_EDITED", "refused"),
(1, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "c", "_TUNED_B", "allowed"),
(1, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "b", "_MODELS_ONLY_B", "allowed"),
(1, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "b", "_TUNED_B", "allowed"),
(None, ["a", "b"], {"a": "_TUNED_A_EDITED", "b": "_TUNED_B"}, "b", "_TUNED_B_EDITED", "allowed"),
(1, [], {}, "c", "_TUNED_B", "allowed"),
@ -6138,6 +6258,7 @@ class TestStrategyRouterWriteValidation:
"_TUNED_A_EDITED": self._TUNED_A_EDITED,
"_TUNED_B": self._TUNED_B,
"_TUNED_B_EDITED": self._TUNED_B_EDITED,
"_MODELS_ONLY_B": self._MODELS_ONLY_B,
}
baselines = snapshot_tuning_baselines(
[self._db_router_row(row_id, configs["_TUNED_A" if row_id == "a" else "_TUNED_B"]) for row_id in baseline_rows]
@ -6172,7 +6293,7 @@ class TestStrategyRouterWriteValidation:
async with _auto_router_capability_slot(fake, effective_params=effective_params, model_id=candidate_id):
pass
assert exc_info.value.status_code == 403
assert "changed heuristic scorer settings or tier models" in str(exc_info.value.detail)
assert "changed heuristic scoring rules" in str(exc_info.value.detail)
assert "'auto_router' feature lifts the limit" in str(exc_info.value.detail)
return
async with _auto_router_capability_slot(fake, effective_params=effective_params, model_id=candidate_id) as table:
@ -6222,13 +6343,13 @@ class TestStrategyRouterWriteValidation:
model_params=Deployment(
model_name="second-tuned",
litellm_params=LiteLLM_Params(
model="auto_router/complexity_router", complexity_router_config=self._TUNED_B
model="auto_router/complexity_router", complexity_router_config=self._TUNED_B_EDITED
),
),
user_api_key_dict=admin,
)
assert exc_info.value.code == "403"
assert "changed heuristic scorer settings or tier models" in str(exc_info.value.message)
assert "changed heuristic scoring rules" in str(exc_info.value.message)
fake.tx_obj.litellm_proxymodeltable.create.assert_not_awaited()
fake.litellm_proxymodeltable.create.assert_not_awaited()

View file

@ -0,0 +1,196 @@
from collections.abc import Mapping
from typing import Final
from types import SimpleNamespace
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
import pytest
from litellm.proxy.management_helpers.auto_router_availability import (
auto_router_availability,
build_auto_router_catalog,
)
from litellm.router_utils.auto_router_tuning_baseline import snapshot_tuning_baselines
def deployment(
model_id: str,
classifier: str,
*,
model: str = "solver",
tuned: bool = False,
config: Mapping[str, object] | None = None,
) -> Mapping[str, object]:
return {
"model_name": model_id,
"model_info": {"id": model_id, "db_model": True},
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {
"classifier_type": classifier,
"tiers": {"SIMPLE": [model]},
**({"code_keywords": ["internal-api"]} if tuned else {}),
**(config or {}),
},
},
}
@pytest.mark.parametrize("classifier", ("heuristic_v2", "capability", "llm_v2"))
def test_occupied_allowance_blocks_new_router_but_not_owner(classifier: str) -> None:
existing: Final = deployment("existing", classifier)
candidate: Final = deployment("new", classifier)
new: Final = auto_router_availability(others=(existing,), existing=None, candidate=candidate, baselines={}, limit=1)
edit: Final = auto_router_availability(others=(), existing=existing, candidate=existing, baselines={}, limit=1)
new_slot: Final = next(slot for slot in new.allowances if slot.key == classifier)
edit_slot: Final = next(slot for slot in edit.allowances if slot.key == classifier)
assert (new_slot.remaining, new_slot.used_by_this_router, new.error is not None) == (0, False, True)
assert (edit_slot.remaining, edit_slot.used_by_this_router, edit.error) == (1, True, None)
def test_edit_does_not_exempt_another_classifier_allowance() -> None:
existing: Final = deployment("existing", "capability")
result: Final = auto_router_availability(
others=(deployment("other", "llm_v2"),),
existing=existing,
candidate=deployment("existing", "llm_v2"),
baselines={},
limit=1,
)
assert result.error is not None
assert next(slot for slot in result.allowances if slot.key == "llm_v2").remaining == 0
def test_model_selection_does_not_claim_occupied_scoring_allowance() -> None:
original: Final = deployment("legacy", "heuristic")
changed: Final = deployment("other", "heuristic", tuned=True)
baselines: Final = snapshot_tuning_baselines((original,))
unchanged: Final = auto_router_availability(
others=(changed,),
existing=original,
candidate=original,
baselines=baselines,
limit=1,
)
edited: Final = auto_router_availability(
others=(changed,),
existing=original,
candidate=deployment("legacy", "heuristic", model="new"),
baselines=baselines,
limit=1,
)
assert unchanged.error is None
assert next(slot for slot in unchanged.allowances if slot.key == "heuristic_tuning").remaining == 0
assert edited.error is None
tuned: Final = auto_router_availability(
others=(changed,),
existing=original,
candidate=deployment("legacy", "heuristic", tuned=True),
baselines=baselines,
limit=1,
)
assert tuned.error is not None
assert "weights, thresholds, keywords, and custom dimensions" in tuned.error
def test_missing_baselines_are_reported_as_unknown() -> None:
result: Final = auto_router_availability(
others=(),
existing=None,
candidate=deployment("new", "heuristic"),
baselines=None,
limit=1,
)
slot: Final = next(slot for slot in result.allowances if slot.key == "heuristic_tuning")
assert (slot.available, slot.remaining, slot.limit) == (False, None, 1)
def test_unlimited_entitlement_does_not_report_exhausted_allowances() -> None:
result: Final = auto_router_availability(
others=(deployment("other", "heuristic_v2"),),
existing=None,
candidate=deployment("new", "heuristic_v2"),
baselines=None,
limit=None,
)
assert all(slot.available and slot.limit is None and slot.remaining is None for slot in result.allowances)
assert result.error is None
@pytest.mark.parametrize(
"customization",
(
{"tier_definitions": [{"name": "SIMPLE"}, {"name": "AUDIT", "description": "Review risks"}]},
{"classification_prompt": "Use the simplest sufficient tier"},
{"classification_examples": "Review this code -> COMPLEX"},
{"classifier_llm_config": {"model": "judge", "system_prompt": "Route by urgency"}},
),
)
def test_customization_owner_can_edit_models_and_restoring_defaults_clears_the_gate(
customization: Mapping[str, object],
) -> None:
owner: Final = deployment("owner", "llm", config=customization)
blocked: Final = auto_router_availability(
others=(owner,), existing=None, candidate=deployment("new", "llm", config=customization), baselines={}, limit=1
)
assert blocked.error is not None
assert "Custom tiers or classifier instructions" in blocked.error
edited: Final = auto_router_availability(
others=(),
existing=owner,
candidate=deployment("owner", "llm", model="new", config=customization),
baselines={},
limit=1,
)
assert edited.error is None
assert next(slot for slot in edited.allowances if slot.key == "tier_or_classifier_prompt").used_by_this_router
restored: Final = auto_router_availability(
others=(owner,), existing=None, candidate=deployment("new", "llm"), baselines={}, limit=1
)
assert restored.error is None
assert next(slot for slot in restored.allowances if slot.key == "tier_or_classifier_prompt").remaining == 0
def test_restoring_tiers_does_not_exempt_a_retained_custom_prompt() -> None:
prompt: Final = {"classification_prompt": "Use the simplest sufficient tier"}
result: Final = auto_router_availability(
others=(deployment("owner", "llm", config=prompt),),
existing=None,
candidate=deployment("new", "llm", config=prompt),
baselines={},
limit=1,
)
assert result.error is not None
assert "Custom tiers or classifier instructions" in result.error
@pytest.mark.parametrize("blocked", (False, True))
def test_catalog_keeps_unloaded_routers_and_ownership_without_provider_credentials(blocked: bool, monkeypatch) -> None:
monkeypatch.setenv("LITELLM_SALT_KEY", "catalog-test-key")
source: Final = SimpleNamespace(
model_id="saved",
created_by="owner",
model_info={"team_id": "team"},
blocked=blocked,
litellm_params={
"model": encrypt_value_helper("auto_router/complexity_router"),
"api_key": "private-key",
"complexity_router_config": {"classifier_type": "heuristic_v2"},
},
)
provider: Final = SimpleNamespace(model_id="provider", litellm_params={"model": "openai/model"})
catalog: Final = build_auto_router_catalog((source, provider))
assert catalog is not None and len(catalog) == 1
assert (catalog[0].model_id, catalog[0].team_id, catalog[0].created_by) == ("saved", "team", "owner")
assert catalog[0].deployment == {
"model_info": {"id": "saved", "db_model": True},
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"classifier_type": "heuristic_v2"},
},
}
def test_catalog_distinguishes_missing_data_from_an_empty_model_table() -> None:
assert build_auto_router_catalog(()) == ()
assert build_auto_router_catalog((SimpleNamespace(model_id="incomplete"),)) is None

View file

@ -920,7 +920,7 @@ def test_proxy_startup_event_warns_for_global_budget_without_database():
@pytest.mark.asyncio
async def test_tuning_baseline_v2_is_created_alongside_the_legacy_row():
async def test_tuning_baseline_v3_is_created_alongside_the_legacy_row():
from litellm.router_utils.auto_router_tuning_baseline import DEFAULT_TUNING_FINGERPRINT
prisma_client = MagicMock()
@ -935,11 +935,61 @@ async def test_tuning_baseline_v2_is_created_alongside_the_legacy_row():
assert result == {'yaml:["a",[]]': DEFAULT_TUNING_FINGERPRINT}
assert prisma_client.db.litellm_config.create.await_args.kwargs["data"] == {
"param_name": "auto_router_tuning_baseline_v2",
"param_name": "auto_router_tuning_baseline_v3",
"param_value": json.dumps(dict(result)),
}
@pytest.mark.asyncio
async def test_scorer_baseline_upgrade_preserves_existing_routers_and_is_not_refreshed_on_restart():
from litellm.router_utils.auto_router_tuning_baseline import mutable_tuned_identities, snapshot_tuning_baselines
deployments = [
{
"model_name": name,
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"tiers": {"SIMPLE": name}, "code_keywords": [name]},
},
}
for name in ("a", "b")
]
prisma_client = MagicMock()
prisma_client.db.litellm_config.find_unique = AsyncMock(
side_effect=lambda where: (
MagicMock(param_value='{"legacy-router":"old-combined-hash"}')
if where["param_name"] == "auto_router_tuning_baseline_v2"
else None
)
)
prisma_client.db.litellm_config.create = AsyncMock()
baseline = await ProxyStartupEvent._load_heuristic_v1_tuning_baselines(prisma_client, deployments)
assert baseline == snapshot_tuning_baselines(deployments)
assert mutable_tuned_identities(deployments, baseline) == frozenset()
prisma_client.db.litellm_config.create.assert_awaited_once_with(
data={"param_name": "auto_router_tuning_baseline_v3", "param_value": json.dumps(dict(baseline))}
)
prisma_client.db.litellm_config.find_unique.side_effect = None
prisma_client.db.litellm_config.find_unique.return_value = MagicMock(param_value=json.dumps(dict(baseline)))
changed = [
{
"model_name": "a",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {"tiers": {"SIMPLE": "different-model"}, "code_keywords": ["new-rule"]},
},
}
]
reloaded = await ProxyStartupEvent._load_heuristic_v1_tuning_baselines(prisma_client, changed)
assert reloaded == baseline
assert mutable_tuned_identities(changed, reloaded) == frozenset({'yaml:["a",[]]'})
prisma_client.db.litellm_config.create.assert_awaited_once()
@pytest.mark.asyncio
async def test_tuning_baseline_waits_for_a_complete_db_model_census(monkeypatch):
prisma_client = MagicMock()

View file

@ -875,9 +875,7 @@ class _ConfigTable:
await asyncio.sleep(0)
return _ConfigRow(param_value=value) if value is not None else None
async def upsert(
self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]
) -> _ConfigRow:
async def upsert(self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> _ConfigRow:
param_name: Final = where["param_name"]
value: Final = _CONFIG_VALUE.validate_json(data["update"]["param_value"])
self.rows[param_name] = value
@ -926,7 +924,9 @@ class _ConfigPrisma:
self.db.litellm_config.upserted_param_names.append(param_name)
def _db_backed_proxy_config(monkeypatch, rows: Mapping[str, Mapping[str, JsonValue]]) -> tuple[ProxyConfig, _ConfigTable]:
def _db_backed_proxy_config(
monkeypatch, rows: Mapping[str, Mapping[str, JsonValue]]
) -> tuple[ProxyConfig, _ConfigTable]:
table: Final = _ConfigTable(rows)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", _ConfigPrisma(db=_ConfigDb(litellm_config=table)))
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
@ -4750,9 +4750,7 @@ def test_validate_deployment_access_windows_rejects_malformed_time():
"model_name": "gpt-4o-shared",
"litellm_params": {"model": "gpt-4o"},
"model_info": {
"access_windows": [
{"start": "25:00", "end": "06:00", "timezone": "America/New_York", "team_ids": ["t"]}
]
"access_windows": [{"start": "25:00", "end": "06:00", "timezone": "America/New_York", "team_ids": ["t"]}]
},
}
@ -4767,9 +4765,7 @@ def test_validate_deployment_access_windows_rejects_unknown_timezone():
"model_name": "gpt-4o-shared",
"litellm_params": {"model": "gpt-4o"},
"model_info": {
"access_windows": [
{"start": "22:00", "end": "06:00", "timezone": "Mars/Olympus", "team_ids": ["t"]}
]
"access_windows": [{"start": "22:00", "end": "06:00", "timezone": "Mars/Olympus", "team_ids": ["t"]}]
},
}
@ -4799,3 +4795,28 @@ def test_validate_deployment_access_windows_accepts_valid_and_absent():
)
is None
)
@pytest.mark.asyncio
async def test_model_refresh_updates_availability_catalog_and_retains_it_on_db_failure():
pc = ProxyConfig()
row = SimpleNamespace(
model_id="gated",
created_by="owner",
model_info={},
litellm_params={
"model": "auto_router/complexity_router",
"complexity_router_config": {"classifier_type": "heuristic_v2"},
},
)
find_many = AsyncMock(side_effect=[[row], RuntimeError("database unavailable"), []])
client = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many)))
assert pc.auto_router_db_catalog is None
assert await pc._get_models_from_db(client) == [row]
loaded = pc.auto_router_db_catalog
assert loaded is not None and loaded[0].model_id == "gated"
assert await pc._get_models_from_db(client) is None
assert pc.auto_router_db_catalog == loaded
assert await pc._get_models_from_db(client) == []
assert pc.auto_router_db_catalog == ()
assert find_many.await_count == 3

View file

@ -1218,13 +1218,13 @@ def test_get_autorouter_presets_local_mode_serves_bundled_catalog(
assert "anthropic_family" in payload
assert payload["1m_context"]["complexity_router_config"]["classifier_type"] == "heuristic_v2"
assert payload["1m_context"]["complexity_router_config"]["tiers"] == {
"SIMPLE": ["gpt-5.6-luna"],
"SIMPLE": ["gpt-6-luna"],
"MEDIUM": ["gpt-5.6-terra"],
"COMPLEX": ["gpt-5.6-sol"],
"REASONING": ["claude-opus-5"],
"COMPLEX": ["gpt-6-sol"],
"REASONING": ["claude-opus-5-5"],
}
assert payload["1m_context"]["complexity_router_config"]["tier_model_configs"] == {
"REASONING": [{"model_name": "claude-opus-5", "litellm_params": {"reasoning_effort": "high"}}]
"REASONING": [{"model_name": "claude-opus-5-5", "litellm_params": {"reasoning_effort": "high"}}]
}
for preset in payload.values():
assert isinstance(preset["label"], str)

View file

@ -239,7 +239,6 @@ async def test_arerank_error_is_mapped_to_litellm_exception(respx_mock: respx.Mo
@pytest.mark.asyncio
@pytest.mark.timeout(300)
async def test_arerank_declared_authenticating_provider_skips_resolution(monkeypatch):
"""Regression for the event-loop hazard in arerank's provider pre-resolution:
get_llm_provider runs the blocking OAuth device flow for github_copilot/chatgpt,
@ -257,6 +256,9 @@ async def test_arerank_declared_authenticating_provider_skips_resolution(monkeyp
raise BaseLLMException(status_code=401, message='{"error":"bad key"}')
monkeypatch.setattr(litellm, "get_llm_provider", record_resolution)
monkeypatch.setattr(
"litellm.litellm_core_utils.llm_response_utils.get_api_base.get_llm_provider", record_resolution
)
monkeypatch.setattr("litellm.rerank_api.main.rerank", rerank_raises_provider_error)
with pytest.raises(litellm.AuthenticationError) as exc_info:

View file

@ -33,21 +33,6 @@ _HISTORICAL_FINGERPRINTS: Final = (
{"custom_dimensions": [{"name": "sqlDdl", "weight": 0.4, "patterns": [r"\bCREATE\s{1,4}TABLE\b"]}]},
"814ce0017fc7f60a160b262f658d910e9bdf784e6139a4ba4f1e2657aa203950",
),
(
{
"tiers": _TIERS,
"dimension_weights": {"codePresence": 0.3},
"custom_dimensions": [
{
"name": "internalFrameworks",
"weight": 0.2,
"keywords": ["orbitmesh", "fluxgate"],
"patterns": [r"\bALTER\s{1,4}TABLE\b"],
}
],
},
"38970dc9224e265ab38c89674563d8d0537822591f9239b45251db6f5ca6cc39",
),
)
@ -78,11 +63,9 @@ class TestTuningFingerprint:
{"tiers": {"SIMPLE": "x"}}
)
@pytest.mark.parametrize("field", sorted(set(HEURISTIC_V1_TUNING_FIELDS) - {"tier_model_configs"}))
@pytest.mark.parametrize("field", HEURISTIC_V1_TUNING_FIELDS)
def test_every_tuning_field_changes_the_fingerprint(self, field: str) -> None:
samples: dict[str, object] = {
"tiers": _ALT_TIERS,
"classifier_type": "heuristic_first",
"tier_boundaries": {"simple_medium": 0.2, "medium_complex": 0.4, "complex_reasoning": 0.7},
"reasoning_override_min_score": 0.05,
"token_thresholds": {"simple": 20, "complex": 500},
@ -97,9 +80,6 @@ class TestTuningFingerprint:
"keyword_tier_rules": [{"keywords": ["urgent"], "tier": "COMPLEX"}],
}
config: dict[str, object] = {field: samples[field]}
if field == "classifier_type":
config["heuristic_first_max_tier"] = "MEDIUM"
config["classifier_llm_config"] = {"model": "judge"}
assert tuning_fingerprint(config) != DEFAULT_TUNING_FINGERPRINT
def test_explicit_empty_tier_model_configs_follow_omission(self) -> None:
@ -126,12 +106,33 @@ class TestTuningFingerprint:
!= historical
)
def test_tier_model_overrides_change_the_fingerprint(self) -> None:
def test_tier_model_overrides_do_not_change_the_fingerprint(self) -> None:
plain = tuning_fingerprint({"tiers": {"SIMPLE": "x"}})
with_override = tuning_fingerprint(
{"tiers": {"SIMPLE": {"model_name": "x", "litellm_params": {"temperature": 0.1}}}}
)
assert plain != with_override
assert plain == with_override == DEFAULT_TUNING_FINGERPRINT
@pytest.mark.parametrize("classifier_type", ("heuristic", "heuristic_first", "hybrid"))
def test_model_selection_and_classifier_switching_do_not_claim_tuning(self, classifier_type: str) -> None:
config: Final = {
"classifier_type": classifier_type,
**({"classifier_llm_config": {"model": "judge"}} if classifier_type != "heuristic" else {}),
**({"heuristic_first_max_tier": "MEDIUM"} if classifier_type == "heuristic_first" else {}),
**({"hybrid_boundary_margin": 0.1} if classifier_type == "hybrid" else {}),
"tiers": _ALT_TIERS,
"escalation_keywords": ["LITELLM ESCALATE"],
"tier_model_configs": {"COMPLEX": [{"model_name": "other-strong", "litellm_params": {"temperature": 0.1}}]},
}
tuned: Final = _router("tuned", {"dimension_weights": {"codePresence": 0.9}})
model_only: Final = _router("model-only", {"tiers": _TIERS})
candidate: Final = _router("another", config)
assert tuning_fingerprint(config) == DEFAULT_TUNING_FINGERPRINT
assert tuning_quota_violation(candidate=candidate, others=(tuned, model_only), baselines={}, limit=1) is None
def test_disabling_or_replacing_escalation_is_still_a_custom_rule(self) -> None:
assert tuning_fingerprint({"escalation_keywords": []}) != DEFAULT_TUNING_FINGERPRINT
assert tuning_fingerprint({"escalation_keywords": ["USE A STRONGER MODEL"]}) != DEFAULT_TUNING_FINGERPRINT
def test_non_tuning_fields_do_not_change_the_fingerprint(self) -> None:
assert (
@ -230,17 +231,18 @@ class TestQuota:
def test_router_added_after_snapshot_is_mutable_only_when_tuned(self) -> None:
baselines = snapshot_tuning_baselines([_router("a", {"tiers": _TIERS})])
assert mutable_tuned_identities([_router("new", {})], baselines) == frozenset()
assert mutable_tuned_identities([_router("new", {"tiers": _TIERS})], baselines) == {
router_identity(_router("new", {}))
}
assert mutable_tuned_identities([_router("new", {"tiers": _TIERS})], baselines) == frozenset()
assert mutable_tuned_identities(
[_router("new", {"tiers": _TIERS, "code_keywords": ["internal-api"]})], baselines
) == {router_identity(_router("new", {}))}
def test_quota_matrix(self) -> None:
legacy_a = _router("a", {"tiers": _TIERS})
legacy_b = _router("b", {"tiers": _ALT_TIERS})
baselines = snapshot_tuning_baselines([legacy_a, legacy_b])
edited_a = _router("a", {"tiers": _TIERS, "dimension_weights": {"codePresence": 0.9}})
edited_b = _router("b", {"tiers": _TIERS})
new_c = _router("c", {"tiers": _TIERS})
edited_b = _router("b", {"tiers": _TIERS, "code_keywords": ["internal-api"]})
new_c = _router("c", {"tiers": _TIERS, "code_keywords": ["internal-api"]})
assert tuning_quota_violation(candidate=edited_a, others=[legacy_b], baselines=baselines, limit=1) is None
assert (
@ -260,7 +262,7 @@ class TestQuota:
legacy_a = _router("a", {"tiers": _TIERS})
legacy_b = _router("b", {"tiers": _ALT_TIERS})
baselines = snapshot_tuning_baselines([legacy_a, legacy_b])
edited_b = _router("b", {"tiers": _TIERS})
edited_b = _router("b", {"tiers": _TIERS, "code_keywords": ["internal-api"]})
assert tuning_quota_violation(candidate=edited_b, others=[legacy_a], baselines=baselines, limit=1) is None
assert (
tuning_quota_violation(candidate=edited_b, others=[legacy_a, edited_b], baselines=baselines, limit=1)
@ -304,5 +306,6 @@ class TestQuota:
assert message is not None
assert "At most 1 auto-router(s)" in message
assert "revert the other changed router to its baseline" in message
assert "Selecting models does not use this allowance" in message
assert tuning_limit_violation(held=1, limit=1) is None
assert tuning_limit_violation(held=5, limit=None) is None

View file

@ -595,6 +595,8 @@ def test_completion_cost_image_generation_registered_deployment_price_keeps_map_
deployment_id,
{"mode": "image_generation", "litellm_provider": "gemini", "output_cost_per_image": 0.1},
)
map_model: Final = "gemini/gemini-3.1-flash-image"
row: Final = litellm.model_cost[map_model]
usage: Final = ImageUsage(
input_tokens=10,
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=0, text_tokens=10),
@ -604,7 +606,7 @@ def test_completion_cost_image_generation_registered_deployment_price_keeps_map_
cost = completion_cost(
completion_response=ImageResponse(data=[ImageObject(url="https://example.com/img.png")], usage=usage),
model="gemini/gemini-3.1-flash-image-preview",
model=map_model,
custom_llm_provider="gemini",
call_type="image_generation",
custom_pricing=True,
@ -612,7 +614,10 @@ def test_completion_cost_image_generation_registered_deployment_price_keeps_map_
litellm_logging_obj=SimpleNamespace(litellm_params={"metadata": {"model_info": {"id": deployment_id}}}),
)
assert cost == pytest.approx(10 * 5e-07 + 1290 * 6e-05)
expected: Final = (
usage.input_tokens * row["input_cost_per_token"] + usage.output_tokens * row["output_cost_per_image_token"]
)
assert cost == pytest.approx(expected)
def test_completion_cost_image_generation_ignores_deployment_model_info_without_custom_pricing(

View file

@ -2,7 +2,7 @@
Tests for _redact_string usage in error/logging paths.
Covers actual execution of redaction in:
- WebSocket close reasons in realtime handlers (openai, azure, bedrock)
- WebSocket close reasons in realtime handlers (openai, bedrock)
- Gemini RAG ingestion x-goog-api-key header usage
- Traceback redaction pattern used in proxy streaming
- Router fallback-failure traceback redaction
@ -72,25 +72,6 @@ class TestOpenAIRealtimeRedaction:
api_key="test-key",
)
@pytest.mark.asyncio
async def test_invalid_status_code_redacts_reason(self):
import websockets.exceptions
from litellm.llms.openai.realtime.handler import OpenAIRealtime
handler = OpenAIRealtime()
exc = websockets.exceptions.InvalidStatusCode(403, None)
exc.status_code = 403
kwargs = self._call_kwargs()
mock_ws = kwargs["websocket"]
p1, p2, p3 = self._make_patches(handler)
with p1, p2, p3, patch("websockets.connect", side_effect=exc):
await handler.async_realtime(**kwargs)
mock_ws.close.assert_called_once()
assert mock_ws.close.call_args[1]["code"] == 403
@pytest.mark.asyncio
async def test_generic_exception_redacts_reason(self):
from litellm.llms.openai.realtime.handler import OpenAIRealtime
@ -111,41 +92,6 @@ class TestOpenAIRealtimeRedaction:
assert "sk-1234567890abcdefghij" not in mock_ws.close.call_args[1]["reason"]
class TestAzureRealtimeRedaction:
"""Test that Azure realtime handler redacts secrets in websocket close reasons."""
@pytest.mark.asyncio
async def test_invalid_status_code_redacts_reason(self):
import websockets.exceptions
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
handler = AzureOpenAIRealtime()
mock_ws = AsyncMock()
exc = websockets.exceptions.InvalidStatusCode(403, None)
exc.status_code = 403
with (
patch.object(
handler,
"_construct_url",
return_value="wss://test.openai.azure.com/openai/realtime",
),
patch("websockets.connect", side_effect=exc),
):
await handler.async_realtime(
model="gpt-4",
websocket=mock_ws,
logging_obj=MagicMock(),
api_base="https://test.openai.azure.com/",
api_key="test-key",
api_version="2024-10-01-preview",
)
mock_ws.close.assert_called_once()
assert mock_ws.close.call_args[1]["code"] == 403
class TestBedrockRealtimeRedaction:
"""Test that _redact_string produces safe close reasons for Bedrock-style errors."""

View file

@ -0,0 +1,74 @@
import os
import subprocess
import sys
from pathlib import Path
from types import MappingProxyType
from typing import Final
import pytest
import yaml
_REPO_ROOT: Final = Path(__file__).resolve().parents[2]
_BASE_WORKFLOW: Final = _REPO_ROOT / ".github" / "workflows" / "_test-unit-base.yml"
_SHARD_ENV: Final = MappingProxyType(
{"MAX_FAILURES": "10", "RERUNS": "0", "DIST": "loadscope", "TEST_TIMEOUT_SECONDS": "60", "COVERAGE_CORE": "sysmon"}
)
_UV_SHIM: Final = f'#!/usr/bin/env bash\nshift 2\nexec "{sys.executable}" -m "$@"\n'
_PASSING_TEST: Final = "def test_passes():\n assert True\n"
_FAILING_TEST: Final = "def test_fails():\n assert False\n"
def _run_tests_script() -> str:
workflow: Final = yaml.safe_load(_BASE_WORKFLOW.read_text())
return next(step["run"] for step in workflow["jobs"]["run"]["steps"] if step.get("name") == "Run tests")
def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.CompletedProcess[str]:
shim_dir: Final = tmp_path / "bin"
shim_dir.mkdir()
(shim_dir / "uv").write_text(_UV_SHIM)
(shim_dir / "uv").chmod(0o755)
(tmp_path / "pyproject.toml").write_text("[tool.pytest.ini_options]\naddopts = '-p no:cacheprovider'\n")
return subprocess.run(
("bash", "--noprofile", "--norc", "-eo", "pipefail", "-c", _run_tests_script()),
cwd=tmp_path,
env={
**os.environ,
**_SHARD_ENV,
"PATH": f"{shim_dir}{os.pathsep}{os.environ['PATH']}",
"TEST_PATH": test_path,
"WORKERS": workers,
},
capture_output=True,
text=True,
timeout=120,
check=False,
)
def _write_passing_test(tmp_path: Path) -> Path:
present: Final = tmp_path / "tests" / "present"
present.mkdir(parents=True)
(present / "test_present.py").write_text(_PASSING_TEST)
return present
@pytest.mark.parametrize("workers", ("0", "2"), ids=("serial", "xdist"))
def test_a_missing_path_is_dropped_and_the_existing_paths_still_run(tmp_path: Path, workers: str) -> None:
_write_passing_test(tmp_path)
result: Final = _run_shard(tmp_path, "tests/gone tests/present", workers)
assert result.returncode == 0, result.stdout + result.stderr
assert "1 passed" in result.stdout, result.stdout
assert "::warning::tests/gone does not exist" in result.stdout
def test_ignore_flags_survive_the_path_filter(tmp_path: Path) -> None:
present: Final = _write_passing_test(tmp_path)
(present / "test_ignored.py").write_text(_FAILING_TEST)
result: Final = _run_shard(tmp_path, "tests/present --ignore=tests/present/test_ignored.py", "0")
assert result.returncode == 0, result.stdout + result.stderr
assert "1 passed" in result.stdout, result.stdout

View file

@ -32,27 +32,35 @@ from litellm._logging import (
from litellm.caching.caching import Cache
from litellm.caching.caching_handler import _PENDING_CACHE_WRITES
from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
from litellm.proxy.utils import is_valid_api_key
from litellm.types.integrations.custom_logger import HEADROOM_CONVERTED_STREAM_KEY
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams
from litellm.types.utils import (
ADDRESSED_RESPONSE_ID_FIELD,
CallTypes,
Choices,
Delta,
EmbeddingResponse,
ImageResponse,
LlmProviders,
LLMResponseTypes,
ModelResponse,
ModelResponseStream,
PromptTokensDetailsWrapper,
RerankResponse,
StreamingChoices,
TranscriptionResponse,
Usage,
all_litellm_params,
bedrock_batch_litellm_params,
)
from litellm.types.videos.main import VideoObject
from litellm.utils import (
CustomStreamWrapper,
ProviderConfigManager,
@ -631,6 +639,9 @@ def validate_model_cost_values(model_data, exceptions=None):
"output_cost_per_character",
"input_cost_per_image",
"output_cost_per_image",
"output_cost_per_image_512",
"output_cost_per_image_1024",
"output_cost_per_image_1536",
"input_cost_per_pixel",
"output_cost_per_pixel",
"input_cost_per_second",
@ -858,6 +869,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"output_cost_per_character": {"type": "number"},
"output_cost_per_character_above_128k_tokens": {"type": "number"},
"output_cost_per_image": {"type": "number"},
"output_cost_per_image_512": {"type": "number"},
"output_cost_per_image_1024": {"type": "number"},
"output_cost_per_image_1536": {"type": "number"},
"output_cost_per_image_token": {"type": "number"},
"output_cost_per_video_token": {"type": "number"},
"output_cost_per_pixel": {"type": "number"},
@ -4406,6 +4420,111 @@ async def test_converted_chat_stream_hook_skips_unhandled_wrappers(
assert wrapper.completion_stream is completion_stream
class _ChatShapedSuccessDeploymentHook(CustomLogger):
async def async_post_call_success_deployment_hook(
self, request_data: dict[str, object], response: object, call_type: CallTypes | None
) -> None:
raise AttributeError(f"{type(response).__name__!r} object has no attribute 'choices'")
class _RecordingSuccessDeploymentHook(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.seen_responses: tuple[object, ...] = ()
async def async_post_call_success_deployment_hook(
self, request_data: dict[str, object], response: object, call_type: CallTypes | None
) -> None:
self.seen_responses = (*self.seen_responses, response)
_SUCCESS_RESPONSES_BY_CALL_TYPE: Final = (
pytest.param(
VideoObject(id="video_abc", object="video", status="queued", model="sora-2", seconds="4", size="720x1280"),
CallTypes.avideo_generation,
id="video",
),
pytest.param(EmbeddingResponse(model="text-embedding-3-small"), CallTypes.aembedding, id="embedding"),
pytest.param(
ResponsesAPIResponse(
id="resp_abc", created_at=1, output=[], parallel_tool_calls=False, tool_choice="auto", tools=[], model="gpt-5.6"
),
CallTypes.aresponses,
id="responses",
),
pytest.param(ImageResponse(), CallTypes.aimage_generation, id="image"),
pytest.param(RerankResponse(id="rerank_abc"), CallTypes.arerank, id="rerank"),
pytest.param(TranscriptionResponse(text="hi"), CallTypes.atranscription, id="transcription"),
pytest.param(ModelResponse(model="gpt-5.6"), CallTypes.acompletion, id="chat"),
pytest.param(ModelResponse(model="claude-sonnet-4-5"), CallTypes.aanthropic_messages, id="anthropic_messages"),
)
@pytest.mark.asyncio
@pytest.mark.parametrize(("response", "call_type"), _SUCCESS_RESPONSES_BY_CALL_TYPE)
async def test_success_deployment_hook_raising_keeps_response_and_runs_later_hooks(
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, response: object, call_type: CallTypes
) -> None:
second_hook: Final = _RecordingSuccessDeploymentHook()
monkeypatch.setattr(litellm, "callbacks", [_ChatShapedSuccessDeploymentHook(), second_hook])
with caplog.at_level(logging.ERROR, logger=verbose_logger.name):
result: Final = await async_post_call_success_deployment_hook(
request_data={"model": "m"}, response=response, call_type=call_type
)
assert result is response
assert second_hook.seen_responses == (response,)
failure_logs: Final = tuple(r for r in caplog.records if "async_post_call_success_deployment_hook error" in r.message)
assert len(failure_logs) == 1
assert "_ChatShapedSuccessDeploymentHook" in failure_logs[0].message
assert str(call_type) in failure_logs[0].message
assert failure_logs[0].exc_info is not None
@pytest.mark.asyncio
async def test_success_deployment_hook_raising_keeps_earlier_hook_rewrite(monkeypatch: pytest.MonkeyPatch) -> None:
rewriter: Final = _RewritingSuccessDeploymentHook()
trailing_hook: Final = _RecordingSuccessDeploymentHook()
monkeypatch.setattr(litellm, "callbacks", [rewriter, _ChatShapedSuccessDeploymentHook(), trailing_hook])
original: Final = ModelResponse(model="gpt-5.6")
result: Final = await async_post_call_success_deployment_hook(
request_data={"model": "gpt-5.6"}, response=original, call_type=CallTypes.acompletion
)
assert isinstance(result, ModelResponse)
assert result is not original
assert result.choices[0].message.content == "rewritten by deployment hook"
assert trailing_hook.seen_responses == (result,)
class _GuardrailBlocked(Exception):
pass
class _BlockingSuccessDeploymentGuardrail(CustomGuardrail):
async def async_post_call_success_deployment_hook(
self, request_data: dict, response: LLMResponseTypes, call_type: CallTypes | None
) -> LLMResponseTypes | None:
raise _GuardrailBlocked("Violated moderation policy")
@pytest.mark.asyncio
async def test_success_deployment_hook_still_propagates_guardrail_block(monkeypatch: pytest.MonkeyPatch) -> None:
later_hook: Final = _RewritingSuccessDeploymentHook()
monkeypatch.setattr(
litellm, "callbacks", [_BlockingSuccessDeploymentGuardrail(guardrail_name="blocking"), later_hook]
)
with pytest.raises(_GuardrailBlocked):
await async_post_call_success_deployment_hook(
request_data={"model": "gpt-5.6"}, response=ModelResponse(model="gpt-5.6"), call_type=CallTypes.acompletion
)
assert later_hook.seen_responses == ()
@pytest.mark.asyncio
@respx.mock
async def test_wrapper_async_leaves_success_deployment_hook_off_requested_fake_stream(

View file

@ -2120,6 +2120,7 @@ class TestPrismaTableRepository:
"litellm_prompttable",
"litellm_searchtoolstable",
"litellm_ssoconfig",
"litellm_uisettings",
}
)

View file

@ -13,6 +13,7 @@
"@headlessui/tailwindcss": "0.2.2",
"@heroicons/react": "1.0.6",
"@hookform/resolvers": "5.4.0",
"@shadcn/react": "0.3.1",
"@tanstack/react-pacer": "0.22.1",
"@tanstack/react-query": "5.100.7",
"@tanstack/react-table": "8.21.3",
@ -2990,6 +2991,24 @@
"dev": true,
"license": "MIT"
},
"node_modules/@shadcn/react": {
"version": "0.3.1",
"resolved": "https://registry.npmjs.org/@shadcn/react/-/react-0.3.1.tgz",
"integrity": "sha512-2gOR0HDMtWeRsCZfNDaU0YDFdgH3zsDQ6lz67Fv/y/qjY9y+R8kJqAn6q56phqv7/zHi0wqURjntrNO9zL7vnQ==",
"license": "MIT",
"peerDependencies": {
"@types/react": ">=19",
"react": ">=19"
},
"peerDependenciesMeta": {
"@types/react": {
"optional": true
},
"react": {
"optional": true
}
}
},
"node_modules/@standard-schema/spec": {
"version": "1.1.0",
"resolved": "https://registry.npmjs.org/@standard-schema/spec/-/spec-1.1.0.tgz",

View file

@ -29,6 +29,7 @@
"@headlessui/tailwindcss": "0.2.2",
"@heroicons/react": "1.0.6",
"@hookform/resolvers": "5.4.0",
"@shadcn/react": "0.3.1",
"@tanstack/react-pacer": "0.22.1",
"@tanstack/react-query": "5.100.7",
"@tanstack/react-table": "8.21.3",

View file

@ -1,4 +1,4 @@
import { fetchProxySettings } from "@/utils/proxyUtils";
import { getProxyBaseUrl, getProxyUISettings } from "@/components/networking";
import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
@ -18,11 +18,19 @@ const EMPTY_PROXY_SETTINGS: ProxySettings = {
LITELLM_UI_API_DOC_BASE_URL: null,
};
export default function useProxySettings(accessToken: string | null): ProxySettings {
const { data } = useQuery({
queryKey: [...proxySettingsKeys.all, accessToken],
queryFn: () => fetchProxySettings(accessToken),
export function useProxySettingsQuery(accessToken: string | null) {
const managementBaseUrl = getProxyBaseUrl();
return useQuery({
queryKey: [...proxySettingsKeys.all, managementBaseUrl, accessToken],
queryFn: () => {
if (getProxyBaseUrl() !== managementBaseUrl) throw new Error("Gateway changed while loading settings.");
return accessToken ? getProxyUISettings(accessToken) : null;
},
enabled: Boolean(accessToken),
});
}
export default function useProxySettings(accessToken: string | null): ProxySettings {
const { data } = useProxySettingsQuery(accessToken);
return data ?? EMPTY_PROXY_SETTINGS;
}

View file

@ -13,6 +13,7 @@ import { NoRedisWarningBanner } from "@/components/NoRedisWarningBanner";
import { EnvCredentialLoginWarningBanner } from "@/components/EnvCredentialLoginWarningBanner";
import { LicenseExpiryBanner } from "@/components/LicenseExpiryBanner";
import { UserBanner } from "@/components/UserBanner";
import LiteAdmin from "@/components/liteadmin/LiteAdmin";
import { UpgradeBanner } from "@/components/UpgradeBanner";
import { uiHref } from "@/utils/uiHref";
import { PluginModeProvider, usePluginMode } from "@/contexts/PluginModeContext";
@ -141,6 +142,7 @@ function DashboardShell({ children }: { children: React.ReactNode }) {
<UserBanner accessToken={accessToken} />
<UpgradeBanner accessToken={accessToken} />
<main className="min-w-0 flex-1 overflow-y-auto">{children}</main>
<LiteAdmin />
</div>
</div>
);

View file

@ -56,14 +56,16 @@ export function ChatComposer({
{showSuggestions && suggestions.length > 0 && (
<div className="flex w-full flex-col gap-1.5" data-testid="chat-suggested-actions">
{suggestions.map((suggestion) => (
<button
<Button
key={suggestion}
type="button"
className="w-full truncate rounded-lg border border-border/50 bg-card/30 px-3 py-1.5 text-left text-[12px] leading-snug text-muted-foreground transition-colors hover:bg-card/60 hover:text-foreground"
variant="outline"
size="sm"
className="w-full justify-start overflow-hidden text-xs text-muted-foreground"
onClick={() => onSuggestionSelect?.(suggestion)}
>
{suggestion}
</button>
<span className="truncate">{suggestion}</span>
</Button>
))}
</div>
)}

View file

@ -0,0 +1,165 @@
import { createContext, useContext, useEffect, useState } from "react";
import { useQuery, type UseQueryOptions } from "@tanstack/react-query";
import { apiClient } from "@/components/networking";
import { Popover, PopoverContent, PopoverTitle, PopoverTrigger } from "@/components/ui/popover";
import type { components } from "@/lib/http/schema";
type Availability = components["schemas"]["AutoRouterAvailabilityResponse"];
type Request = components["schemas"]["AutoRouterAvailabilityRequest"];
export type Allowance = components["schemas"]["AutoRouterAllowance"];
type AvailabilityState = {
data?: Availability;
isPending: boolean;
isError: boolean;
isChecking?: boolean;
refetch?: () => unknown;
};
export const AutoRouterAvailabilityContext = createContext<AvailabilityState>({ isPending: true, isError: false });
export const useAutoRouterAvailability = (accessToken: string, body: Request, enabled = true) => {
const serialized = JSON.stringify(body.complexity_router_config ?? null);
const [debounced, setDebounced] = useState(serialized);
useEffect(() => {
const timeout = setTimeout(() => setDebounced(serialized), 300);
return () => clearTimeout(timeout);
}, [serialized]);
const options: UseQueryOptions<Availability> = {
queryKey: ["autoRouterAvailability", accessToken, body.team_id, body.saved_model_id, debounced],
queryFn: ({ signal }) =>
apiClient.post<Availability>("/auto_router/availability", {
accessToken,
body: { ...body, complexity_router_config: JSON.parse(debounced) },
signal,
}),
enabled: enabled && Boolean(accessToken),
placeholderData: (previous, previousQuery) => {
const key = previousQuery?.queryKey;
return key?.[1] === accessToken && key[2] === body.team_id && key[3] === body.saved_model_id
? previous
: undefined;
},
refetchOnMount: "always",
staleTime: 0,
retry: false,
};
const query = useQuery(options);
const isChecking = query.isFetching || query.isPlaceholderData || serialized !== debounced;
const saveBlockedReason = () => {
if (!enabled) return null;
if (query.isPending || isChecking) return "Checking availability";
if (query.isError || !query.data) return "Could not check availability. Retry before saving.";
return query.data.error ?? null;
};
return {
...query,
isPending: query.isPending || (query.isFetching && !query.isFetchedAfterMount),
isChecking,
saveBlockedReason: saveBlockedReason(),
};
};
export const allowanceLabel = (allowance?: Allowance): string | null => {
if (!allowance?.available) return "Availability unavailable";
if (allowance.limit == null) return null;
if (allowance.used_by_this_router) return "Used by this router";
return `${allowance.remaining} of ${allowance.limit} available`;
};
const availabilityLabel = (state: AvailabilityState, key: string) => {
if (state.isPending || state.isChecking) return "Checking availability";
if (state.isError) return "Availability unavailable";
return allowanceLabel(state.data?.allowances.find((entry) => entry.key === key));
};
export const useAllowanceLabel = (key: string) => availabilityLabel(useContext(AutoRouterAvailabilityContext), key);
export const isAllowanceExhausted = (allowance?: Allowance) =>
Boolean(allowance?.available && allowance.limit != null && allowance.remaining === 0) &&
!allowance?.used_by_this_router;
export const AUTO_ROUTER_CONTACT_URL = "https://calendly.com/tin-berri/litellm-auto-router-pricing-discussion";
export const AutoRouterContactLink = ({ features, message }: { features?: string[]; message?: string }) => {
const state = useContext(AutoRouterAvailabilityContext);
if (state.isPending || state.isError || state.isChecking) return null;
const exhausted = state.data?.allowances.some(
(entry) => (!features || features.includes(entry.key)) && isAllowanceExhausted(entry),
);
if (!exhausted) return null;
return (
<span className="inline-flex flex-wrap items-baseline gap-x-1 text-xs leading-5 text-muted-foreground">
{message}
<a
href={AUTO_ROUTER_CONTACT_URL}
target="_blank"
rel="noopener noreferrer"
className="font-medium text-blue-600 hover:underline dark:text-blue-400"
>
Talk to our team
</a>
</span>
);
};
export const AutoRouterAllowanceLabel = ({ feature }: { feature: string }) => {
const label = useAllowanceLabel(feature);
return label ? (
<span className="shrink-0 whitespace-nowrap text-xs leading-5 tabular-nums text-muted-foreground">{label}</span>
) : null;
};
export const AutoRouterAllowanceNote = ({ feature, label }: { feature: string; label: string }) => {
const availability = useAllowanceLabel(feature);
return availability ? (
<p className="text-xs leading-5 text-muted-foreground">
{label}: {availability} <AutoRouterContactLink features={[feature]} />
</p>
) : null;
};
export const AutoRouterLimits = () => {
const state = useContext(AutoRouterAvailabilityContext);
const limits = [
["heuristic_v2", "Heuristic v2 routers"],
["capability", "Capability routers"],
["llm_v2", "Fuse v2 routers"],
["tier_or_classifier_prompt", "Custom tiers or prompts"],
["heuristic_tuning", "Rule-based tuning"],
];
return (
<Popover>
<PopoverTrigger className="shrink-0 whitespace-nowrap text-xs font-normal text-muted-foreground underline underline-offset-4 hover:text-foreground">
View limits
</PopoverTrigger>
<PopoverContent align="end" className="w-96 max-w-[calc(100vw-2rem)] gap-3">
<PopoverTitle>Routing and customization limits</PopoverTitle>
<p className="text-xs leading-5 text-muted-foreground">
Rule-based, Complexity, and Jev are unlimited with built-in settings. Choose or change tier models freely.
Customization allowances are shared across this proxy.
</p>
<dl className="space-y-2 text-xs">
{limits.map(([key, label]) => (
<div key={key} className="flex items-center justify-between gap-3">
<dt>{label}</dt>
<dd className="shrink-0 tabular-nums text-muted-foreground">
{availabilityLabel(state, key) ?? "Unlimited"}
</dd>
</div>
))}
</dl>
<p className="text-xs leading-5 text-muted-foreground">
Custom tier definitions and written classifier instructions share one allowance. Built-in prompts and
display-name changes do not use it.
</p>
<p className="text-xs leading-5 text-muted-foreground">
Changing scoring rules, such as weights, thresholds, keywords, or custom dimensions, uses the Rule-based
tuning allowance. It also applies to Heuristic first and Hybrid. Recorded settings on existing routers are
preserved; new routers start from built-in rules.
</p>
<AutoRouterContactLink message="Need a higher limit?" />
</PopoverContent>
</Popover>
);
};

View file

@ -1,7 +1,9 @@
import React, { useState } from "react";
import { describe, expect, it, vi } from "vitest";
import { fireEvent, renderWithProviders, screen } from "../../../tests/test-utils";
import { fireEvent, renderWithProviders, screen, waitFor, within } from "../../../tests/test-utils";
import { selectAutoRouterOption } from "../../../tests/autoRouterSetup";
import AutoRouterClassifierTabs from "./AutoRouterClassifierTabs";
import { AutoRouterAllowanceNote, AutoRouterAvailabilityContext } from "./AutoRouterAvailability";
import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
const initial: ComplexityRouterConfigValue = {
@ -9,28 +11,67 @@ const initial: ComplexityRouterConfigValue = {
tiers: { SIMPLE: ["efficient"], MEDIUM: [], COMPLEX: [], REASONING: ["capable"] },
};
function Form({ initialValue = initial }: { initialValue?: ComplexityRouterConfigValue }) {
function Form({
initialValue = initial,
remaining = 1,
limit = 1,
ownedFeature,
availabilityState,
}: {
initialValue?: ComplexityRouterConfigValue;
remaining?: number;
limit?: number | null;
ownedFeature?: string;
availabilityState?: Partial<React.ContextType<typeof AutoRouterAvailabilityContext>>;
}) {
const [value, setValue] = useState(initialValue);
return (
<AutoRouterClassifierTabs value={value} onChange={setValue}>
<output aria-label="Classifier type">{value.classifier_type}</output>
</AutoRouterClassifierTabs>
<AutoRouterAvailabilityContext.Provider
value={{
isPending: false,
isError: false,
data: {
allowances: ["heuristic_v2", "capability", "llm_v2", "tier_or_classifier_prompt", "heuristic_tuning"].map(
(key) => ({
key,
limit,
remaining,
available: true,
used_by_this_router: key === ownedFeature,
}),
),
error: null,
},
...availabilityState,
}}
>
<AutoRouterClassifierTabs value={value} onChange={setValue}>
<output aria-label="Classifier type">{value.classifier_type}</output>
</AutoRouterClassifierTabs>
</AutoRouterAvailabilityContext.Provider>
);
}
describe("AutoRouterClassifierTabs", () => {
it.each(["heuristic", "heuristic_v2", "llm", "heuristic_first", "hybrid"] as const)(
"groups %s under Complexity without resetting its configuration",
(classifier_type) => {
describe("Auto-router classifier selection", () => {
it.each(["heuristic", "heuristic_v2", "llm", "heuristic_first", "hybrid", "jev"] as const)(
"shows saved %s without changing its configuration",
async (classifier_type) => {
const onChange = vi.fn();
renderWithProviders(
<AutoRouterClassifierTabs value={{ ...initial, classifier_type }} onChange={onChange}>
Existing classifier settings
Existing settings
</AutoRouterClassifierTabs>,
);
expect(screen.getByRole("tab", { name: "Complexity" })).toHaveAttribute("aria-selected", "true");
expect(screen.getByRole("tabpanel", { name: "Complexity" })).toHaveTextContent("Existing classifier settings");
fireEvent.click(screen.getByRole("tab", { name: "Complexity" }));
const family = {
heuristic: "Heuristics",
heuristic_v2: "Heuristics",
llm: "LLM",
heuristic_first: "LLM",
hybrid: "LLM",
jev: "Jev",
}[classifier_type];
expect(screen.getByRole("radio", { name: new RegExp(`^${family}$`) })).toBeChecked();
fireEvent.click(screen.getByRole("radio", { name: new RegExp(`^${family}$`) }));
expect(onChange).not.toHaveBeenCalled();
},
);
@ -38,38 +79,186 @@ describe("AutoRouterClassifierTabs", () => {
it.each([
["capability", "Capability"],
["llm_v2", "Fuse v2"],
] as const)("opens saved %s settings and switches back to local Complexity", (classifier_type, label) => {
renderWithProviders(<Form initialValue={{ ...initial, classifier_type }} />);
expect(screen.getByRole("tab", { name: label })).toHaveAttribute("aria-selected", "true");
fireEvent.click(screen.getByRole("tab", { name: "Complexity" }));
expect(screen.getByRole("tab", { name: "Complexity" })).toHaveAttribute("aria-selected", "true");
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent("heuristic");
] as const)(
"opens saved %s and retains the LLM family when switching to Complexity",
async (classifier_type, label) => {
renderWithProviders(<Form initialValue={{ ...initial, classifier_type }} />);
expect(screen.getByRole("button", { name: "Routing approach" })).toHaveTextContent(label);
await selectAutoRouterOption("Routing approach", "Complexity");
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent("llm");
},
);
it.each([
[1, "heuristic"],
[0, "heuristic"],
])("defaults to Rule-based when %s v2 slots remain", async (remaining, classifier) => {
renderWithProviders(<Form remaining={Number(remaining)} />);
fireEvent.click(screen.getByRole("radio", { name: /^Heuristics$/ }));
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent(String(classifier));
fireEvent.click(screen.getByRole("button", { name: "Heuristic" }));
expect(screen.getByRole("menuitemradio", { name: /^Heuristic v2/ })).toHaveTextContent(
`${remaining} of 1 available`,
);
});
it("keeps custom tiers editable under Complexity and explains why forecast tabs are disabled", () => {
it.each([
{ data: undefined },
{ isPending: true },
{ isError: true },
{ isChecking: true },
{ data: { allowances: [], error: null } },
{ data: { allowances: [{ key: "heuristic_v2", limit: 1, remaining: null, available: false }], error: null } },
])("uses Rule-based when v2 availability is unverified: %j", async (availabilityState) => {
renderWithProviders(<Form availabilityState={availabilityState} />);
fireEvent.click(screen.getByRole("radio", { name: "Heuristics" }));
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent(/^heuristic$/);
});
it("does not present Rule-based as having a classifier quota", () => {
renderWithProviders(<Form initialValue={{ ...initial, classifier_type: "heuristic" }} remaining={0} />);
expect(screen.getByRole("button", { name: "Heuristic" })).toHaveTextContent(/^Rule-based/);
fireEvent.click(screen.getByRole("button", { name: "Heuristic" }));
expect(screen.getAllByRole("menuitemradio")[0]).toHaveTextContent(/^Rule-based/);
expect(screen.getByRole("menuitemradio", { name: /^Rule-based/ })).not.toHaveTextContent("of 1 available");
expect(screen.getByRole("menuitemradio", { name: /^Heuristic v2/ })).toHaveTextContent("0 of 1 available");
});
it("omits allowance labels with an unlimited entitlement", async () => {
renderWithProviders(<Form initialValue={{ ...initial, classifier_type: "heuristic_v2" }} limit={null} />);
fireEvent.click(screen.getByRole("button", { name: "Heuristic" }));
expect(screen.getByRole("menuitemradio", { name: /^Heuristic v2/ })).not.toHaveTextContent("available");
});
it.each([
["heuristic", "Heuristic", "Heuristic v2"],
["llm", "Routing approach", "Capability"],
["llm", "Routing approach", "Fuse v2"],
] as const)("blocks exhausted %s options: %s / %s", (classifier_type, field, option) => {
renderWithProviders(<Form initialValue={{ ...initial, classifier_type }} remaining={0} />);
fireEvent.click(screen.getByRole("button", { name: field }));
const unavailable = screen.getByRole("menuitemradio", { name: new RegExp(`^${option}`) });
expect(unavailable).toHaveAttribute("aria-disabled", "true");
fireEvent.click(unavailable);
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent(classifier_type);
});
it.each([
["heuristic_v2", "heuristic", "Heuristic", "Heuristic v2"],
["capability", "llm", "Routing approach", "Capability"],
["llm_v2", "llm", "Routing approach", "Fuse v2"],
] as const)("lets a saved router reselect its own %s allowance", async (feature, classifier_type, field, option) => {
renderWithProviders(<Form initialValue={{ ...initial, classifier_type }} remaining={0} ownedFeature={feature} />);
await selectAutoRouterOption(field, option);
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent(feature);
expect(screen.getByRole("button", { name: field })).toHaveTextContent("Used by this router");
});
it("shows Jev's single Complexity approach without changing saved configuration", () => {
const onChange = vi.fn();
renderWithProviders(
<AutoRouterClassifierTabs
value={{
<AutoRouterClassifierTabs value={{ ...initial, classifier_type: "jev" }} onChange={onChange}>
Existing settings
</AutoRouterClassifierTabs>,
);
expect(screen.getByRole("button", { name: "Routing approach" })).toHaveTextContent("ComplexityUnlimited");
fireEvent.click(screen.getByRole("button", { name: "Routing approach" }));
expect(screen.getAllByRole("menuitemradio")).toHaveLength(1);
fireEvent.click(screen.getByRole("menuitemradio", { name: /^Complexity/ }));
expect(onChange).not.toHaveBeenCalled();
});
it("keeps custom tiers editable and disables incompatible choices", async () => {
renderWithProviders(
<Form
initialValue={{
...initial,
custom_tier_set: {
tiers: [{ id: "review", name: "REVIEW", definition: "Code reviews", models: ["capable"] }],
fallback_tier_id: "review",
},
}}
onChange={onChange}
>
Custom tiers
</AutoRouterClassifierTabs>,
/>,
);
expect(screen.getByRole("tabpanel", { name: "Complexity" })).toHaveTextContent("Custom tiers");
expect(screen.getByRole("radio", { name: /^Heuristics$/ })).toHaveAttribute("aria-disabled", "true");
fireEvent.click(screen.getByRole("button", { name: "Routing approach" }));
for (const name of ["Capability", "Fuse v2"]) {
const tab = screen.getByRole("tab", { name });
expect(tab).toHaveAttribute("aria-disabled", "true");
expect(tab).toHaveAccessibleDescription("Restore standard tiers to use Capability or Fuse v2.");
fireEvent.click(tab);
expect(screen.getByRole("menuitemradio", { name: new RegExp(`^${name}`) })).toHaveAttribute(
"aria-disabled",
"true",
);
}
expect(onChange).not.toHaveBeenCalled();
expect(screen.getByText("Restore standard tiers to use Capability or Fuse v2.")).toBeVisible();
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent("llm");
});
});
describe("Gated routing contact action", () => {
it("offers a pricing discussion in View limits", async () => {
renderWithProviders(<Form remaining={0} />);
expect(screen.queryByRole("link", { name: "Talk to our team" })).not.toBeInTheDocument();
fireEvent.click(screen.getByRole("button", { name: "View limits" }));
const link = within(screen.getByRole("dialog")).getByRole("link", { name: "Talk to our team" });
await waitFor(() => expect(link).toBeVisible());
expect(link).toHaveAttribute("href", "https://calendly.com/tin-berri/litellm-auto-router-pricing-discussion");
expect(link).toHaveAttribute("target", "_blank");
expect(link).toHaveAttribute("rel", "noopener noreferrer");
});
it.each([
["heuristic", "Heuristic", "Heuristic v2"],
["llm", "Routing approach", "Capability"],
] as const)(
"keeps the contact action available beside the disabled %s choice",
async (classifier_type, field, option) => {
renderWithProviders(<Form initialValue={{ ...initial, classifier_type }} remaining={0} />);
expect(screen.queryByText(/Need more/)).not.toBeInTheDocument();
fireEvent.click(screen.getByRole("button", { name: field }));
const disabled = screen.getByRole("menuitemradio", { name: new RegExp(`^${option}`) });
expect(disabled).toHaveAttribute("aria-disabled", "true");
const link = screen.getByRole("menuitem", { name: `Talk to our team about ${option}` });
await waitFor(() => expect(link).toBeVisible());
expect(link).toHaveAttribute("href", "https://calendly.com/tin-berri/litellm-auto-router-pricing-discussion");
expect(link).toHaveAttribute("target", "_blank");
fireEvent.click(link);
expect(screen.getByRole("status", { name: "Classifier type" })).toHaveTextContent(classifier_type);
},
);
it.each([
{ remaining: 1 },
{ remaining: 0, limit: null },
{ remaining: 0, availabilityState: { isPending: true } },
{ remaining: 0, availabilityState: { isError: true } },
{ remaining: 0, availabilityState: { isChecking: true } },
])("does not pitch an upgrade for a free or unverified option: %j", (props) => {
renderWithProviders(<Form {...props} />);
fireEvent.click(screen.getByRole("button", { name: "Routing approach" }));
expect(screen.queryByRole("menuitem", { name: /Talk to our team/ })).not.toBeInTheDocument();
});
it("does not pitch an upgrade for the saved heuristic's own slot", () => {
renderWithProviders(
<Form initialValue={{ ...initial, classifier_type: "heuristic_v2" }} remaining={0} ownedFeature="heuristic_v2" />,
);
fireEvent.click(screen.getByRole("button", { name: "Heuristic" }));
expect(screen.queryByRole("menuitem", { name: /Talk to our team/ })).not.toBeInTheDocument();
});
it("includes the sales action beside customization limits and blocked changes", () => {
const allowance = { key: "tier_or_classifier_prompt", limit: 1, remaining: 0, available: true };
const state = {
isPending: false,
isError: false,
data: { allowances: [allowance], error: "Custom tiers have no available allowance" },
};
renderWithProviders(
<AutoRouterAvailabilityContext.Provider value={state}>
<AutoRouterClassifierTabs value={initial} onChange={vi.fn()}>
<AutoRouterAllowanceNote feature="tier_or_classifier_prompt" label="Custom tiers" />
</AutoRouterClassifierTabs>
</AutoRouterAvailabilityContext.Provider>,
);
expect(screen.getByText(/Custom tiers: 0 of 1 available/)).toHaveTextContent("Talk to our team");
expect(within(screen.getByRole("alert")).getByRole("link", { name: "Talk to our team" })).toBeVisible();
});
});

View file

@ -1,8 +1,119 @@
import React, { useId } from "react";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
import { effectiveClassifierType, type ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
import React, { useContext, useId } from "react";
import { Label } from "@/components/ui/label";
import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group";
import { ChevronDownIcon } from "lucide-react";
import { Button } from "@/components/ui/button";
import {
DropdownMenu,
DropdownMenuContent,
DropdownMenuItem,
DropdownMenuRadioGroup,
DropdownMenuRadioItem,
DropdownMenuTrigger,
} from "@/components/ui/dropdown-menu";
import {
effectiveClassifierType,
type ClassifierType,
type ComplexityRouterConfigValue,
} from "./ComplexityRouterConfig";
import { transitionClassifierType } from "./classifier_type_transition";
import { isForecastClassifier } from "./forecast_classifier_config";
import {
AutoRouterAllowanceLabel,
AutoRouterAvailabilityContext,
AutoRouterLimits,
AutoRouterContactLink,
isAllowanceExhausted,
AUTO_ROUTER_CONTACT_URL,
} from "./AutoRouterAvailability";
function ClassifierOption({
value,
label,
description,
feature,
disabled,
unlimited = true,
}: {
value: string;
label: string;
description: string;
feature?: string;
disabled?: boolean;
unlimited?: boolean;
}) {
const state = useContext(AutoRouterAvailabilityContext);
const allowance = state.data?.allowances.find((entry) => entry.key === feature);
const fresh = !state.isPending && !state.isError && !state.isChecking;
const exhausted = isAllowanceExhausted(allowance);
return (
<div className="relative">
<DropdownMenuRadioItem value={value} disabled={disabled || exhausted} closeOnClick className="py-3">
<span className="grid w-full min-w-0 gap-1 whitespace-normal">
<span className="flex items-center justify-between gap-3">
<span className="font-medium">{label}</span>
{feature ? (
<AutoRouterAllowanceLabel feature={feature} />
) : (
unlimited && <span className="shrink-0 text-xs leading-5 text-muted-foreground">Unlimited</span>
)}
</span>
<span className={`text-xs leading-5 text-muted-foreground ${exhausted ? "pr-28" : ""}`}>{description}</span>
</span>
</DropdownMenuRadioItem>
{fresh && exhausted && (
<DropdownMenuItem
render={<a href={AUTO_ROUTER_CONTACT_URL} target="_blank" rel="noopener noreferrer" />}
aria-label={`Talk to our team about ${label}`}
className="absolute top-9 right-8 cursor-pointer px-0 py-0 text-xs leading-5 font-medium text-blue-600 focus:text-blue-600 hover:underline dark:text-blue-400 dark:focus:text-blue-400"
>
Talk to our team
</DropdownMenuItem>
)}
</div>
);
}
function ClassifierMenu({
id,
label,
value,
selectedLabel,
feature,
onValueChange,
children,
}: {
id: string;
label: string;
value: string;
selectedLabel: string;
feature?: string;
onValueChange: (value: string) => void;
children: React.ReactNode;
}) {
return (
<DropdownMenu>
<DropdownMenuTrigger
id={id}
aria-label={label}
render={<Button variant="outline" className="w-full justify-between font-normal" />}
>
<span className="flex-1 text-left">{selectedLabel}</span>
{feature ? (
<AutoRouterAllowanceLabel feature={feature} />
) : (
<span className="text-xs text-muted-foreground">Unlimited</span>
)}
<ChevronDownIcon className="size-4 text-muted-foreground" />
</DropdownMenuTrigger>
<DropdownMenuContent>
<DropdownMenuRadioGroup value={value} onValueChange={onValueChange}>
{children}
</DropdownMenuRadioGroup>
</DropdownMenuContent>
</DropdownMenu>
);
}
interface AutoRouterClassifierTabsProps {
value: ComplexityRouterConfigValue;
@ -11,47 +122,174 @@ interface AutoRouterClassifierTabsProps {
}
const AutoRouterClassifierTabs: React.FC<AutoRouterClassifierTabsProps> = ({ value, onChange, children }) => {
const restrictionId = useId();
const id = useId();
const availability = useContext(AutoRouterAvailabilityContext);
const classifierType = effectiveClassifierType(value);
const selected = isForecastClassifier(classifierType) ? classifierType : "complexity";
const familyByType: Record<ClassifierType, string> = {
heuristic: "heuristics",
heuristic_v2: "heuristics",
llm: "llm",
heuristic_first: "llm",
hybrid: "llm",
capability: "llm",
llm_v2: "llm",
jev: "jev",
custom: "custom",
};
const family = familyByType[classifierType];
const hasCustomTiers = Boolean(value.custom_tier_set);
const handleChange = (tab: unknown) => {
if (tab === selected) return;
if (tab === "complexity") {
onChange(transitionClassifierType(value, isForecastClassifier(classifierType) ? "heuristic" : classifierType));
} else if (!hasCustomTiers && (tab === "capability" || tab === "llm_v2")) {
onChange(transitionClassifierType(value, tab));
}
const changeType = (next: ClassifierType) => {
if (next !== classifierType) onChange(transitionClassifierType(value, next));
};
const changeFamily = (next: unknown) => {
if (next === family) return;
if (next === "heuristics") changeType("heuristic");
if (next === "llm") changeType("llm");
if (next === "jev") changeType("jev");
};
const approachLabels: Partial<Record<ClassifierType, string>> = { capability: "Capability", llm_v2: "Fuse v2" };
const approachDescription: Partial<Record<ClassifierType, string>> = {
capability: "Use the efficient model when it is likely to succeed",
llm_v2: "Use the efficient model when its predicted quality is close enough to the capable model",
};
return (
<Tabs value={selected} onValueChange={handleChange}>
<p className="text-sm font-medium">Classifier type</p>
<TabsList aria-label="Classifier type" className="w-full">
<TabsTrigger value="complexity">Complexity</TabsTrigger>
<TabsTrigger
value="capability"
disabled={hasCustomTiers}
aria-describedby={hasCustomTiers ? restrictionId : undefined}
>
Capability
</TabsTrigger>
<TabsTrigger
value="llm_v2"
disabled={hasCustomTiers}
aria-describedby={hasCustomTiers ? restrictionId : undefined}
>
Fuse v2
</TabsTrigger>
</TabsList>
<div className="flex flex-col gap-6">
<fieldset>
<legend className="mb-3 flex w-full items-center justify-between gap-3 text-sm font-medium">
What classifies your requests?
<AutoRouterLimits />
</legend>
<RadioGroup value={family} onValueChange={changeFamily} className="grid gap-3 sm:grid-cols-3">
{[
{ value: "heuristics", label: "Heuristics", description: "Classify locally, with no API call" },
{ value: "llm", label: "LLM", description: "Use a judge model to choose a solver" },
{ value: "jev", label: "Jev", description: "Use TypeSafe System One Choice to choose a tier" },
].map((option) => (
<Label
key={option.value}
className="cursor-pointer items-start rounded-lg border p-4 transition-colors hover:bg-muted/50 has-data-checked:border-primary has-data-checked:bg-primary/5"
>
<RadioGroupItem
aria-describedby={`${id}-${option.value}-description`}
value={option.value}
disabled={option.value === "heuristics" && hasCustomTiers}
/>
<span className="space-y-1">
<span className="block font-medium">{option.label}</span>
<span
id={`${id}-${option.value}-description`}
aria-hidden="true"
className="block text-xs font-normal text-muted-foreground"
>
{option.description}
</span>
</span>
</Label>
))}
</RadioGroup>
</fieldset>
{family === "custom" && (
<p className="text-sm text-muted-foreground">This router uses a custom classifier plugin</p>
)}
{family === "heuristics" && (
<div className="space-y-2">
<Label htmlFor={`${id}-heuristic`}>Heuristic</Label>
<ClassifierMenu
id={`${id}-heuristic`}
label="Heuristic"
selectedLabel={classifierType === "heuristic_v2" ? "Heuristic v2" : "Rule-based"}
feature={classifierType === "heuristic_v2" ? "heuristic_v2" : undefined}
value={classifierType}
onValueChange={(next) => {
if (next === "heuristic" || next === "heuristic_v2") changeType(next);
}}
>
<ClassifierOption
value="heuristic"
label="Rule-based"
description="Score requests with local rules to choose a tier, with no API call"
/>
<ClassifierOption
value="heuristic_v2"
label="Heuristic v2"
description="Use calibrated probabilities to match requests to a tier, with no API call"
feature="heuristic_v2"
/>
</ClassifierMenu>
<p className="text-xs text-muted-foreground">
{classifierType === "heuristic_v2"
? "Use calibrated probabilities to match requests to a tier"
: "Match requests using scoring rules. Choose or change tier models freely"}
</p>
</div>
)}
{(family === "llm" || family === "jev") && (
<div className="space-y-2">
<Label htmlFor={`${id}-approach`}>Routing approach</Label>
<ClassifierMenu
id={`${id}-approach`}
label="Routing approach"
selectedLabel={approachLabels[classifierType] ?? "Complexity"}
feature={isForecastClassifier(classifierType) ? classifierType : undefined}
value={isForecastClassifier(classifierType) ? classifierType : "llm"}
onValueChange={(next) => {
if (next === "llm" || next === "capability" || next === "llm_v2") {
if (next === "llm" && !isForecastClassifier(classifierType)) return;
changeType(next);
}
}}
>
<ClassifierOption
value="llm"
label="Complexity"
description="Match task difficulty to a tier, then use one of that tier's models"
/>
{family === "llm" && (
<>
<ClassifierOption
value="capability"
label="Capability"
description="Use the efficient model when it is likely to succeed; otherwise use the capable model"
feature="capability"
disabled={hasCustomTiers}
/>
<ClassifierOption
value="llm_v2"
label="Fuse v2"
description="Use the efficient model when its predicted quality is close enough to the capable model"
feature="llm_v2"
disabled={hasCustomTiers}
/>
</>
)}
</ClassifierMenu>
<p className="text-xs text-muted-foreground">
{approachDescription[classifierType] ?? "Match task difficulty to a tier"}
</p>
</div>
)}
{hasCustomTiers && (
<p id={restrictionId} className="text-sm text-muted-foreground">
Restore standard tiers to use Capability or Fuse v2.
<p className="text-xs text-muted-foreground">
Restore standard tiers to use Heuristics, Capability, or Fuse v2
</p>
)}
<TabsContent value={selected}>{children}</TabsContent>
</Tabs>
{availability.data?.error && !availability.isChecking && (
<div role="alert" className="space-y-1">
<p className="text-sm text-destructive">{availability.data.error}</p>
<AutoRouterContactLink />
</div>
)}
{availability.isError && (
<p className="text-xs text-muted-foreground">
Could not check availability.{" "}
<button type="button" className="underline" onClick={() => availability.refetch?.()}>
Retry
</button>
</p>
)}
{children}
</div>
);
};

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