mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): preserve current bridge auth in configuration scope
This commit is contained in:
commit
b0f1eb656c
139 changed files with 9189 additions and 1791 deletions
1
.github/e2e-stack/select_tests.py
vendored
1
.github/e2e-stack/select_tests.py
vendored
|
|
@ -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)$"
|
||||
|
|
|
|||
32
.github/workflows/_test-unit-base.yml
vendored
32
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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[@]}" \
|
||||
|
|
|
|||
8
.github/workflows/test-unit.yml
vendored
8
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
148
litellm/proxy/management_helpers/auto_router_availability.py
Normal file
148
litellm/proxy/management_helpers/auto_router_availability.py
Normal 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
|
||||
),
|
||||
)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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`).
|
||||
|
|
|
|||
145
tests/e2e/batches/bedrock_env_gateway.py
Normal file
145
tests/e2e/batches/bedrock_env_gateway.py
Normal 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)
|
||||
109
tests/e2e/batches/test_bedrock_blank_s3_env_e2e.py
Normal file
109
tests/e2e/batches/test_bedrock_blank_s3_env_e2e.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -65,6 +65,7 @@ LlmCapability = Literal[
|
|||
"assume_role",
|
||||
"basic",
|
||||
"batch_deployment",
|
||||
"blank_s3_env",
|
||||
"count_tokens",
|
||||
"govcloud_partition",
|
||||
"split_s3_credentials",
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
8
tests/e2e/gateway/secret_manager_cyberark_ci_config.yml
Normal file
8
tests/e2e/gateway/secret_manager_cyberark_ci_config.yml
Normal 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/
|
||||
|
|
@ -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/
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
76
tests/e2e/secret_manager/backend.sh
Executable file
76
tests/e2e/secret_manager/backend.sh
Executable 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
|
||||
57
tests/e2e/secret_manager/conftest.py
Normal file
57
tests/e2e/secret_manager/conftest.py
Normal 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()
|
||||
25
tests/e2e/secret_manager/secret_backends.py
Normal file
25
tests/e2e/secret_manager/secret_backends.py
Normal 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
|
||||
29
tests/e2e/secret_manager/secret_store.py
Normal file
29
tests/e2e/secret_manager/secret_store.py
Normal 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"
|
||||
118
tests/e2e/secret_manager/secret_store_cyberark.py
Normal file
118
tests/e2e/secret_manager/secret_store_cyberark.py
Normal 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())
|
||||
113
tests/e2e/secret_manager/secret_store_hashicorp_vault.py
Normal file
113
tests/e2e/secret_manager/secret_store_hashicorp_vault.py
Normal 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"}),
|
||||
)
|
||||
128
tests/e2e/secret_manager/test_secret_manager_e2e.py
Normal file
128
tests/e2e/secret_manager/test_secret_manager_e2e.py
Normal 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}")
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
189
tests/integration/spend/test_shutdown_flush.py
Normal file
189
tests/integration/spend/test_shutdown_flush.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
},
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"]},
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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") == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
74
tests/test_litellm/test_unit_shard_missing_paths.py
Normal file
74
tests/test_litellm/test_unit_shard_missing_paths.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -2120,6 +2120,7 @@ class TestPrismaTableRepository:
|
|||
"litellm_prompttable",
|
||||
"litellm_searchtoolstable",
|
||||
"litellm_ssoconfig",
|
||||
"litellm_uisettings",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
19
ui/litellm-dashboard/package-lock.json
generated
19
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue