From 13691426c4610f10da9627388a53419196a989ee Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 16:52:05 -0700 Subject: [PATCH 01/16] test(cost_calculator): point image-generation deployment price test at a live gemini row (#42615) The test priced gemini/gemini-3.1-flash-image-preview, which #42435 removed from the cost map as deprecated, so the calculator had no per-token rates to keep and the hardcoded expected value no longer matched. Price the live gemini/gemini-3.1-flash-image row instead and derive the expected cost from that row in litellm.model_cost, so a rate change on it cannot break the test while a calculator that drops the map's token rates still fails it. Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- tests/test_litellm/test_cost_calculator.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index e7ce9a74797..c6b57fa7604 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -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( From cf08cb89e81d4e744644428a09241da08b2996ce Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:04:13 -0700 Subject: [PATCH 02/16] test(utils): accept the per-size image cost keys in the price-map schema check (#42612) The cost map's fal_ai/fal-ai/trellis-2 entry prices its output by resolution with output_cost_per_image_512, output_cost_per_image_1024, and output_cost_per_image_1536, which litellm/types/utils.py types and the fal_ai cost calculator reads, but INTENDED_SCHEMA in test_aaamodel_prices_and_context_window_json_is_valid never allowed them, so the test fails on main with "Additional properties are not allowed". Add the three keys next to output_cost_per_image in the schema and in the cost-under-1 field list so a per-size image price is validated like the per-size video ones Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- tests/test_litellm/test_utils.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 79462d16a8c..283cff97ce0 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -631,6 +631,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 +861,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"}, From 9bad2c35e51f34abd0557877fe03f565f9657067 Mon Sep 17 00:00:00 2001 From: Pawan Shahane <110886433+Pawan-Shahane@users.noreply.github.com> Date: Wed, 23 Sep 2026 05:39:39 +0530 Subject: [PATCH 03/16] fix(ollama): send PNG and JPEG images without requiring Pillow. (#41979) * fix(ollama): send PNG and JPEG images without requiring Pillow The ollama/ completion transport imported Pillow before it looked at the image, so every image request failed with a 500 on installs without Pillow. That includes the Docker image, where Pillow is only a CI dependency Detect PNG and JPEG from their leading bytes and pass them through untouched. Pillow is now imported only when another format has to be re-encoded as JPEG, and that case still raises the same install hint * fix(ollama): address Greptile findings on image conversion Catch all exceptions on Pillow import, not just ImportError, so the helpful install hint always appears. Break a line that exceeded 120 characters --- litellm/llms/ollama/common_utils.py | 36 ++++----- .../test_ollama_completion_transformation.py | 74 +++++++++++++++++++ 2 files changed, 92 insertions(+), 18 deletions(-) diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index ed4bab22a84..9f46cbc5cd5 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -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 diff --git a/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py index b2071155f3f..28e86e40944 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py @@ -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] From ca95fc2bd4185483e91cdd62093b6fdb339f82f3 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:23:44 -0700 Subject: [PATCH 04/16] fix: answer get_api_base for github_copilot and chatgpt without running the login flow (#42602) * fix: answer get_api_base for github_copilot and chatgpt without running the login flow * refactor(get_api_base): dispatch the provider helpers with if-chains --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../llm_response_utils/get_api_base.py | 39 +++++--- litellm/llms/chatgpt/chat/transformation.py | 5 +- .../github_copilot/chat/transformation.py | 15 +-- .../llm_response_utils/test_get_api_base.py | 93 +++++++++++++++++++ tests/test_litellm/rerank_api/test_main.py | 4 +- 5 files changed, 134 insertions(+), 22 deletions(-) create mode 100644 tests/test_litellm/litellm_core_utils/llm_response_utils/test_get_api_base.py diff --git a/litellm/litellm_core_utils/llm_response_utils/get_api_base.py b/litellm/litellm_core_utils/llm_response_utils/get_api_base.py index 26e79fa0ea8..3815ea91b51 100644 --- a/litellm/litellm_core_utils/llm_response_utils/get_api_base.py +++ b/litellm/litellm_core_utils/llm_response_utils/get_api_base.py @@ -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) diff --git a/litellm/llms/chatgpt/chat/transformation.py b/litellm/llms/chatgpt/chat/transformation.py index e35408b0829..1b110704c8b 100644 --- a/litellm/llms/chatgpt/chat/transformation.py +++ b/litellm/llms/chatgpt/chat/transformation.py @@ -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: diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 8634b374f1b..169b9a037a5 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -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: diff --git a/tests/test_litellm/litellm_core_utils/llm_response_utils/test_get_api_base.py b/tests/test_litellm/litellm_core_utils/llm_response_utils/test_get_api_base.py new file mode 100644 index 00000000000..63977c30270 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/llm_response_utils/test_get_api_base.py @@ -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 diff --git a/tests/test_litellm/rerank_api/test_main.py b/tests/test_litellm/rerank_api/test_main.py index aca0c970dd5..f56673fec30 100644 --- a/tests/test_litellm/rerank_api/test_main.py +++ b/tests/test_litellm/rerank_api/test_main.py @@ -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: From 4f93e2c3da75289393af1b9e2ddb28935b4e11da Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 22 Sep 2026 17:28:34 -0700 Subject: [PATCH 05/16] test: point CircleCI-only suites at models still in the cost map (#42617) * test: point CircleCI-only suites at models still in the cost map #42435 removed cost map entries past their deprecation date and #42437 added litellm_uisettings to the config-synced tables, but both only updated tests/test_litellm. The CircleCI-only suites (local_testing, llm_translation, logging_callback_tests, litellm_utils_tests, unit) kept using the removed models or the old table list and went red on main. Each test keeps its assertions and swaps the removed model for a current one with the same provider and capabilities. The fireworks tests pick a vision model from the cost map because #34941 set supports_vision false on minimax-m3, and the vertex image provider test injects the image model set because #42435 removed every vertex_ai-image-models entry. * test(vertex_ai): register the image model through add_known_models in the provider test --- tests/litellm_utils_tests/test_utils.py | 2 +- .../test_fireworks_ai_translation.py | 14 ++++-- .../test_gemini_image_usage.py | 2 +- tests/llm_translation/test_groq.py | 2 +- tests/llm_translation/test_optional_params.py | 8 +-- tests/llm_translation/test_xai.py | 2 +- .../test_amazing_vertex_completion.py | 6 +-- tests/local_testing/test_completion_cost.py | 50 +++++++++---------- tests/local_testing/test_exceptions.py | 2 +- .../test_function_call_parsing.py | 2 +- tests/local_testing/test_get_llm_provider.py | 18 +++++-- tests/local_testing/test_get_model_info.py | 4 +- .../local_testing/test_lowest_cost_routing.py | 2 +- .../test_openai_moderations_hook.py | 6 +-- tests/local_testing/test_router_utils.py | 24 ++++----- .../test_spend_calculate_endpoint.py | 4 +- .../completion_with_vertex_call.json | 10 ++-- tests/logging_callback_tests/test_alerting.py | 6 +-- .../test_langfuse_e2e_test.py | 4 +- tests/unit/repositories/test_repositories.py | 1 + 20 files changed, 93 insertions(+), 76 deletions(-) diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index f7575b969c4..fb20cdf7e0e 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -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"}, diff --git a/tests/llm_translation/test_fireworks_ai_translation.py b/tests/llm_translation/test_fireworks_ai_translation.py index e20134fc1bf..a7dc913c388 100644 --- a/tests/llm_translation/test_fireworks_ai_translation.py +++ b/tests/llm_translation/test_fireworks_ai_translation.py @@ -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", diff --git a/tests/llm_translation/test_gemini_image_usage.py b/tests/llm_translation/test_gemini_image_usage.py index 096f9c4796c..0be8b6c23e1 100644 --- a/tests/llm_translation/test_gemini_image_usage.py +++ b/tests/llm_translation/test_gemini_image_usage.py @@ -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 = { diff --git a/tests/llm_translation/test_groq.py b/tests/llm_translation/test_groq.py index c720f818eaf..fbecbeab08b 100644 --- a/tests/llm_translation/test_groq.py +++ b/tests/llm_translation/test_groq.py @@ -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""" diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py index 997f5b3b73f..58446014bdf 100644 --- a/tests/llm_translation/test_optional_params.py +++ b/tests/llm_translation/test_optional_params.py @@ -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"], diff --git a/tests/llm_translation/test_xai.py b/tests/llm_translation/test_xai.py index 7a121afc3fa..d6d42ed215e 100644 --- a/tests/llm_translation/test_xai.py +++ b/tests/llm_translation/test_xai.py @@ -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"}], } diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py index 3d66064f5c0..8b45ae08813 100644 --- a/tests/local_testing/test_amazing_vertex_completion.py +++ b/tests/local_testing/test_amazing_vertex_completion.py @@ -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 diff --git a/tests/local_testing/test_completion_cost.py b/tests/local_testing/test_completion_cost.py index f40818b9bf1..3ce99f893d8 100644 --- a/tests/local_testing/test_completion_cost.py +++ b/tests/local_testing/test_completion_cost.py @@ -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]) diff --git a/tests/local_testing/test_exceptions.py b/tests/local_testing/test_exceptions.py index e6392cda406..813146f8ace 100644 --- a/tests/local_testing/test_exceptions.py +++ b/tests/local_testing/test_exceptions.py @@ -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 diff --git a/tests/local_testing/test_function_call_parsing.py b/tests/local_testing/test_function_call_parsing.py index c98f170a98f..ebb13e0018d 100644 --- a/tests/local_testing/test_function_call_parsing.py +++ b/tests/local_testing/test_function_call_parsing.py @@ -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): diff --git a/tests/local_testing/test_get_llm_provider.py b/tests/local_testing/test_get_llm_provider.py index ebad0fbafc5..4ac7cecb97a 100644 --- a/tests/local_testing/test_get_llm_provider.py +++ b/tests/local_testing/test_get_llm_provider.py @@ -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(): diff --git a/tests/local_testing/test_get_model_info.py b/tests/local_testing/test_get_model_info.py index 37f4ece611d..1e46a1bf853 100644 --- a/tests/local_testing/test_get_model_info.py +++ b/tests/local_testing/test_get_model_info.py @@ -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) diff --git a/tests/local_testing/test_lowest_cost_routing.py b/tests/local_testing/test_lowest_cost_routing.py index 5bf3a3ee98b..631271ca710 100644 --- a/tests/local_testing/test_lowest_cost_routing.py +++ b/tests/local_testing/test_lowest_cost_routing.py @@ -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"}, }, ] diff --git a/tests/local_testing/test_openai_moderations_hook.py b/tests/local_testing/test_openai_moderations_hook.py index 530ab714eae..7ce4bc2e4bf 100644 --- a/tests/local_testing/test_openai_moderations_hook.py +++ b/tests/local_testing/test_openai_moderations_hook.py @@ -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"), }, } diff --git a/tests/local_testing/test_router_utils.py b/tests/local_testing/test_router_utils.py index 1b3e361bb1f..635bda55144 100644 --- a/tests/local_testing/test_router_utils.py +++ b/tests/local_testing/test_router_utils.py @@ -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 diff --git a/tests/local_testing/test_spend_calculate_endpoint.py b/tests/local_testing/test_spend_calculate_endpoint.py index 3bedab794e2..054dc398039 100644 --- a/tests/local_testing/test_spend_calculate_endpoint.py +++ b/tests/local_testing/test_spend_calculate_endpoint.py @@ -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": { diff --git a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json index b6c11f96953..5998c52659c 100644 --- a/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json +++ b/tests/logging_callback_tests/langfuse_expected_request_body/completion_with_vertex_call.json @@ -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, diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py index 3074e973a8e..0a3e1a0e982 100644 --- a/tests/logging_callback_tests/test_alerting.py +++ b/tests/logging_callback_tests/test_alerting.py @@ -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() diff --git a/tests/logging_callback_tests/test_langfuse_e2e_test.py b/tests/logging_callback_tests/test_langfuse_e2e_test.py index 5682d3720d8..76ebd2b9a28 100644 --- a/tests/logging_callback_tests/test_langfuse_e2e_test.py +++ b/tests/logging_callback_tests/test_langfuse_e2e_test.py @@ -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"]}, diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index 87cf2fc4268..e185d95ffb8 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -2120,6 +2120,7 @@ class TestPrismaTableRepository: "litellm_prompttable", "litellm_searchtoolstable", "litellm_ssoconfig", + "litellm_uisettings", } ) From 7688f56256484f866687e147d0a994fb7649dd6d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:28:48 -0700 Subject: [PATCH 06/16] fix(pricing): drop the duplicate cache_read_input_token_cost_batches key from 23 entries (#42623) Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 69 +++++++------------ model_prices_and_context_window.json | 69 +++++++------------ 2 files changed, 46 insertions(+), 92 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 224129fbd25..ad66cb76f8b 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 224129fbd25..ad66cb76f8b 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, From b4ccb5b7474270bd69f1b491a3febf0bd2c44210 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Tue, 22 Sep 2026 20:33:44 -0400 Subject: [PATCH 07/16] fix(s3): replace colons in generated log filenames (#40452) * fix(s3): replace colons in generated log filenames Bedrock and Vertex AI batch file uploads use s3:// and gs:// URIs as response ids. The shared filename sanitizer replaced slashes but kept the scheme colon, producing log object keys that Hadoop-style consumers reject as a relative path in an absolute URI. Fixes #40234 * test(s3): drop docstrings flagged by review --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/integrations/s3.py | 2 +- tests/test_litellm/integrations/test_s3_v2.py | 28 ++++++++++++++++--- 2 files changed, 25 insertions(+), 5 deletions(-) diff --git a/litellm/integrations/s3.py b/litellm/integrations/s3.py index 796784fb993..54ec876fd0d 100644 --- a/litellm/integrations/s3.py +++ b/litellm/integrations/s3.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 diff --git a/tests/test_litellm/integrations/test_s3_v2.py b/tests/test_litellm/integrations/test_s3_v2.py index 52fbbe40b0e..a9d13038180 100644 --- a/tests/test_litellm/integrations/test_s3_v2.py +++ b/tests/test_litellm/integrations/test_s3_v2.py @@ -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.""" From 38f0eb876bf2f05b4850adb8e80bcb48885e557e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:42:12 -0700 Subject: [PATCH 08/16] test(realtime): drop legacy InvalidStatusCode tests and pin websockets imports (#42624) The two redaction tests raised the deprecated InvalidStatusCode, which the websockets 15 asyncio client never raises, and asserted the raw 403 close code that the handshake refusal path replaced with 1008. The refusal path builds its close reason from the status code alone, so there is no secret to redact there, and the handshake refusal tests already cover the error event and the 1008 close. Those refusal tests only passed when run after a sibling test had imported websockets.asyncio.client, since websockets lazy-loads its exceptions submodule. Importing InvalidStatus, Response, and Headers from their own submodules makes them pass in any order. Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../llms/azure/realtime/test_handler.py | 8 +-- .../realtime/test_openai_realtime_handler.py | 8 +-- .../test_redact_string_in_error_paths.py | 56 +------------------ 3 files changed, 9 insertions(+), 63 deletions(-) diff --git a/tests/test_litellm/llms/azure/realtime/test_handler.py b/tests/test_litellm/llms/azure/realtime/test_handler.py index edf1b8b290f..e9d24b459d8 100644 --- a/tests/test_litellm/llms/azure/realtime/test_handler.py +++ b/tests/test_litellm/llms/azure/realtime/test_handler.py @@ -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 diff --git a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py index 7cd2b9e259c..f7a88b5ba63 100644 --- a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py @@ -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 diff --git a/tests/test_litellm/test_redact_string_in_error_paths.py b/tests/test_litellm/test_redact_string_in_error_paths.py index 07d1ec5f523..a5128a87b0d 100644 --- a/tests/test_litellm/test_redact_string_in_error_paths.py +++ b/tests/test_litellm/test_redact_string_in_error_paths.py @@ -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.""" From 8ee6bab52931798d9aec00003797c25d0190e144 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:55:22 -0700 Subject: [PATCH 09/16] fix(bedrock): treat blank AWS_S3_* env vars as unset for batch jobs (#42528) * test(e2e): pin bedrock batch create with blank S3 env vars Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(bedrock): treat blank S3 env vars as unset for batch jobs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): trim blank S3 env gateway config Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): register blank_s3_env capability and clean gateway tempdir Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): move blank S3 env batch test to its own module Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/bedrock/common_utils.py | 2 +- tests/e2e/batches/COVERAGE.md | 1 + tests/e2e/batches/bedrock_env_gateway.py | 145 ++++++++++++++++++ .../batches/test_bedrock_blank_s3_env_e2e.py | 109 +++++++++++++ .../llm_nonconversational.yaml | 1 + tests/e2e/coverage_registry/schema.py | 1 + .../bedrock/batches/test_transformation.py | 41 ++++- 7 files changed, 297 insertions(+), 3 deletions(-) create mode 100644 tests/e2e/batches/bedrock_env_gateway.py create mode 100644 tests/e2e/batches/test_bedrock_blank_s3_env_e2e.py diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index d9fc813a594..2e20aafffcb 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -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: diff --git a/tests/e2e/batches/COVERAGE.md b/tests/e2e/batches/COVERAGE.md index 1731b4c620d..862eef5c0f4 100644 --- a/tests/e2e/batches/COVERAGE.md +++ b/tests/e2e/batches/COVERAGE.md @@ -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`). diff --git a/tests/e2e/batches/bedrock_env_gateway.py b/tests/e2e/batches/bedrock_env_gateway.py new file mode 100644 index 00000000000..fb3ec60c87c --- /dev/null +++ b/tests/e2e/batches/bedrock_env_gateway.py @@ -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 "" + 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) diff --git a/tests/e2e/batches/test_bedrock_blank_s3_env_e2e.py b/tests/e2e/batches/test_bedrock_blank_s3_env_e2e.py new file mode 100644 index 00000000000..77eb8427e59 --- /dev/null +++ b/tests/e2e/batches/test_bedrock_blank_s3_env_e2e.py @@ -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) diff --git a/tests/e2e/coverage_registry/llm_nonconversational.yaml b/tests/e2e/coverage_registry/llm_nonconversational.yaml index c58c8af44ff..3e389acc2a9 100644 --- a/tests/e2e/coverage_registry/llm_nonconversational.yaml +++ b/tests/e2e/coverage_registry/llm_nonconversational.yaml @@ -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"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index f3ac1ef8a83..e009b02b69c 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -65,6 +65,7 @@ LlmCapability = Literal[ "assume_role", "basic", "batch_deployment", + "blank_s3_env", "count_tokens", "govcloud_partition", "split_s3_credentials", diff --git a/tests/test_litellm/llms/bedrock/batches/test_transformation.py b/tests/test_litellm/llms/bedrock/batches/test_transformation.py index eb08c19cbdf..347c459a369 100644 --- a/tests/test_litellm/llms/bedrock/batches/test_transformation.py +++ b/tests/test_litellm/llms/bedrock/batches/test_transformation.py @@ -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( From 21a2d828dffd70b0c868d2dba487d3c372239475 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 00:56:32 +0000 Subject: [PATCH 10/16] ci(test-unit): drop dead misc shard paths and skip missing paths with a warning (#42603) Eight directories the misc shard named moved to tests/unit on 2026-09-20, and one missing path makes pytest-xdist collect [0 items] for the whole shard, which the exit-5 tolerance turned into a green required check running nothing. The shared Run tests step now drops a path that does not exist with a ::warning:: and runs pytest over the rest, keeping option tokens verbatim. Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .github/workflows/_test-unit-base.yml | 32 +++++--- .github/workflows/test-unit.yml | 8 -- .../test_unit_shard_missing_paths.py | 74 +++++++++++++++++++ 3 files changed, 97 insertions(+), 17 deletions(-) create mode 100644 tests/test_litellm/test_unit_shard_missing_paths.py diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index f4d4fb54ba6..8faddd11df8 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -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[@]}" \ diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 49e6d7040d4..ac60a9b5f05 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -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 diff --git a/tests/test_litellm/test_unit_shard_missing_paths.py b/tests/test_litellm/test_unit_shard_missing_paths.py new file mode 100644 index 00000000000..b91c2cff764 --- /dev/null +++ b/tests/test_litellm/test_unit_shard_missing_paths.py @@ -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 From a80379baf8910f5aa72990cd1bdbd8088d5e6b2c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:58:08 -0700 Subject: [PATCH 11/16] fix(proxy): keep the in-flight daily spend batch when shutdown cancels the flush (#42593) * fix(proxy): keep the in-flight daily spend batch when shutdown cancels the flush A daily spend batch drained from the in-memory queue was dropped for good when the scheduler tick was cancelled by shutdown, because asyncio.CancelledError bypasses the except Exception requeue. The flush now requeues the drained rows on cancellation and re-raises, and each daily batch upsert runs in an interactive transaction so a statement that already reached Postgres is rolled back with the cancel instead of committing behind the requeue Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): requeue the cancelled daily spend batch before its rollback returns Behind a lock the rollback of the cancelled interactive transaction only returns once the blocked statement does, which is after the shutdown flush has already run. The commit now runs as a shielded task so the cancelled tick requeues the batch at once and lets the rollback finish in the background. The final flush then finds the rows and writes them exactly once Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): give the recording db a transaction seam for the bulk upsert tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): route the mocked daily tag spend upsert through the transaction seam Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): restore the drained Redis tag batch when shutdown cancels its commit Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 25 +- tests/integration/_support/process.py | 37 ++- tests/integration/contracts.json | 6 + .../integration/spend/test_shutdown_flush.py | 189 ++++++++++++ .../test_update_daily_tag_spend.py | 1 + .../proxy/db/test_daily_spend_bulk_upsert.py | 9 + .../proxy/db/test_db_spend_update_writer.py | 270 +++++++++++++++++- 7 files changed, 520 insertions(+), 17 deletions(-) create mode 100644 tests/integration/spend/test_shutdown_flush.py diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index d9c8b271646..e38214c98a6 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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( diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index e0923c9d055..0798927c6d1 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -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) diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json index 8c4ce4debbc..cfe31e2d8ec 100644 --- a/tests/integration/contracts.json +++ b/tests/integration/contracts.json @@ -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" ], diff --git a/tests/integration/spend/test_shutdown_flush.py b/tests/integration/spend/test_shutdown_flush.py new file mode 100644 index 00000000000..d27275f5213 --- /dev/null +++ b/tests/integration/spend/test_shutdown_flush.py @@ -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, + ) diff --git a/tests/proxy_unit_tests/test_update_daily_tag_spend.py b/tests/proxy_unit_tests/test_update_daily_tag_spend.py index 35e9c6796eb..530c0d16767 100644 --- a/tests/proxy_unit_tests/test_update_daily_tag_spend.py +++ b/tests/proxy_unit_tests/test_update_daily_tag_spend.py @@ -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": { diff --git a/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py b/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py index 510f77cecec..7893fb82281 100644 --- a/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py +++ b/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py @@ -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: diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index a997939ed34..daa07c8224a 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -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") == [] From 944f44d82be735616448f92b534dee1f063008b7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:59:50 -0700 Subject: [PATCH 12/16] fix(utils): isolate callback errors in async_post_call_success_deployment_hook (#42535) * fix(utils): isolate callback errors in async_post_call_success_deployment_hook A callback that raises inside async_post_call_success_deployment_hook no longer fails the completed request. The exception is logged with the callback class and call_type, the response stays as it was, and later callbacks still run. Guardrail callbacks are exempt because raising is how a post-call guardrail blocks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(utils): drop unrelated ruff autofixes from test_utils Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(utils): drop fastapi import from guardrail propagation regression Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(utils): cover every success deployment hook call type with a raising hook Parametrize the unit regression over video, embedding, responses, image, rerank, transcription, chat and anthropic messages responses and assert the failure log names the callback and call type. Run the integration test through a real proxy for /v1/chat/completions, /v1/embeddings, /v1/responses and /v1/videos Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): move raising success hook cases into the existing callback delivery file Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/utils.py | 17 +- tests/integration/contracts.json | 12 ++ .../observability/test_callback_delivery.py | 166 +++++++++++++++++- tests/test_litellm/test_utils.py | 113 ++++++++++++ 4 files changed, 303 insertions(+), 5 deletions(-) diff --git a/litellm/utils.py b/litellm/utils.py index 01f6fee3594..dd35c17809f 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json index cfe31e2d8ec..ee9b81f63dc 100644 --- a/tests/integration/contracts.json +++ b/tests/integration/contracts.json @@ -1782,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": { diff --git a/tests/integration/observability/test_callback_delivery.py b/tests/integration/observability/test_callback_delivery.py index c44c1f30b80..7a5f420d2cc 100644 --- a/tests/integration/observability/test_callback_delivery.py +++ b/tests/integration/observability/test_callback_delivery.py @@ -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 diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 283cff97ce0..89472cbd20f 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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, @@ -4412,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( From 630c4624f6027c4c075aa91439702878690ece78 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Tue, 22 Sep 2026 18:02:32 -0700 Subject: [PATCH 13/16] test(e2e): add secret manager lanes for HashiCorp Vault and CyberArk Conjur (#42503) * test(e2e): add a HashiCorp Vault secret manager lane key_management_system had no end-to-end coverage: the Rust crates and the Python unit tests all run against mocked managers. This adds a secret_manager suite that drives a proxy configured with hashicorp_vault against a real Vault. The tests seed a fresh secret name per test with the runner's OPENAI_API_KEY and register a deployment pointing at os.environ/. The proxy's env never holds that name, so get_secret's os.environ fallback cannot mask a broken manager, and a bogus value in Vault must come back as the provider's 401. Virtual keys are checked written to and removed from Vault under prefix_for_stored_virtual_keys. The setting is global to the proxy, so the lane has its own config and the secret_manager_vault opt-in marker, and stays out of the per-PR selector. Co-Authored-By: Claude Opus 5 * test(e2e): make the secret manager suite backend-agnostic One marker and opt-in (secret_manager / E2E_SECRET_MANAGER=) pick the backend from secret_backends.BACKENDS. The tests reach the manager through a SecretStore protocol, and each backend contributes a secret_store_.py module, a registry entry, and gateway/secret_manager__ci_config.yml. requires_capability deselects tests a backend cannot support (CyberArk does not delete), and test_secret_backends.py checks every lane config against its backend without a live stack. Co-Authored-By: Claude Opus 5 * test(e2e): add a CyberArk Conjur secret manager lane Adds cyberark as the second secret_manager backend: a Conjur store over its REST API (policy-declared variables, raw-text values, policy-patch teardown), its lane config, and a registry entry without deletes_stored_keys, since the proxy's CyberArk delete answers not_supported and Conjur keeps the key. secret_manager/backend.sh up|down boots any backend in Docker and writes proxy.env and tests.env, so every lane runs the same way; the registry test checks the script boots exactly the registered backends. e2e_http gains send_text_external for APIs that speak raw text rather than JSON. Co-Authored-By: Claude Opus 5 * fix(e2e): give the secret manager suite a client with .proxy and address review The shared resources fixture reads client.proxy, so a bare ProxyClient errored every live test at setup. backend.sh now writes its env under a per-user directory with umask 077, the markerless unit tests are gone per tests/e2e/AGENTS.md, and routine comments are trimmed. Co-Authored-By: Claude Opus 5.5 --------- Co-authored-by: Claude Opus 5 --- .github/e2e-stack/select_tests.py | 1 + tests/e2e/AGENTS.md | 1 + tests/e2e/CONTRIBUTING.md | 27 +++- tests/e2e/conftest.py | 7 + tests/e2e/coverage_registry/other.yaml | 5 +- tests/e2e/e2e_config.py | 1 + tests/e2e/e2e_http.py | 30 +++- .../secret_manager_cyberark_ci_config.yml | 8 ++ ...cret_manager_hashicorp_vault_ci_config.yml | 8 ++ tests/e2e/pytest.ini | 1 + tests/e2e/secret_manager/backend.sh | 76 +++++++++++ tests/e2e/secret_manager/conftest.py | 57 ++++++++ tests/e2e/secret_manager/secret_backends.py | 25 ++++ tests/e2e/secret_manager/secret_store.py | 29 ++++ .../secret_manager/secret_store_cyberark.py | 118 ++++++++++++++++ .../secret_store_hashicorp_vault.py | 113 ++++++++++++++++ .../secret_manager/test_secret_manager_e2e.py | 128 ++++++++++++++++++ 17 files changed, 630 insertions(+), 5 deletions(-) create mode 100644 tests/e2e/gateway/secret_manager_cyberark_ci_config.yml create mode 100644 tests/e2e/gateway/secret_manager_hashicorp_vault_ci_config.yml create mode 100755 tests/e2e/secret_manager/backend.sh create mode 100644 tests/e2e/secret_manager/conftest.py create mode 100644 tests/e2e/secret_manager/secret_backends.py create mode 100644 tests/e2e/secret_manager/secret_store.py create mode 100644 tests/e2e/secret_manager/secret_store_cyberark.py create mode 100644 tests/e2e/secret_manager/secret_store_hashicorp_vault.py create mode 100644 tests/e2e/secret_manager/test_secret_manager_e2e.py diff --git a/.github/e2e-stack/select_tests.py b/.github/e2e-stack/select_tests.py index 492c52233cd..e425c313d6a 100644 --- a/.github/e2e-stack/select_tests.py +++ b/.github/e2e-stack/select_tests.py @@ -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)$" diff --git a/tests/e2e/AGENTS.md b/tests/e2e/AGENTS.md index d948c6fd1d9..94035cfe849 100644 --- a/tests/e2e/AGENTS.md +++ b/tests/e2e/AGENTS.md @@ -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/` 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=` opts in and picks the backend from `secret_backends.BACKENDS`, the proxy is booted from `gateway/secret_manager__ci_config.yml` against the live manager, and the tests reach that manager through the backend's `SecretStore` (`secret_store_.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 ` 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` diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index 15c2d6763d6..7e1f516422e 100644 --- a/tests/e2e/CONTRIBUTING.md +++ b/tests/e2e/CONTRIBUTING.md @@ -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__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_.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__*` 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__ci_config.yml`, a copy of an existing lane's with only `key_management_system` changed +4. an `up_` 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 ` (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 diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index f4de4ac01ba..603591006d9 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -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__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: diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml index 1706b2f8f25..f695af3cb11 100644 --- a/tests/e2e/coverage_registry/other.yaml +++ b/tests/e2e/coverage_registry/other.yaml @@ -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/ 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"} diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index ac26b2a3875..ca3b74281ae 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -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")) diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index 1b62a1dbd8c..24bd76d4429 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -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) diff --git a/tests/e2e/gateway/secret_manager_cyberark_ci_config.yml b/tests/e2e/gateway/secret_manager_cyberark_ci_config.yml new file mode 100644 index 00000000000..32602783b02 --- /dev/null +++ b/tests/e2e/gateway/secret_manager_cyberark_ci_config.yml @@ -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/ diff --git a/tests/e2e/gateway/secret_manager_hashicorp_vault_ci_config.yml b/tests/e2e/gateway/secret_manager_hashicorp_vault_ci_config.yml new file mode 100644 index 00000000000..8e6b7e74f8e --- /dev/null +++ b/tests/e2e/gateway/secret_manager_hashicorp_vault_ci_config.yml @@ -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/ diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini index 4866511af3f..d01caeff3ea 100644 --- a/tests/e2e/pytest.ini +++ b/tests/e2e/pytest.ini @@ -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__ci_config.yml against that live secret manager; deselected unless E2E_SECRET_MANAGER names the backend (see secret_manager/secret_backends.py) diff --git a/tests/e2e/secret_manager/backend.sh b/tests/e2e/secret_manager/backend.sh new file mode 100755 index 00000000000..70272c237c5 --- /dev/null +++ b/tests/e2e/secret_manager/backend.sh @@ -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 diff --git a/tests/e2e/secret_manager/conftest.py b/tests/e2e/secret_manager/conftest.py new file mode 100644 index 00000000000..46ef598d711 --- /dev/null +++ b/tests/e2e/secret_manager/conftest.py @@ -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() diff --git a/tests/e2e/secret_manager/secret_backends.py b/tests/e2e/secret_manager/secret_backends.py new file mode 100644 index 00000000000..578b56970a3 --- /dev/null +++ b/tests/e2e/secret_manager/secret_backends.py @@ -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 diff --git a/tests/e2e/secret_manager/secret_store.py b/tests/e2e/secret_manager/secret_store.py new file mode 100644 index 00000000000..cff5005b73a --- /dev/null +++ b/tests/e2e/secret_manager/secret_store.py @@ -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" diff --git a/tests/e2e/secret_manager/secret_store_cyberark.py b/tests/e2e/secret_manager/secret_store_cyberark.py new file mode 100644 index 00000000000..87bcd2ffb1c --- /dev/null +++ b/tests/e2e/secret_manager/secret_store_cyberark.py @@ -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()) diff --git a/tests/e2e/secret_manager/secret_store_hashicorp_vault.py b/tests/e2e/secret_manager/secret_store_hashicorp_vault.py new file mode 100644 index 00000000000..ccf8cefe716 --- /dev/null +++ b/tests/e2e/secret_manager/secret_store_hashicorp_vault.py @@ -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"}), +) diff --git a/tests/e2e/secret_manager/test_secret_manager_e2e.py b/tests/e2e/secret_manager/test_secret_manager_e2e.py new file mode 100644 index 00000000000..a9c9024718d --- /dev/null +++ b/tests/e2e/secret_manager/test_secret_manager_e2e.py @@ -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}") From 2157351004d63e8a3db0f51c32bd8fe9c69ad6d5 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Tue, 22 Sep 2026 18:03:32 -0700 Subject: [PATCH 14/16] feat(ui): simplify auto-router setup and clarify feature limits (#42625) * feat(ui): simplify auto-router setup and clarify feature limits * fix(ui): validate auto-router drafts before saving * fix: keep auto-router allowances consistent after deletes and refreshes --- litellm/proxy/_types.py | 1 + .../auto_router_endpoints.py | 50 ++ .../model_management_endpoints.py | 13 +- .../auto_router_availability.py | 148 +++++ litellm/proxy/proxy_server.py | 9 + .../public_endpoints/autorouter_presets.json | 36 +- .../complexity_router/fuse_presets.json | 23 +- .../auto_router_tuning_baseline.py | 16 +- .../auto_router_endpoints.py | 19 + .../test_auto_router_endpoints.py | 153 ++++- .../test_model_management_endpoints.py | 145 ++++- .../test_auto_router_availability.py | 196 +++++++ .../proxy/proxy_server/test_lifecycle.py | 54 +- .../proxy/proxy_server/test_proxy_config.py | 41 +- .../public_endpoints/test_public_endpoints.py | 8 +- .../test_auto_router_tuning_baseline.py | 61 +- .../add_model/AutoRouterAvailability.tsx | 165 ++++++ ...oRouterClassifierTabs.integration.test.tsx | 255 ++++++-- .../add_model/AutoRouterClassifierTabs.tsx | 310 ++++++++-- .../add_model/ClassificationMethodConfig.tsx | 143 ++--- .../add_model/ClassifierPrimarySettings.tsx | 98 ++++ .../add_model/ClassifierTypeRadios.tsx | 2 +- .../ComplexityRouterAdvancedSections.tsx | 118 +++- ...mplexityRouterConfig.integration.test.tsx} | 461 ++++++++------- .../add_model/ComplexityRouterConfig.tsx | 36 +- ...plexityRouterFastMode.integration.test.tsx | 9 +- ...ecastClassifierConfig.integration.test.tsx | 23 +- .../add_model/ForecastClassifierConfig.tsx | 357 ++++++------ .../JevClassifierConfig.integration.test.tsx | 32 +- .../add_model/JevClassifierConfig.tsx | 45 +- .../JevConnectionTest.integration.test.tsx | 6 +- .../add_model/NonReasoningTierToggle.tsx | 2 +- .../components/add_model/RoutingOptions.tsx | 25 +- .../components/add_model/TierConfigIntro.tsx | 31 +- ... add_auto_router_tab.integration.test.tsx} | 442 ++++++++++---- .../add_model/add_auto_router_tab.tsx | 542 +++++++++--------- .../add_model/auto_router_connection_test.tsx | 10 +- ...dit_auto_router_modal.integration.test.tsx | 214 +++++-- .../edit_auto_router_modal.tsx | 267 +++++---- .../src/lib/autorouter_presets.test.ts | 30 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 87 +++ ui/litellm-dashboard/tests/autoRouterSetup.ts | 34 ++ 42 files changed, 3364 insertions(+), 1353 deletions(-) create mode 100644 litellm/proxy/management_helpers/auto_router_availability.py create mode 100644 tests/test_litellm/proxy/management_helpers/test_auto_router_availability.py create mode 100644 ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx create mode 100644 ui/litellm-dashboard/src/components/add_model/ClassifierPrimarySettings.tsx rename ui/litellm-dashboard/src/components/add_model/{ComplexityRouterConfig.test.tsx => ComplexityRouterConfig.integration.test.tsx} (85%) rename ui/litellm-dashboard/src/components/add_model/{add_auto_router_tab.test.tsx => add_auto_router_tab.integration.test.tsx} (77%) create mode 100644 ui/litellm-dashboard/tests/autoRouterSetup.ts diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 76df8790b92..c7273738fd0 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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", diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 03fc58622ce..2ae8639fe61 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -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, diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index fcadcfe2cae..615a528b552 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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. diff --git a/litellm/proxy/management_helpers/auto_router_availability.py b/litellm/proxy/management_helpers/auto_router_availability.py new file mode 100644 index 00000000000..52cd9f0499b --- /dev/null +++ b/litellm/proxy/management_helpers/auto_router_availability.py @@ -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 + ), + ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index acb36fc4a59..dab7decd4dc 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, @@ -5068,6 +5069,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 @@ -7691,6 +7693,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. @@ -7711,6 +7719,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( diff --git a/litellm/proxy/public_endpoints/autorouter_presets.json b/litellm/proxy/public_endpoints/autorouter_presets.json index 7a251afc076..1c2f96350c6 100644 --- a/litellm/proxy/public_endpoints/autorouter_presets.json +++ b/litellm/proxy/public_endpoints/autorouter_presets.json @@ -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": { diff --git a/litellm/router_strategy/complexity_router/fuse_presets.json b/litellm/router_strategy/complexity_router/fuse_presets.json index 4006366dc25..d3f9c48c0d4 100644 --- a/litellm/router_strategy/complexity_router/fuse_presets.json +++ b/litellm/router_strategy/complexity_router/fuse_presets.json @@ -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": [ diff --git a/litellm/router_utils/auto_router_tuning_baseline.py b/litellm/router_utils/auto_router_tuning_baseline.py index 9699ab886b9..e87548bf6de 100644 --- a/litellm/router_utils/auto_router_tuning_baseline.py +++ b/litellm/router_utils/auto_router_tuning_baseline.py @@ -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." ) diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index 93ea925bd9e..334ca0dfb08 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -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. diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index a282ae731dd..ff3d19e8637 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index bd252169131..8d618f4699b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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() diff --git a/tests/test_litellm/proxy/management_helpers/test_auto_router_availability.py b/tests/test_litellm/proxy/management_helpers/test_auto_router_availability.py new file mode 100644 index 00000000000..027930d9ebb --- /dev/null +++ b/tests/test_litellm/proxy/management_helpers/test_auto_router_availability.py @@ -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 diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 9954351fa2e..36d2e16d261 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -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() diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index b47cce43dcc..89cb8356289 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -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 diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index 0f1ff3b024d..0dec44af402 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -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) diff --git a/tests/test_litellm/router_utils/test_auto_router_tuning_baseline.py b/tests/test_litellm/router_utils/test_auto_router_tuning_baseline.py index 75c115cd3ab..9307668d66f 100644 --- a/tests/test_litellm/router_utils/test_auto_router_tuning_baseline.py +++ b/tests/test_litellm/router_utils/test_auto_router_tuning_baseline.py @@ -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 diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx new file mode 100644 index 00000000000..b1233ef03b8 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterAvailability.tsx @@ -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({ 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 = { + queryKey: ["autoRouterAvailability", accessToken, body.team_id, body.saved_model_id, debounced], + queryFn: ({ signal }) => + apiClient.post("/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 ( + + {message} + + Talk to our team + + + ); +}; + +export const AutoRouterAllowanceLabel = ({ feature }: { feature: string }) => { + const label = useAllowanceLabel(feature); + return label ? ( + {label} + ) : null; +}; + +export const AutoRouterAllowanceNote = ({ feature, label }: { feature: string; label: string }) => { + const availability = useAllowanceLabel(feature); + return availability ? ( +

+ {label}: {availability} +

+ ) : 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 ( + + + View limits + + + Routing and customization limits +

+ Rule-based, Complexity, and Jev are unlimited with built-in settings. Choose or change tier models freely. + Customization allowances are shared across this proxy. +

+
+ {limits.map(([key, label]) => ( +
+
{label}
+
+ {availabilityLabel(state, key) ?? "Unlimited"} +
+
+ ))} +
+

+ Custom tier definitions and written classifier instructions share one allowance. Built-in prompts and + display-name changes do not use it. +

+

+ 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. +

+ +
+
+ ); +}; diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx index ac6851349ea..f843a472d15 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.integration.test.tsx @@ -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>; +}) { const [value, setValue] = useState(initialValue); return ( - - {value.classifier_type} - + ({ + key, + limit, + remaining, + available: true, + used_by_this_router: key === ownedFeature, + }), + ), + error: null, + }, + ...availabilityState, + }} + > + + {value.classifier_type} + + ); } -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( - Existing classifier settings + Existing settings , ); - 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(
); - 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(); + 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(); + 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(); + 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(); + 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(); + 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(); + 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(); + 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( - + Existing settings + , + ); + 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( + - Custom tiers - , + />, ); - 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(); + 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(); + 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(); + 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( + , + ); + 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( + + + + + , + ); + 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(); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx index 98c0d4aab2f..92fc8d2a335 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterClassifierTabs.tsx @@ -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 ( +
+ + + + {label} + {feature ? ( + + ) : ( + unlimited && Unlimited + )} + + {description} + + + {fresh && exhausted && ( + } + 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 + + )} +
+ ); +} + +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 ( + + } + > + {selectedLabel} + {feature ? ( + + ) : ( + Unlimited + )} + + + + + {children} + + + + ); +} interface AutoRouterClassifierTabsProps { value: ComplexityRouterConfigValue; @@ -11,47 +122,174 @@ interface AutoRouterClassifierTabsProps { } const AutoRouterClassifierTabs: React.FC = ({ 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 = { + 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> = { capability: "Capability", llm_v2: "Fuse v2" }; + const approachDescription: Partial> = { + 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 ( - -

Classifier type

- - Complexity - - Capability - - - Fuse v2 - - +
+
+ + What classifies your requests? + + + + {[ + { 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) => ( + + ))} + +
+ {family === "custom" && ( +

This router uses a custom classifier plugin

+ )} + {family === "heuristics" && ( +
+ + { + if (next === "heuristic" || next === "heuristic_v2") changeType(next); + }} + > + + + +

+ {classifierType === "heuristic_v2" + ? "Use calibrated probabilities to match requests to a tier" + : "Match requests using scoring rules. Choose or change tier models freely"} +

+
+ )} + {(family === "llm" || family === "jev") && ( +
+ + { + if (next === "llm" || next === "capability" || next === "llm_v2") { + if (next === "llm" && !isForecastClassifier(classifierType)) return; + changeType(next); + } + }} + > + + {family === "llm" && ( + <> + + + + )} + +

+ {approachDescription[classifierType] ?? "Match task difficulty to a tier"} +

+
+ )} {hasCustomTiers && ( -

- Restore standard tiers to use Capability or Fuse v2. +

+ Restore standard tiers to use Heuristics, Capability, or Fuse v2

)} - {children} - + {availability.data?.error && !availability.isChecking && ( +
+

{availability.data.error}

+ +
+ )} + {availability.isError && ( +

+ Could not check availability.{" "} + +

+ )} + {children} +
); }; diff --git a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx index b7a0fd67443..ccd204aa521 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx @@ -1,9 +1,10 @@ +import ClassifierPrimarySettings from "./ClassifierPrimarySettings"; +import { AutoRouterAllowanceNote } from "./AutoRouterAvailability"; import { transitionClassifierType } from "./classifier_type_transition"; import JevClassifierConfig from "./JevClassifierConfig"; import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; import { MultiSelect } from "@/components/shared/MultiSelect"; -import { SearchSelect } from "@/components/shared/SearchSelect"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Card, CardContent } from "@/components/ui/card"; import { Button } from "@/components/ui/button"; @@ -25,13 +26,10 @@ import ClassifierTypeRadios from "./ClassifierTypeRadios"; import type { ReasoningEffort } from "./complexity_router_tiers"; import { useComplexityScorerDefaults } from "@/app/(dashboard)/hooks/autoRouter/useComplexityScorerDefaults"; import { - ClassificationFrequency, ClassifierFallback, ClassifierLLMConfig, ClassifierType, ComplexityRouterConfigValue, - classificationFrequency, - withClassificationFrequency, DEFAULT_CLASSIFIER_CONTEXT_BUDGET_CHARS, MIN_QUOTED_CONTEXT_TURN_CHARS, DEFAULT_CLASSIFIER_CONTEXT_WINDOW_SIZE, @@ -171,6 +169,7 @@ interface ClassificationMethodConfigProps { showValidationErrors?: boolean; /** The resolved default model - see resolveComplexityDefaultModel. Names and gates the radio. */ defaultModel?: string; + advancedOnly?: boolean; } export const InactiveHeuristicV2Threshold: React.FC> = ({ @@ -215,13 +214,11 @@ const ClassificationMethodConfig: React.FC = ({ onCustomTechnicalKeywordsChange, showValidationErrors = false, defaultModel, + advancedOnly = false, }) => { const [draft, setDraft] = React.useState<{ id: string; raw: string } | null>(null); const hasDefaultModel = Boolean(defaultModel); const classifierType = effectiveClassifierType(value); - const sessionFrequencyRestriction = restrictedBy(value, "sessionAffinity"); - const classifierModelMissing = - showValidationErrors && usesLlmClassifier(classifierType) && !value.classifier_llm_config?.model; const usesCustomPrompt = Boolean(value.classifier_llm_config?.system_prompt?.trim()); const contextBudget = value.classifier_context_budget_chars ?? DEFAULT_CLASSIFIER_CONTEXT_BUDGET_CHARS; const contextBudgetQuotesNothing = contextBudget > 0 && contextBudget < MIN_QUOTED_CONTEXT_TURN_CHARS; @@ -282,23 +279,6 @@ const ClassificationMethodConfig: React.FC = ({ onChange(nextValue); }; - const handleClassifierModelChange = (model: string | null) => { - if (model === null) return; - if (model === value.classifier_llm_config?.model) return; - const { reasoning_effort: _reasoningEffort, ...classifierLlmConfig } = value.classifier_llm_config ?? { - model: "", - timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS, - }; - onChange({ - ...value, - classifier_llm_config: { - ...classifierLlmConfig, - model, - timeout_ms: classifierLlmConfig.timeout_ms, - }, - }); - }; - const handleClassifierReasoningEffortChange = (reasoningEffort: ReasoningEffort | undefined) => { if (!value.classifier_llm_config) return; const { reasoning_effort: _reasoningEffort, ...classifierLlmConfig } = value.classifier_llm_config; @@ -350,10 +330,6 @@ const ClassificationMethodConfig: React.FC = ({ onChange({ ...value, classifier_fallback: fallback }); }; - const handleClassificationFrequencyChange = (frequency: ClassificationFrequency) => { - onChange(withClassificationFrequency(value, frequency)); - }; - const handleClassifierContextWindowSizeChange = (windowSize: number) => { onChange({ ...value, @@ -389,7 +365,50 @@ const ClassificationMethodConfig: React.FC = ({ return ( <> - + {!advancedOnly && ( + <> + + + + )} + {advancedOnly && ["llm", "heuristic_first", "hybrid"].includes(classifierType) && ( +
+ + +
+ )} {classifierType === "custom" && ( @@ -474,66 +493,9 @@ const ClassificationMethodConfig: React.FC = ({ )} -
- How often to classify - - handleClassificationFrequencyChange(frequency as ClassificationFrequency) - } - > -
- - - -
-
-

- Holding the tier keeps an agent on one model for a whole tool loop and cuts scoring cost. A turn the router - cannot match to a held decision, such as one with no session id or an expired one, is scored again -

-
- {classifierType === "jev" && } {usesLlmClassifier(classifierType) && (
-
- Classifier Model - - {classifierModelMissing && A classifier model is required} -
= ({
+ {!value.custom_tier_set && usesCustomPrompt ? ( = ({ /> Number of prior user turns sent to the classifier provider, excluding tool output and harness reminders. - LLM and JEV default to 3 turns; JEV sends them to the configured TypeSafe endpoint. Set to 0 to omit + LLM and Jev default to 3 turns; Jev sends them to the configured TypeSafe endpoint. Set to 0 to omit conversation history. The current message and selected system text are still sent. @@ -769,6 +735,9 @@ const ClassificationMethodConfig: React.FC = ({ )} + {["heuristic", "heuristic_first", "hybrid"].includes(classifierType) && ( + + )} diff --git a/ui/litellm-dashboard/src/components/add_model/ClassifierPrimarySettings.tsx b/ui/litellm-dashboard/src/components/add_model/ClassifierPrimarySettings.tsx new file mode 100644 index 00000000000..32cb4851580 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/ClassifierPrimarySettings.tsx @@ -0,0 +1,98 @@ +import React from "react"; +import { Label } from "@/components/ui/label"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { + classificationFrequency, + withClassificationFrequency, + effectiveClassifierType, + usesLlmClassifier, + DEFAULT_CLASSIFIER_TIMEOUT_MS, + type ComplexityRouterConfigValue, + type ClassificationFrequency, +} from "./ComplexityRouterConfig"; +import { restrictedBy } from "./TierRestrictions"; + +export default function ClassifierPrimarySettings({ + value, + onChange, + modelOptions, + showValidationErrors = false, +}: { + value: ComplexityRouterConfigValue; + onChange: (value: ComplexityRouterConfigValue) => void; + modelOptions: { value: string; label: string }[]; + showValidationErrors?: boolean; +}) { + const id = React.useId(); + const restriction = restrictedBy(value, "sessionAffinity"); + const frequency = classificationFrequency(value); + const frequencyDescription = { + every_request: "Choose a model again for every request", + user_turn: "Reclassify when the user sends a new message", + session: "Keep the same tier for the session. Requires a client session ID", + }[frequency]; + const usesJudge = usesLlmClassifier(effectiveClassifierType(value)); + const missingJudge = showValidationErrors && usesJudge && !value.classifier_llm_config?.model; + return ( +
+
+ + +

{restriction?.reason ?? frequencyDescription}

+
+ {usesJudge && ( +
+ + { + if (!model || model === value.classifier_llm_config?.model) return; + onChange({ + ...value, + classifier_llm_config: { + ...value.classifier_llm_config, + model, + timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS, + reasoning_effort: undefined, + }, + }); + }} + /> + {missingJudge && ( +

+ A judge model is required +

+ )} +
+ )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx index 1602e19069a..1fd6dfa6a20 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassifierTypeRadios.tsx @@ -54,7 +54,7 @@ const ClassifierTypeRadios: React.FC = ({ value, clas diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterAdvancedSections.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterAdvancedSections.tsx index 44adecb9e1d..a4e833152a9 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterAdvancedSections.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterAdvancedSections.tsx @@ -6,6 +6,7 @@ import type { ModelGroup } from "@/components/llm_calls/fetch_models"; import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig"; import ClassificationMethodConfig from "./ClassificationMethodConfig"; +import ForecastClassifierConfig from "./ForecastClassifierConfig"; import ContextWindowEscalationConfig from "./ContextWindowEscalationConfig"; import ResponseFormatControls from "./ResponseFormatControls"; import StallEscalationConfig from "./StallEscalationConfig"; @@ -79,13 +80,31 @@ const ComplexityRouterAdvancedSections: React.FC { const sections = [ + ...(forecast + ? [ + { + key: "classifier", + label: Classifier tuning, + children: ( + + ), + }, + ] + : []), ...(!forecast ? [ { key: "classifier", - label: Advanced: Classification Method, + label: Classification Method, children: ( Advanced: Heuristic Keyword Overrides, + label: Heuristic Keyword Overrides, children: , }, ] : []), { key: "adaptive", - label: Advanced: Adaptive Routing, + label: Adaptive Routing, children: ( @@ -119,39 +138,39 @@ const ComplexityRouterAdvancedSections: React.FCAdvanced: Affinity, + label: Affinity, children: , }, { key: "modality", - label: Advanced: Modality Routing, + label: Modality Routing, children: , }, { key: "plan-mode", - label: Advanced: Plan-Mode Override, + label: Plan-Mode Override, children: ( ), }, { key: "housekeeping", - label: Advanced: Housekeeping Routing, + label: Housekeeping Routing, children: , }, { key: "reminder-markers", - label: Advanced: Ignore Custom Tags, + label: Ignore Custom Tags, children: , }, { key: "context-window", - label: Advanced: Context Window Escalation, + label: Context Window Escalation, children: , }, { key: "stall-escalation", - label: Advanced: Stalled Task Escalation, + label: Stalled Task Escalation, children: ( @@ -160,14 +179,14 @@ const ComplexityRouterAdvancedSections: React.FCAdvanced: Response Format, + label: Response Format, children: , }, ...(onEscalationKeywordsChange ? [ { key: "escalation", - label: Advanced: Escalation Keywords, + label: Escalation Keywords, children: ( @@ -180,7 +199,7 @@ const ComplexityRouterAdvancedSections: React.FCAdvanced: Compression, + label: Compression, children: , }, ] @@ -189,7 +208,7 @@ const ComplexityRouterAdvancedSections: React.FCAdvanced: Keyword/Semantic Matching, + label: Keyword/Semantic Matching, children: ( <> {onKeywordTierRulesChange && ( @@ -220,20 +239,65 @@ const ComplexityRouterAdvancedSections: React.FC(() => + showValidationErrors ? groups.map((group) => group.label) : [], + ); + const [previousValidation, setPreviousValidation] = React.useState(showValidationErrors); + if (previousValidation !== showValidationErrors) { + setPreviousValidation(showValidationErrors); + if (showValidationErrors) setOpenGroups(groups.map((group) => group.label)); + } return ( - <> - {sections - .filter(({ key }) => !forecast || !["adaptive", "context-window", "escalation"].includes(key)) - .map(({ key, label, children }) => ( - - - - {label} - - {children} - - ))} - +
+ {groups.map((group) => ( + + setOpenGroups((current) => + open ? [...current, group.label] : current.filter((label) => label !== group.label), + ) + } + className="border-b border-border last:border-b-0" + > + + + {group.label} + + + {sections + .filter( + ({ key }) => + group.keys.includes(key) && + (!forecast || !["adaptive", "context-window", "escalation"].includes(key)), + ) + .map(({ key, label, children }) => ( +
+ {key !== "classifier" &&

{label}

} + {children} +
+ ))} +
+
+ ))} +
); }; diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.integration.test.tsx similarity index 85% rename from ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx rename to ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.integration.test.tsx index 756e505997c..4a00a469f97 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.integration.test.tsx @@ -1,8 +1,15 @@ +import { openAutoRouterAdvanced, selectAutoRouterOption } from "../../../tests/autoRouterSetup"; import { fireEvent, renderWithProviders, screen, within } from "../../../tests/test-utils"; import userEvent from "@testing-library/user-event"; import React from "react"; -import { vi, type Mock } from "vitest"; -import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; +import { describe, it, expect, vi, type Mock } from "vitest"; +import ComplexityRouterConfigView, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; +import AutoRouterClassifierTabs from "./AutoRouterClassifierTabs"; +const ComplexityRouterConfig = (props: React.ComponentProps) => ( + + + +); vi.mock( "@/app/(dashboard)/hooks/autoRouter/useComplexityScorerDefaults", async () => await import("../../../tests/mocks/complexityScorerDefaults"), @@ -45,12 +52,12 @@ const baseProps = { }; describe("ComplexityRouterConfig", () => { - it("should render", () => { + it("should render", async () => { renderWithProviders(); - expect(screen.getByText("Complexity Tier Configuration")).toBeInTheDocument(); + expect(screen.getByText("Models by tier")).toBeInTheDocument(); }); - it("should display all four tier labels", () => { + it("should display all four tier labels", async () => { renderWithProviders(); expect(screen.getByText("Simple Tier")).toBeInTheDocument(); expect(screen.getByText("Medium Tier")).toBeInTheDocument(); @@ -58,7 +65,7 @@ describe("ComplexityRouterConfig", () => { expect(screen.getByText("Reasoning Tier")).toBeInTheDocument(); }); - it("should show example queries for each tier", () => { + it("should show example queries for each tier", async () => { renderWithProviders(); expect(screen.getByText(/Hello!/)).toBeInTheDocument(); expect(screen.getByText(/Explain how REST APIs work/)).toBeInTheDocument(); @@ -66,46 +73,51 @@ describe("ComplexityRouterConfig", () => { expect(screen.getByText(/Think step by step/)).toBeInTheDocument(); }); - it("should display the how classification works section", () => { + it("should display the how classification works section", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("How Classification Works")).toBeInTheDocument(); }); - it("should show score thresholds in the classification section", () => { + it("should show score thresholds in the classification section", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText(/Score < 0.15/)).toBeInTheDocument(); expect(screen.getByText(/Score 0.15 - 0.35/)).toBeInTheDocument(); expect(screen.getByText(/Score 0.35 - 0.60/)).toBeInTheDocument(); expect(screen.getByText(/Score > 0.60/)).toBeInTheDocument(); }); - it("leaves the score threshold list color to the theme instead of an inline style", () => { + it("leaves the score threshold list color to the theme instead of an inline style", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); const list = screen.getByText(/Score < 0.15/).closest("ul"); expect(list).toBeInTheDocument(); expect(list).toHaveClass("text-muted-foreground"); expect(list?.style.color).toBe(""); }); - it("should default to heuristic and hide classifier model/timeout fields", () => { + it("should default to heuristic and hide classifier model/timeout fields", async () => { renderWithProviders(); - expect(screen.getByText("Advanced: Classification Method")).toBeInTheDocument(); - expect(screen.queryByText("Classifier Model")).not.toBeInTheDocument(); + openAutoRouterAdvanced("Classification Method"); + expect(screen.getByText("Classifier tuning")).toBeInTheDocument(); + expect(screen.queryByText("Judge model")).not.toBeInTheDocument(); }); - it("shows heuristic advanced sections and hides keyword overrides for capability classifiers", () => { + it("shows heuristic advanced sections and hides keyword overrides for capability classifiers", async () => { const { rerender } = renderWithProviders(); - expect(screen.getByText("Advanced: Heuristic Keyword Overrides")).toBeInTheDocument(); - expect(screen.getByText("Advanced: Housekeeping Routing")).toBeInTheDocument(); - expect(screen.getByText("Advanced: Ignore Custom Tags")).toBeInTheDocument(); + openAutoRouterAdvanced("Heuristic Keyword Overrides"); + + expect(screen.getByText("Heuristic Keyword Overrides")).toBeInTheDocument(); + openAutoRouterAdvanced("Housekeeping Routing"); + expect(screen.getByText("Housekeeping Routing")).toBeInTheDocument(); + openAutoRouterAdvanced("Ignore Custom Tags"); + expect(screen.getByText("Ignore Custom Tags")).toBeInTheDocument(); const capabilityValue = { ...defaultValue, classifier_type: "capability" as const }; rerender(); - expect(screen.queryByText("Advanced: Heuristic Keyword Overrides")).not.toBeInTheDocument(); + expect(screen.queryByText("Heuristic Keyword Overrides")).not.toBeInTheDocument(); }); it.each([ @@ -115,7 +127,7 @@ describe("ComplexityRouterConfig", () => { renderWithProviders( , ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); if (visible) { expect(screen.getByLabelText("Classifier plugin timeout (ms)")).toBeInTheDocument(); } else { @@ -128,7 +140,7 @@ describe("ComplexityRouterConfig", () => { renderWithProviders( , ); - fireEvent.click(screen.getByText("Advanced: Ignore Custom Tags")); + openAutoRouterAdvanced("Ignore Custom Tags"); const validation = screen.queryByText(/needs both/i); if (showValidationErrors) { expect(validation).toBeInTheDocument(); @@ -137,11 +149,11 @@ describe("ComplexityRouterConfig", () => { } }); - it("disables housekeeping sentinels when cheapest-tier routing is off", () => { + it("disables housekeeping sentinels when cheapest-tier routing is off", async () => { renderWithProviders( , ); - fireEvent.click(screen.getByText("Advanced: Housekeeping Routing")); + openAutoRouterAdvanced("Housekeeping Routing"); const sentinelInput = screen.getByRole("combobox", { name: "e.g., conversation title" }); expect(sentinelInput).toBeDisabled(); }); @@ -151,7 +163,7 @@ describe("ComplexityRouterConfig", () => { const onChange = vi.fn(); renderWithProviders(); - await user.click(screen.getByText("Advanced: Response Format")); + openAutoRouterAdvanced("Response Format"); await user.click(screen.getByRole("switch", { name: "Return raw model name" })); expect(onChange).toHaveBeenCalledWith({ @@ -160,13 +172,13 @@ describe("ComplexityRouterConfig", () => { }); }); - it("should reveal classifier model and timeout fields when llm is selected", () => { + it("should reveal classifier model and timeout fields when llm is selected", async () => { const onChange = vi.fn(); renderWithProviders(); // Collapse panel content isn't rendered until first expanded. - fireEvent.click(screen.getByText("Advanced: Classification Method")); - fireEvent.click(screen.getByText("LLM Classifier")); + openAutoRouterAdvanced("Classification Method"); + fireEvent.click(screen.getByRole("radio", { name: /^LLM$/ })); const expectedValue: ComplexityRouterConfigValue = { ...defaultValue, @@ -178,14 +190,14 @@ describe("ComplexityRouterConfig", () => { expect(onChange).toHaveBeenCalledWith(expectedValue); }); - it("selects heuristic v2 without requiring a classifier model or showing weighted scoring", () => { + it("selects heuristic v2 without requiring a classifier model or showing weighted scoring", async () => { const onChange = vi.fn(); const { rerender } = renderWithProviders( , ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - fireEvent.click(screen.getByText("Heuristic v2")); + openAutoRouterAdvanced("Classification Method"); + await selectAutoRouterOption("Heuristic", "Heuristic v2"); expect(onChange).toHaveBeenCalledWith( expect.objectContaining({ @@ -197,7 +209,7 @@ describe("ComplexityRouterConfig", () => { const heuristicV2Value: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "heuristic_v2" }; rerender(); - expect(screen.queryByText("Classifier Model")).not.toBeInTheDocument(); + expect(screen.queryByText("Judge model")).not.toBeInTheDocument(); expect(screen.queryByText("Advanced scoring")).not.toBeInTheDocument(); expect(screen.getByText(/estimates success probability for all four tiers/)).toBeInTheDocument(); expect(screen.queryByText(/Score < 0.15/)).not.toBeInTheDocument(); @@ -233,7 +245,7 @@ describe("ComplexityRouterConfig", () => { expect(onChange).toHaveBeenCalledWith({ ...value, heuristic_v2_success_threshold: undefined }); }); - it("shows an inactive zero threshold until explicitly cleared and hides the summary for active or absent values", () => { + it("shows an inactive zero threshold until explicitly cleared and hides the summary for active or absent values", async () => { const onChange = vi.fn(); const value = { ...defaultValue, heuristic_v2_success_threshold: 0 }; const { rerender } = renderWithProviders( @@ -248,7 +260,7 @@ describe("ComplexityRouterConfig", () => { expect(screen.queryByRole("region", { name: "Inactive Heuristic v2 threshold" })).not.toBeInTheDocument(); }); - it("should show classifier fields and use the configured values when classifier_type is llm", () => { + it("should show classifier fields and use the configured values when classifier_type is llm", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -258,9 +270,9 @@ describe("ComplexityRouterConfig", () => { }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); - expect(screen.getByText("Classifier Model")).toBeInTheDocument(); + expect(screen.getByText("Judge model")).toBeInTheDocument(); expect(screen.getByLabelText("Timeout (ms)")).toHaveValue("750"); expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).toBeChecked(); expect(screen.getByLabelText("Circuit breaker cooldown (seconds)")).toHaveValue("30"); @@ -268,7 +280,7 @@ describe("ComplexityRouterConfig", () => { expect(screen.queryByText("Context Per-Turn Character Limit")).not.toBeInTheDocument(); }); - it("should allow the default-on classifier circuit breaker to be disabled", () => { + it("should allow the default-on classifier circuit breaker to be disabled", async () => { const onChange = vi.fn(); const llmValue: ComplexityRouterConfigValue = { ...defaultValue, @@ -276,7 +288,7 @@ describe("ComplexityRouterConfig", () => { classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" })); @@ -287,7 +299,7 @@ describe("ComplexityRouterConfig", () => { ); }); - it("should default the context window and budget when llm is selected", () => { + it("should default the context window and budget when llm is selected", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -295,13 +307,13 @@ describe("ComplexityRouterConfig", () => { }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByLabelText("Context Window Size")).toHaveValue("3"); expect(screen.getByLabelText("Context Character Budget")).toHaveValue("8000"); }); - it("should warn when the budget is too small to quote any turn that does not already fit", () => { + it("should warn when the budget is too small to quote any turn that does not already fit", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -310,12 +322,12 @@ describe("ComplexityRouterConfig", () => { }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText(/no room to quote a turn/i)).toBeInTheDocument(); }); - it("should not warn on a budget large enough to quote a turn, nor on a deliberate zero", () => { + it("should not warn on a budget large enough to quote a turn, nor on a deliberate zero", async () => { for (const budget of [120, 8000, 0]) { const { unmount } = renderWithProviders( { onChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.queryByText(/no room to quote a turn/i)).not.toBeInTheDocument(); unmount(); } }); - it("should show the assistant-turns switch with its configured value when classifier_type is llm", () => { + it("should show the assistant-turns switch with its configured value when classifier_type is llm", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -344,13 +356,13 @@ describe("ComplexityRouterConfig", () => { }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("Include Assistant Turns")).toBeInTheDocument(); expect(screen.getByRole("switch", { name: "Include Assistant Turns" })).toBeChecked(); }); - it("should render the assistant-turns switch off when it is not set", () => { + it("should render the assistant-turns switch off when it is not set", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -358,18 +370,18 @@ describe("ComplexityRouterConfig", () => { }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByRole("switch", { name: "Include Assistant Turns" })).not.toBeChecked(); }); - it("should hide the assistant-turns switch when classifier_type is heuristic", () => { + it("should hide the assistant-turns switch when classifier_type is heuristic", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.queryByText("Include Assistant Turns")).not.toBeInTheDocument(); }); - it("should call onChange when the assistant-turns switch is toggled", () => { + it("should call onChange when the assistant-turns switch is toggled", async () => { const onChange = vi.fn(); const llmValue: ComplexityRouterConfigValue = { ...defaultValue, @@ -378,7 +390,7 @@ describe("ComplexityRouterConfig", () => { }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); fireEvent.click(screen.getByRole("switch", { name: "Include Assistant Turns" })); expect(onChange).toHaveBeenCalledWith( @@ -386,9 +398,9 @@ describe("ComplexityRouterConfig", () => { ); }); - it("should hide classifier context fields when classifier_type is heuristic", () => { + it("should hide classifier context fields when classifier_type is heuristic", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.queryByText("Context Window Size")).not.toBeInTheDocument(); expect(screen.queryByText("Context Per-Turn Character Limit")).not.toBeInTheDocument(); }); @@ -416,7 +428,7 @@ describe("ComplexityRouterConfig", () => { classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); const input = screen.getByLabelText(label); fireEvent.change(input, { target: { value: "" } }); @@ -429,14 +441,14 @@ describe("ComplexityRouterConfig", () => { expect(onChange).toHaveBeenLastCalledWith({ ...llmValue, ...expected }); }); - it("restores the committed context window size after an empty field loses focus", () => { + it("restores the committed context window size after an empty field loses focus", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); const input = screen.getByLabelText("Context Window Size"); fireEvent.change(input, { target: { value: "" } }); @@ -445,13 +457,13 @@ describe("ComplexityRouterConfig", () => { expect(input).toHaveValue("3"); }); - it("should render the custom technical keywords field", () => { + it("should render the custom technical keywords field", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("Custom Technical Keywords")).toBeInTheDocument(); }); - it("should display existing custom technical keywords as tags", () => { + it("should display existing custom technical keywords as tags", async () => { renderWithProviders( { onCustomTechnicalKeywordsChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("udp")).toBeInTheDocument(); expect(screen.getByText("kafka")).toBeInTheDocument(); }); @@ -474,7 +486,7 @@ describe("ComplexityRouterConfig", () => { onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement; await user.type(within(keywordsSection).getByRole("combobox"), "udp"); await user.click(await screen.findByText('Create "udp"')); @@ -491,35 +503,35 @@ describe("ComplexityRouterConfig", () => { onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement; await user.type(within(keywordsSection).getByRole("combobox"), "udp, kafka ,terraform"); await user.click(await screen.findByText('Create "udp, kafka ,terraform"')); expect(onCustomTechnicalKeywordsChange).toHaveBeenCalledWith(["udp", "kafka", "terraform"]); }); - it("should render an empty state when no keyword tier rules exist", () => { + it("should render an empty state when no keyword tier rules exist", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); expect(screen.getByText("Keyword Tier Overrides")).toBeInTheDocument(); expect(screen.getByText("No keyword tier overrides configured")).toBeInTheDocument(); }); - it("hides the keyword-tier and semantic sections when their change handlers are absent (edit modal)", () => { + it("hides the keyword-tier and semantic sections when their change handlers are absent (edit modal)", async () => { // The edit-auto-router modal renders ComplexityRouterConfig without these handlers; // the sections must stay hidden rather than render interactive-but-dead controls. renderWithProviders(); expect(screen.queryByText("Keyword Tier Overrides")).not.toBeInTheDocument(); expect(screen.queryByText("Semantic keyword matching")).not.toBeInTheDocument(); // Core tier config still renders. - expect(screen.getByText("Complexity Tier Configuration")).toBeInTheDocument(); + expect(screen.getByText("Models by tier")).toBeInTheDocument(); }); it("should call onKeywordTierRulesChange with a new rule when 'Add keyword rule' is clicked", async () => { const user = userEvent.setup(); const onKeywordTierRulesChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); expect(onKeywordTierRulesChange).toHaveBeenCalledTimes(1); const newRules = onKeywordTierRulesChange.mock.calls[0][0]; @@ -537,7 +549,7 @@ describe("ComplexityRouterConfig", () => { onKeywordTierRulesChange={onKeywordTierRulesChange} />, ); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); const field = screen.getByText("Keywords 1").closest("div") as HTMLElement; await user.type(within(field).getByRole("combobox"), "invoice"); @@ -556,7 +568,7 @@ describe("ComplexityRouterConfig", () => { onKeywordTierRulesChange={onKeywordTierRulesChange} />, ); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); expect(screen.getByText("invoice")).toBeInTheDocument(); expect(screen.getByText("refund")).toBeInTheDocument(); @@ -564,17 +576,17 @@ describe("ComplexityRouterConfig", () => { expect(onKeywordTierRulesChange).toHaveBeenCalledWith([]); }); - it("should not show embedding model or match score fields when semantic matching is disabled", () => { + it("should not show embedding model or match score fields when semantic matching is disabled", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); expect(screen.getByText("Semantic keyword matching")).toBeInTheDocument(); expect(screen.queryByText("Embedding model")).not.toBeInTheDocument(); expect(screen.queryByText("Minimum match score")).not.toBeInTheDocument(); }); - it("should show embedding model and match score fields when semantic matching is enabled", () => { + it("should show embedding model and match score fields when semantic matching is enabled", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); expect(screen.getByText("Embedding model")).toBeInTheDocument(); expect(screen.getByText("Minimum match score")).toBeInTheDocument(); }); @@ -589,7 +601,7 @@ describe("ComplexityRouterConfig", () => { onSemanticMatchingEnabledChange={onSemanticMatchingEnabledChange} />, ); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); await user.click(screen.getByRole("switch", { name: "Semantic keyword matching" })); expect(onSemanticMatchingEnabledChange).toHaveBeenCalledWith(true, expect.anything()); }); @@ -606,34 +618,34 @@ describe("ComplexityRouterConfig", () => { expect(screen.queryAllByText("text-embedding-3-small")).toHaveLength(0); }); - it("does not show tier validation errors by default", () => { + it("does not show tier validation errors by default", async () => { renderWithProviders(); expect(screen.queryByText("This tier is required")).not.toBeInTheDocument(); }); - it("shows an inline error on the classifier model select when llm is selected without a model", () => { + it("shows an inline error on the classifier model select when llm is selected without a model", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", classifier_llm_config: { model: "", timeout_ms: 3000 }, }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByText("A classifier model is required")).toBeInTheDocument(); + openAutoRouterAdvanced("Classification Method"); + expect(screen.getByText("A judge model is required")).toBeInTheDocument(); }); - it("does not show the classifier model error once a classifier model is set", () => { + it("does not show the classifier model error once a classifier model is set", async () => { const llmValue: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.queryByText("A classifier model is required")).not.toBeInTheDocument(); + openAutoRouterAdvanced("Classification Method"); + expect(screen.queryByText("A judge model is required")).not.toBeInTheDocument(); }); - it("shows a validation error only under unfilled tiers when showValidationErrors is true", () => { + it("shows a validation error only under unfilled tiers when showValidationErrors is true", async () => { renderWithProviders( { expect(screen.getAllByText(/tier is required/)).toHaveLength(1); }); - it("renders the escalation keywords section with current keywords when the handler is provided", () => { + it("renders the escalation keywords section with current keywords when the handler is provided", async () => { renderWithProviders( { onEscalationKeywordsChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Escalation Keywords")); - expect(screen.getByText("Escalation Keywords")).toBeInTheDocument(); + openAutoRouterAdvanced("Escalation Keywords"); + expect(screen.getAllByText("Escalation Keywords")).not.toHaveLength(0); expect(screen.getByText("LITELLM ESCALATE")).toBeInTheDocument(); }); - it("hides the escalation keywords section when no handler is provided", () => { + it("hides the escalation keywords section when no handler is provided", async () => { renderWithProviders(); - expect(screen.queryByText("Advanced: Escalation Keywords")).not.toBeInTheDocument(); + expect(screen.queryByText("Escalation Keywords")).not.toBeInTheDocument(); }); }); @@ -671,21 +683,21 @@ describe("ComplexityRouterConfig classifier fallback", () => { classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, }; - it("defaults the fallback to the heuristic, matching the backend field default", () => { + it("defaults the fallback to the heuristic, matching the backend field default", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByRole("radio", { name: /Score with the heuristic/ })).toBeChecked(); }); - it("records a switch to the default model fallback", () => { + it("records a switch to the default model fallback", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); fireEvent.click(screen.getByRole("radio", { name: /Route to the default model/ })); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ classifier_fallback: "default_model" })); }); - it("disables the default model fallback when no tier would produce one", () => { + it("disables the default model fallback when no tier would produce one", async () => { // The deployment's default model is derived from the tiers on submit, so offering the option // with no tiers picked would save a config the backend rejects at startup. const noTiers: ComplexityRouterConfigValue = { @@ -693,17 +705,17 @@ describe("ComplexityRouterConfig classifier fallback", () => { tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] }, }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByRole("radio", { name: /Route to the default model/ })).toHaveAttribute("aria-disabled", "true"); }); - it("hides the fallback choice for the heuristic classifier, which has nothing to fall back from", () => { + it("hides the fallback choice for the heuristic classifier, which has nothing to fall back from", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.queryByText("If the classifier fails")).not.toBeInTheDocument(); }); - it("stops describing the heuristic as the fallback once a custom prompt routes failures to the default model", () => { + it("stops describing the heuristic as the fallback once a custom prompt routes failures to the default model", async () => { // With both set, the heuristic scorer never runs, so the panel must not keep implying a // score decides anything on this router. renderWithProviders( @@ -717,11 +729,11 @@ describe("ComplexityRouterConfig classifier fallback", () => { onChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText(/no longer runs at all/)).toBeInTheDocument(); }); - it("still describes the heuristic as the fallback when a custom prompt keeps heuristic fallback", () => { + it("still describes the heuristic as the fallback when a custom prompt keeps heuristic fallback", async () => { renderWithProviders( { onChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText(/only when the classifier call fails/)).toBeInTheDocument(); }); - it("clears a stored fallback when switching back to the heuristic classifier", () => { + it("clears a stored fallback when switching back to the heuristic classifier", async () => { const onChange = vi.fn(); renderWithProviders( { onChange={onChange} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - fireEvent.click(screen.getByRole("radio", { name: /rule-based scoring/ })); + openAutoRouterAdvanced("Classification Method"); + fireEvent.click(screen.getByRole("radio", { name: /^Heuristics$/ })); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ classifier_fallback: undefined })); }); }); @@ -758,19 +770,21 @@ describe("ComplexityRouterConfig classification frequency", () => { classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 }, }; - it("defaults to every request, matching both backend field defaults", () => { + it("defaults to every request, matching both backend field defaults", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Every request/ })).toBeChecked(); - expect(screen.getByRole("radio", { name: /Every new user message/ })).not.toBeChecked(); - expect(screen.getByRole("radio", { name: /Once per session/ })).not.toBeChecked(); + openAutoRouterAdvanced("Classification Method"); + expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent("Every request"); + expect(screen.getByRole("combobox", { name: "How often to classify" })).not.toHaveTextContent( + "Every new user message", + ); + expect(screen.getByRole("combobox", { name: "How often to classify" })).not.toHaveTextContent("Once per session"); }); - it("writes both wire fields when the frequency moves to every new user message", () => { + it("writes both wire fields when the frequency moves to every new user message", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - fireEvent.click(screen.getByRole("radio", { name: /Every new user message/ })); + openAutoRouterAdvanced("Classification Method"); + await selectAutoRouterOption("How often to classify", "Every new user message"); expect(onChange).toHaveBeenCalledWith({ ...llmValue, classification_mode: "user_turn", @@ -778,11 +792,11 @@ describe("ComplexityRouterConfig classification frequency", () => { }); }); - it("writes session affinity, not a classification mode, when the frequency moves to once per session", () => { + it("writes session affinity, not a classification mode, when the frequency moves to once per session", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - fireEvent.click(screen.getByRole("radio", { name: /Once per session/ })); + openAutoRouterAdvanced("Classification Method"); + await selectAutoRouterOption("How often to classify", "Once per session"); expect(onChange).toHaveBeenCalledWith({ ...llmValue, classification_mode: "every_request", @@ -790,7 +804,7 @@ describe("ComplexityRouterConfig classification frequency", () => { }); }); - it("shows a hand-authored config that sets both fields as once per session, matching the backend", () => { + it("shows a hand-authored config that sets both fields as once per session, matching the backend", async () => { renderWithProviders( { onChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Once per session/ })).toBeChecked(); - expect(screen.getByRole("radio", { name: /Every new user message/ })).not.toBeChecked(); + openAutoRouterAdvanced("Classification Method"); + expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent("Once per session"); + expect(screen.getByRole("combobox", { name: "How often to classify" })).not.toHaveTextContent( + "Every new user message", + ); }); - it("records a switch back to every request", () => { + it("records a switch back to every request", async () => { const onChange = vi.fn(); renderWithProviders( { onChange={onChange} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Every new user message/ })).toBeChecked(); - fireEvent.click(screen.getByRole("radio", { name: /Every request/ })); + openAutoRouterAdvanced("Classification Method"); + expect(screen.getByRole("combobox", { name: "How often to classify" })).toHaveTextContent("Every new user message"); + await selectAutoRouterOption("How often to classify", "Every request"); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ classification_mode: "every_request" })); }); - it("offers the frequency on a heuristic router, where holding the tier still pins the model", () => { + it("offers the frequency on a heuristic router, where holding the tier still pins the model", async () => { // The backend pin is gated on the two fields alone, so a heuristic router that switches models // mid tool loop is fixed by this control too. renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Every new user message/ })).toBeInTheDocument(); + openAutoRouterAdvanced("Classification Method"); + expect(screen.getByRole("combobox", { name: "How often to classify" })).toBeVisible(); }); }); @@ -836,11 +852,11 @@ describe("ComplexityRouterConfig classifier rubric", () => { const openClassificationPanel = (value: ComplexityRouterConfigValue, onChange = vi.fn()) => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); return onChange; }; - it("shows an existing router with no stored preset as legacy in the prompt control", () => { + it("shows an existing router with no stored preset as legacy in the prompt control", async () => { // This router predates the setting. Displaying a calibrated preset it does not have would tell the // operator their traffic is graded by examples the classifier never receives, and saving the form // unchanged would then move its tier decisions. @@ -849,19 +865,19 @@ describe("ComplexityRouterConfig classifier rubric", () => { expect(screen.getByRole("button", { name: "Customize prompt" })).toBeInTheDocument(); }); - it("stamps the calibrated preset on a classifier being switched on for the first time", () => { + it("stamps the calibrated preset on a classifier being switched on for the first time", async () => { // A heuristic router turning on the LLM classifier has no prior tier behaviour to preserve, so a // newly configured classifier starts on the calibrated rubric rather than the legacy one. const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - fireEvent.click(screen.getByText("LLM Classifier")); + openAutoRouterAdvanced("Classification Method"); + fireEvent.click(screen.getByRole("radio", { name: /^LLM$/ })); expect(onChange).toHaveBeenCalledWith( expect.objectContaining({ classifier_llm_config: expect.objectContaining({ classification_rubric: "agentic" }) }), ); }); - it("shows the calibrated preset when a router stores one", () => { + it("shows the calibrated preset when a router stores one", async () => { openClassificationPanel({ ...llmValue, classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000, classification_rubric: "agentic" }, @@ -906,7 +922,7 @@ describe("ComplexityRouterConfig classifier rubric", () => { ); }); - it("keeps the rubric out of the legacy whole-prompt editor, which replaces it entirely", () => { + it("keeps the rubric out of the legacy whole-prompt editor, which replaces it entirely", async () => { // The backend rejects both together, so the legacy editor must not offer a rubric to pick. openClassificationPanel({ ...llmValue, @@ -916,7 +932,7 @@ describe("ComplexityRouterConfig classifier rubric", () => { expect(screen.queryByRole("button", { name: "Customize prompt" })).not.toBeInTheDocument(); }); - it("hides the prompt control for the heuristic classifier, which sends no prompt at all", () => { + it("hides the prompt control for the heuristic classifier, which sends no prompt at all", async () => { openClassificationPanel(defaultValue); expect(screen.queryByRole("button", { name: "Customize prompt" })).not.toBeInTheDocument(); }); @@ -928,7 +944,7 @@ describe("ComplexityRouterConfig tier labels", () => { tier_labels: { SIMPLE: "Cheap", MEDIUM: "Standard", COMPLEX: "Premium", REASONING: "Deep" }, }; - it("shows the operator's names in the tier headers instead of the defaults", () => { + it("shows the operator's names in the tier headers instead of the defaults", async () => { renderWithProviders(); expect(screen.getByText("Cheap Tier")).toBeInTheDocument(); expect(screen.getByText("Deep Tier")).toBeInTheDocument(); @@ -936,13 +952,13 @@ describe("ComplexityRouterConfig tier labels", () => { expect(screen.queryByText("Reasoning Tier")).not.toBeInTheDocument(); }); - it("keeps the rung ordinal and canonical name visible under a rename", () => { + it("keeps the rung ordinal and canonical name visible under a rename", async () => { renderWithProviders(); expect(screen.getByText(/Tier 1 of 4/)).toHaveTextContent("Tier 1 of 4 · SIMPLE"); expect(screen.getByText(/Tier 4 of 4/)).toHaveTextContent("Tier 4 of 4 · REASONING"); }); - it("names the renamed tier in the required-field error", () => { + it("names the renamed tier in the required-field error", async () => { renderWithProviders( { expect(screen.getByText("The Deep tier is required")).toBeInTheDocument(); }); - it("reports a typed label back to the caller under its canonical tier key", () => { + it("reports a typed label back to the caller under its canonical tier key", async () => { const onChange = vi.fn(); renderWithProviders(); fireEvent.change(screen.getByLabelText("Display name for the Simple tier"), { target: { value: "Cheap" } }); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ tier_labels: { SIMPLE: "Cheap" } })); }); - it("shows a stored label in its input so an edit round-trips", () => { + it("shows a stored label in its input so an edit round-trips", async () => { renderWithProviders(); expect(screen.getByLabelText("Display name for the Reasoning tier")).toHaveValue("Deep"); }); - it("leaves the label inputs empty when nothing was renamed", () => { + it("leaves the label inputs empty when nothing was renamed", async () => { renderWithProviders(); expect(screen.getByLabelText("Display name for the Simple tier")).toHaveValue(""); }); - it("uses the operator's names in the classification score table", () => { + it("uses the operator's names in the classification score table", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("Cheap")).toBeInTheDocument(); expect(screen.getByText("Deep")).toBeInTheDocument(); }); - it("uses the operator's names in the keyword rule tier picker", () => { + it("uses the operator's names in the keyword rule tier picker", async () => { renderWithProviders( { keywordTierRules={[{ id: "r1", keywords: ["invoice"], tier: "REASONING" }]} />, ); - fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); expect(screen.getByRole("combobox", { name: "Route keyword rule 1 to tier" })).toHaveTextContent("Deep"); }); }); describe("ComplexityRouterConfig modality panel", () => { - it("defaults the image-routing switch off and writes modality_routing through onChange", () => { + it("defaults the image-routing switch off and writes modality_routing through onChange", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Modality Routing")); + openAutoRouterAdvanced("Modality Routing"); const toggle = screen.getByRole("switch", { name: "Route image requests to vision-capable models" }); expect(toggle).not.toBeChecked(); @@ -1003,19 +1019,19 @@ describe("ComplexityRouterConfig modality panel", () => { expect(onChange).toHaveBeenCalledWith({ ...defaultValue, modality_routing: true }); }); - it("renders a stored modality_routing=true as on", () => { + it("renders a stored modality_routing=true as on", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Modality Routing")); + openAutoRouterAdvanced("Modality Routing"); expect(screen.getByRole("switch", { name: "Route image requests to vision-capable models" })).toBeChecked(); }); // The backend ignores modality_pin_override unless modality_routing is on, so offering it while // image routing is off would let an operator save a flag that does nothing. - it("disables the pin-override switch while image routing is off", () => { + it("disables the pin-override switch while image routing is off", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Modality Routing")); + openAutoRouterAdvanced("Modality Routing"); const override = screen.getByRole("switch", { name: "Override session pin for image requests" }); expect(override).toHaveAttribute("aria-disabled", "true"); @@ -1023,11 +1039,11 @@ describe("ComplexityRouterConfig modality panel", () => { expect(onChange).not.toHaveBeenCalled(); }); - it("writes modality_pin_override through onChange once image routing is on", () => { + it("writes modality_pin_override through onChange once image routing is on", async () => { const onChange = vi.fn(); const value = { ...defaultValue, modality_routing: true }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Modality Routing")); + openAutoRouterAdvanced("Modality Routing"); const override = screen.getByRole("switch", { name: "Override session pin for image requests" }); expect(override).not.toBeChecked(); @@ -1036,51 +1052,51 @@ describe("ComplexityRouterConfig modality panel", () => { expect(onChange).toHaveBeenCalledWith({ ...value, modality_pin_override: true }); }); - it("renders a stored modality_pin_override=true as on", () => { + it("renders a stored modality_pin_override=true as on", async () => { renderWithProviders( , ); - fireEvent.click(screen.getByText("Advanced: Modality Routing")); + openAutoRouterAdvanced("Modality Routing"); expect(screen.getByRole("switch", { name: "Override session pin for image requests" })).toBeChecked(); }); }); describe("ComplexityRouterConfig affinity panel", () => { - it("holds the deployment switch at its backend default, session pinning having moved to the frequency choice", () => { + it("holds the deployment switch at its backend default, session pinning having moved to the frequency choice", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Affinity")); + openAutoRouterAdvanced("Affinity"); expect(screen.getByRole("switch", { name: "Pin one model deployment per tier" })).toBeChecked(); expect(screen.queryByRole("switch", { name: "Pin a session to its first model" })).not.toBeInTheDocument(); }); - it("writes deployment_affinity through onChange without touching other keys", () => { + it("writes deployment_affinity through onChange without touching other keys", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Affinity")); + openAutoRouterAdvanced("Affinity"); fireEvent.click(screen.getByRole("switch", { name: "Pin one model deployment per tier" })); expect(onChange).toHaveBeenCalledWith({ ...defaultValue, deployment_affinity: false }); }); - it("renders a stored deployment_affinity=false as off", () => { + it("renders a stored deployment_affinity=false as off", async () => { renderWithProviders( , ); - fireEvent.click(screen.getByText("Advanced: Affinity")); + openAutoRouterAdvanced("Affinity"); expect(screen.getByRole("switch", { name: "Pin one model deployment per tier" })).not.toBeChecked(); }); - it("writes an idle TTL on blur and keeps the partial input as a draft while typing", () => { + it("writes an idle TTL on blur and keeps the partial input as a draft while typing", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Affinity")); + openAutoRouterAdvanced("Affinity"); const ttl = screen.getByLabelText("How long a pin survives idle (seconds)"); expect(ttl).toHaveAttribute("placeholder", "3600"); @@ -1091,11 +1107,11 @@ describe("ComplexityRouterConfig affinity panel", () => { expect(onChange).toHaveBeenCalledWith({ ...defaultValue, session_affinity_ttl_seconds: 300 }); }); - it("clearing the idle TTL returns the router to its backend default", () => { + it("clearing the idle TTL returns the router to its backend default", async () => { const onChange = vi.fn(); const value = { ...defaultValue, session_affinity_ttl_seconds: 300 }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Affinity")); + openAutoRouterAdvanced("Affinity"); const ttl = screen.getByLabelText("How long a pin survives idle (seconds)"); expect(ttl).toHaveValue("300"); @@ -1105,10 +1121,10 @@ describe("ComplexityRouterConfig affinity panel", () => { expect(onChange).toHaveBeenCalledWith({ ...value, session_affinity_ttl_seconds: undefined }); }); - it("clamps a non-positive idle TTL to the backend's minimum", () => { + it("clamps a non-positive idle TTL to the backend's minimum", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Affinity")); + openAutoRouterAdvanced("Affinity"); const ttl = screen.getByLabelText("How long a pin survives idle (seconds)"); fireEvent.change(ttl, { target: { value: "0" } }); @@ -1121,12 +1137,12 @@ describe("ComplexityRouterConfig affinity panel", () => { describe("ComplexityRouterConfig default model", () => { const getDefaultModelSelect = () => screen.getByRole("combobox", { name: "Default model" }); - it("shows what the tiers currently imply, so an untouched router still names its default", () => { + it("shows what the tiers currently imply, so an untouched router still names its default", async () => { renderWithProviders(); expect(getDefaultModelSelect()).toHaveAttribute("placeholder", "Derived from tiers: gpt-3.5-turbo"); }); - it("asks for a model rather than naming a derived one when no tier holds one", () => { + it("asks for a model rather than naming a derived one when no tier holds one", async () => { const noTiers: ComplexityRouterConfigValue = { ...defaultValue, tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] }, @@ -1157,13 +1173,13 @@ describe("ComplexityRouterConfig default model", () => { expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ default_model: undefined })); }); - it("shows a pinned model as the selection instead of the tier-derived one", () => { + it("shows a pinned model as the selection instead of the tier-derived one", async () => { const pinned: ComplexityRouterConfigValue = { ...defaultValue, default_model: "claude-3-opus" }; renderWithProviders(); expect(getDefaultModelSelect()).toHaveValue("claude-3-opus"); }); - it("unlocks the default model fallback on a pin alone, with no tier to derive from", () => { + it("unlocks the default model fallback on a pin alone, with no tier to derive from", async () => { const pinnedNoTiers: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -1172,11 +1188,11 @@ describe("ComplexityRouterConfig default model", () => { default_model: "claude-3-opus", }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByRole("radio", { name: /Route to the default model/ })).not.toHaveAttribute("aria-disabled"); }); - it("names the resolved default on the fallback option, so the destination is not a guess", () => { + it("names the resolved default on the fallback option, so the destination is not a guess", async () => { const pinned: ComplexityRouterConfigValue = { ...defaultValue, classifier_type: "llm", @@ -1184,13 +1200,13 @@ describe("ComplexityRouterConfig default model", () => { default_model: "claude-3-opus", }; renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByRole("radio", { name: /Route to the default model \(claude-3-opus\)/ })).toBeInTheDocument(); }); }); describe("plan-mode override", () => { - const openPanel = () => fireEvent.click(screen.getByText("Advanced: Plan-Mode Override")); + const openPanel = () => openAutoRouterAdvanced("Plan-Mode Override"); const switchName = "Route plan-mode requests to a minimum tier"; it("toggling on floors at the highest tier that has models", async () => { @@ -1248,13 +1264,13 @@ describe("plan-mode override", () => { }); describe("ComplexityRouterConfig per-model reasoning effort", () => { - it("renders one effort select per selected model, defaulting to Default", () => { + it("renders one effort select per selected model, defaulting to Default", async () => { renderWithProviders(); const select = screen.getByRole("combobox", { name: "Reasoning effort for gpt-4 in the Complex tier" }); expect(select).toHaveTextContent("Default"); }); - it("shows the hydrated effort for a model that has one stored", () => { + it("shows the hydrated effort for a model that has one stored", async () => { renderWithProviders( { const renderClassifier = (value: ComplexityRouterConfigValue = llmValue, onChange = vi.fn()) => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); return onChange; }; @@ -1348,7 +1364,7 @@ describe("ComplexityRouterConfig classifier reasoning effort", () => { classifier_llm_config: { model: "gpt-4", timeout_ms: 3000, reasoning_effort: "high" }, }); const user = userEvent.setup(); - await user.click(screen.getByRole("combobox", { name: "Classifier Model" })); + await user.click(screen.getByRole("combobox", { name: "Judge model" })); await user.click(await screen.findByRole("option", { name: "gpt-3.5-turbo" })); expect(onChange).toHaveBeenCalledWith({ ...llmValue, @@ -1364,7 +1380,7 @@ describe("ComplexityRouterConfig classifier reasoning effort", () => { classifier_llm_config: { model: "gpt-4", timeout_ms: 3000, reasoning_effort: "high" }, }); const user = userEvent.setup(); - await user.click(screen.getByRole("combobox", { name: "Classifier Model" })); + await user.click(screen.getByRole("combobox", { name: "Judge model" })); if (action === "click") await user.click(await screen.findByRole("option", { name: "gpt-4" })); else await user.keyboard("{Enter}"); expect(onChange).not.toHaveBeenCalled(); @@ -1397,7 +1413,7 @@ describe("ComplexityRouterConfig classifier reasoning effort", () => { }); describe("ComplexityRouterConfig reasoning effort gating", () => { - it("offers no effort select for a model group without reasoning support", () => { + it("offers no effort select for a model group without reasoning support", async () => { renderWithProviders(); expect( screen.queryByRole("combobox", { name: "Reasoning effort for gpt-3.5-turbo in the Simple tier" }), @@ -1406,7 +1422,7 @@ describe("ComplexityRouterConfig reasoning effort gating", () => { // A stored effort on a model the group info calls non-reasoning must stay visible, or the // operator has no way to clear it. - it("keeps the select for a non-reasoning model that already has a stored effort", () => { + it("keeps the select for a non-reasoning model that already has a stored effort", async () => { renderWithProviders( { // An empty list is the group's own answer that its deployments share no level, which is different // from the field being absent, so the control is dropped rather than falling back to every level. - it("offers no effort at all when the group intersects to nothing", () => { + it("offers no effort at all when the group intersects to nothing", async () => { renderWithProviders( { // Hand-authored configs can carry a level outside the supported set (e.g. max); it must render // and stay clearable rather than being masked as Default. - it("keeps showing a stored effort outside the supported set", () => { + it("keeps showing a stored effort outside the supported set", async () => { renderWithProviders( { describe("ComplexityRouterConfig custom technical keywords", () => { const openClassificationPanel = (value: ComplexityRouterConfigValue) => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); }; const llmConfig = { model: "gpt-3.5-turbo", timeout_ms: 3000 }; @@ -1503,7 +1519,7 @@ describe("ComplexityRouterConfig custom technical keywords", () => { expect(screen.getByText("Custom Technical Keywords")).toBeInTheDocument(); }); - it("hides the keywords when the scorer never runs, so they cannot imply an effect they have none", () => { + it("hides the keywords when the scorer never runs, so they cannot imply an effect they have none", async () => { const llmWithDefaultFallback = { ...defaultValue, classifier_type: "llm" as const, @@ -1547,19 +1563,19 @@ describe("ComplexityRouterConfig tier editing", () => { }, }; - it("offers Edit tiers only when the parent owns the editor flag", () => { + it("offers Edit tiers only when the parent owns the editor flag", async () => { renderWithProviders(); expect(screen.queryByRole("button", { name: "Edit tiers" })).not.toBeInTheDocument(); }); - it("surfaces the caller's orphaned-rule verdict while editing, so Done is not a silent exit", () => { + it("surfaces the caller's orphaned-rule verdict while editing, so Done is not a silent exit", async () => { renderEditor(customValue, { keywordRulesError: "Keyword rule(s) 1 route to a tier this router no longer has" }); expect( screen.getByText("Keyword rule(s) 1 route to a tier this router no longer has", { exact: false }), ).toBeInTheDocument(); }); - it("keeps the orphaned-rule verdict out of the collapsed view, where the submit tooltip owns it", () => { + it("keeps the orphaned-rule verdict out of the collapsed view, where the submit tooltip owns it", async () => { renderWithProviders( { expect(screen.queryByText("route to a tier this router no longer has", { exact: false })).not.toBeInTheDocument(); }); - it("renders the four built-in tiers before any edit, unchanged", () => { + it("renders the four built-in tiers before any edit, unchanged", async () => { renderWithProviders(); expect(screen.getByRole("button", { name: "Edit tiers" })).toBeInTheDocument(); expect(screen.getByText("Tier 1 of 4", { exact: false })).toHaveTextContent("SIMPLE"); }); - it("adds a row and moves the form into an edited tier set, which the built-in record never leaves", () => { + it("adds a row and moves the form into an edited tier set, which the built-in record never leaves", async () => { const { committed } = renderEditor(); fireEvent.click(screen.getByRole("button", { name: "Add tier" })); const next = committed(); @@ -1585,7 +1601,7 @@ describe("ComplexityRouterConfig tier editing", () => { expect(next.tiers).toEqual(defaultValue.tiers); }); - it("renames a built-in tier straight from the editor, which is what makes the set custom", () => { + it("renames a built-in tier straight from the editor, which is what makes the set custom", async () => { const { committed } = renderEditor(); fireEvent.change(screen.getByLabelText("Name for tier 3"), { target: { value: "SECURITY_REVIEW" } }); const next = committed(); @@ -1598,13 +1614,13 @@ describe("ComplexityRouterConfig tier editing", () => { expect(next.tiers).toEqual(defaultValue.tiers); }); - it("opening the editor and changing nothing leaves the router on the built-in tiers", () => { + it("opening the editor and changing nothing leaves the router on the built-in tiers", async () => { const { onChange } = renderEditor(); expect(screen.getByRole("button", { name: "Done" })).toBeEnabled(); expect(onChange).not.toHaveBeenCalled(); }); - it("swaps the display-name field for the tier-name field while the editor is open", () => { + it("swaps the display-name field for the tier-name field while the editor is open", async () => { const { rerender } = renderWithProviders(); expect(screen.getByLabelText("Display name for the Simple tier")).toBeInTheDocument(); rerender(); @@ -1612,23 +1628,23 @@ describe("ComplexityRouterConfig tier editing", () => { expect(screen.getByLabelText("Name for tier 1")).toBeInTheDocument(); }); - it("drops the scorer card entirely once an edited tier set replaces the heuristic", () => { + it("drops the scorer card entirely once an edited tier set replaces the heuristic", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.queryByText("How Classification Works")).not.toBeInTheDocument(); expect( screen.queryByText("scores each request across 7 built-in dimensions", { exact: false }), ).not.toBeInTheDocument(); }); - it("keeps the scorer card on a built-in router, whose tiers the score still decides", () => { + it("keeps the scorer card on a built-in router, whose tiers the score still decides", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("How Classification Works")).toBeInTheDocument(); expect(screen.getByText("scores each request across 7 built-in dimensions", { exact: false })).toBeInTheDocument(); }); - it("says why a custom row is blocked instead of only reddening its border", () => { + it("says why a custom row is blocked instead of only reddening its border", async () => { const missingDefinition: ComplexityRouterConfigValue = { ...customValue, custom_tier_set: { @@ -1652,24 +1668,24 @@ describe("ComplexityRouterConfig tier editing", () => { expect(screen.getByRole("button", { name: "Done" })).toBeDisabled(); }); - it("enables Done once every row carries a name, a definition and a model", () => { + it("enables Done once every row carries a name, a definition and a model", async () => { renderEditor(customValue); expect(screen.getByRole("button", { name: "Done" })).toBeEnabled(); }); - it("refuses to remove a row that would take the set below the backend's minimum", () => { + it("refuses to remove a row that would take the set below the backend's minimum", async () => { renderEditor(customValue); expect(screen.getByRole("button", { name: "Remove the CASUAL tier" })).toBeDisabled(); }); - it("keeps a definition on one line, because the backend rejects a newline in it", () => { + it("keeps a definition on one line, because the backend rejects a newline in it", async () => { const { committed } = renderEditor(customValue); fireEvent.change(screen.getByLabelText("Definition for tier 2"), { target: { value: "audits\nand reviews" } }); const next = committed(); expect(next.custom_tier_set?.tiers[1].definition).toBe("audits and reviews"); }); - it("moves a keyword rule with the tier it points at when that tier is renamed", () => { + it("moves a keyword rule with the tier it points at when that tier is renamed", async () => { const onKeywordTierRulesChange = vi.fn(); renderWithProviders( { expect(onKeywordTierRulesChange).toHaveBeenCalledWith([{ id: "r1", keywords: ["audit"], tier: "AUDIT" }]); }); - it("re-points the fallback tier when the row it named is removed, never leaving it dangling", () => { + it("re-points the fallback tier when the row it named is removed, never leaving it dangling", async () => { const threeRows: ComplexityRouterConfigValue = { ...customValue, custom_tier_set: { @@ -1702,7 +1718,7 @@ describe("ComplexityRouterConfig tier editing", () => { expect(next.custom_tier_set?.tiers.some((row) => row.id === next.custom_tier_set?.fallback_tier_id)).toBe(true); }); - it("turns off a plan-mode floor whose row was removed, rather than leaving it pointing at nothing", () => { + it("turns off a plan-mode floor whose row was removed, rather than leaving it pointing at nothing", async () => { const withFloor: ComplexityRouterConfigValue = { ...customValue, plan_mode_min_tier: "sec", @@ -1719,14 +1735,14 @@ describe("ComplexityRouterConfig tier editing", () => { expect(committed().plan_mode_min_tier).toBeUndefined(); }); - it("replaces the display-name inputs with the reason an edited tier set forbids them", () => { + it("replaces the display-name inputs with the reason an edited tier set forbids them", async () => { renderWithProviders(); expect(screen.queryByLabelText("Display name for the Simple tier")).not.toBeInTheDocument(); expect(screen.getByText("Display names rename the built-in tiers", { exact: false })).toBeInTheDocument(); expect(screen.getByLabelText("Fallback tier")).toBeInTheDocument(); }); - it("disables the once-per-session frequency and says why, rather than letting a stripped value look saved", () => { + it("disables the once-per-session frequency and says why, rather than letting a stripped value look saved", async () => { renderWithProviders( { onEditingTiersChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); - const sessionOption = screen.getByRole("radio", { name: /Once per session/ }); + openAutoRouterAdvanced("Classification Method"); + fireEvent.click(screen.getByRole("combobox", { name: "How often to classify" })); + const sessionOption = screen.getByRole("option", { name: "Once per session" }); expect(sessionOption).toHaveAttribute("aria-disabled", "true"); expect(sessionOption).not.toBeChecked(); expect( @@ -1743,15 +1760,15 @@ describe("ComplexityRouterConfig tier editing", () => { ).toBeInTheDocument(); }); - it("lets an edited tier set write its own opening instructions instead of refusing a prompt outright", () => { + it("lets an edited tier set write its own opening instructions instead of refusing a prompt outright", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("your own calibration examples", { exact: false })).toBeInTheDocument(); expect(screen.getByRole("button", { name: "Customize prompt" })).toBeInTheDocument(); expect(screen.queryByText("A replacement prompt drops the tier bullets", { exact: false })).not.toBeInTheDocument(); }); - it("gives built-in routers the opening-only editor, keeping the tier definitions derived", () => { + it("gives built-in routers the opening-only editor, keeping the tier definitions derived", async () => { renderWithProviders( { onEditingTiersChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByText("The base rubric supplies", { exact: false })).toBeInTheDocument(); expect(screen.getByRole("button", { name: "Customize prompt" })).toBeInTheDocument(); expect(screen.queryByText("Replace the built-in complexity rubric", { exact: false })).not.toBeInTheDocument(); }); - it("keeps the legacy whole-prompt editor only on a router that already stored a replacement prompt", () => { + it("keeps the legacy whole-prompt editor only on a router that already stored a replacement prompt", async () => { renderWithProviders( { onEditingTiersChange={vi.fn()} />, ); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.getByRole("button", { name: "Edit custom prompt" })).toBeInTheDocument(); expect(screen.queryByRole("button", { name: "Customize prompt" })).not.toBeInTheDocument(); }); - it("leaves built-in routers with their display-name inputs and no restriction copy", () => { + it("leaves built-in routers with their display-name inputs and no restriction copy", async () => { renderWithProviders(); expect(screen.getByLabelText("Display name for the Simple tier")).toBeInTheDocument(); expect(screen.queryByText("Display names rename the built-in tiers", { exact: false })).not.toBeInTheDocument(); @@ -1810,9 +1827,9 @@ describe("classifier vision settings", () => { ); }; - it("starts off and reveals the default cap when enabled", () => { + it("starts off and reveals the default cap when enabled", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); const vision = screen.getByRole("switch", { name: "Use images for classification" }); expect(vision).not.toBeChecked(); @@ -1823,10 +1840,10 @@ describe("classifier vision settings", () => { expect(screen.getByLabelText("Maximum images per request")).toHaveValue("1"); }); - it("writes the switch and a clamped image cap into the classifier config", () => { + it("writes the switch and a clamped image cap into the classifier config", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); fireEvent.click(screen.getByRole("switch", { name: "Use images for classification" })); expect(onChange).toHaveBeenLastCalledWith({ @@ -1841,10 +1858,10 @@ describe("classifier vision settings", () => { }); }); - it("keeps the image cap draft empty until a valid value is entered", () => { + it("keeps the image cap draft empty until a valid value is entered", async () => { const onChange = vi.fn(); renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); fireEvent.click(screen.getByRole("switch", { name: "Use images for classification" })); onChange.mockClear(); @@ -1861,9 +1878,9 @@ describe("classifier vision settings", () => { }); }); - it("is absent when the classifier is heuristic", () => { + it("is absent when the classifier is heuristic", async () => { renderWithProviders(); - fireEvent.click(screen.getByText("Advanced: Classification Method")); + openAutoRouterAdvanced("Classification Method"); expect(screen.queryByText("Use images for classification")).not.toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index acf6b62a95a..32f9ebf97ad 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -1,4 +1,6 @@ import RoutingOptions from "./RoutingOptions"; +import ClassifierPrimarySettings from "./ClassifierPrimarySettings"; +import { AutoRouterAllowanceNote } from "./AutoRouterAvailability"; import type { JevClassifierConfig } from "./jev_classifier_config"; import { type ClassifierType } from "./classifier_types"; export { type ClassifierType, usesLlmClassifier, usesClassifierContext } from "./classifier_types"; @@ -6,11 +8,11 @@ import ForecastClassifierConfig, { ForecastSolverModels } from "./ForecastClassi import { isForecastClassifier, type CapabilitySettings, type FuseSettings } from "./forecast_classifier_config"; import { SimpleTooltip } from "@/components/ui/tooltip"; import { MultiSelect } from "@/components/shared/MultiSelect"; +import TierConfigIntro from "./TierConfigIntro"; import DefaultModelField from "./DefaultModelField"; import { Info, Plus, Trash2, X } from "lucide-react"; import NonReasoningTierToggle from "./NonReasoningTierToggle"; -import TierConfigIntro from "./TierConfigIntro"; import TierRowSelect from "./TierRowSelect"; import { Card, CardContent } from "@/components/ui/card"; import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; @@ -226,10 +228,16 @@ const TierSetToolbar: React.FC<{ ) )} + {editing && ( + + )} {editing && ( Add or remove tiers to define your own set. Every custom tier needs a definition the classifier routes on, and - an edited set requires the LLM or JEV classification method + an edited set requires the LLM or Jev classification method )} {editing && keywordRulesError && ( @@ -596,10 +604,14 @@ const ComplexityRouterConfig: React.FC = ({ return (
+
-

- {forecast ? "Solver models" : "Complexity Tier Configuration"} -

+

{forecast ? "Solver models" : "Models by tier"}

{!forecast && ( @@ -619,6 +631,7 @@ const ComplexityRouterConfig: React.FC = ({ fastModeByModel={fastModeByModel} /> = ({ ) : ( <> - {!customTierSet && ( @@ -750,13 +762,19 @@ const ComplexityRouterConfig: React.FC = ({ )} - {!forecast && } + - + {forecast && ( <> - void>(); renderWithProviders(); - await user.click(screen.getByRole("button", { name: "Advanced routing options" })); + openAutoRouterAdvanced("Keyword/Semantic Matching"); expect(screen.getByRole("switch", { name: "Fast mode for secondary in the Medium routing pool tier" })).toBeChecked(); expect(onChange).not.toHaveBeenCalled(); await user.click(screen.getByRole("combobox", { name: "Select medium routing pool models" })); @@ -245,7 +246,7 @@ it.each(["capability", "llm_v2"] as const)( ); const view = renderWithProviders(editor(hydrateComplexityRouterConfig(stored, undefined))); - await user.click(screen.getByRole("button", { name: "Advanced routing options" })); + openAutoRouterAdvanced("Keyword/Semantic Matching"); const select = () => screen.getByRole("combobox", { name: "Default model" }); expect(select()).toHaveValue("legacy-default"); expect(onChange).not.toHaveBeenCalled(); @@ -284,8 +285,8 @@ it.each(["capability", "llm_v2"] as const)("offers only populated keyword target /> ); const view = renderWithProviders(editor([])); - await user.click(screen.getByRole("button", { name: "Advanced routing options" })); - await user.click(screen.getByText("Advanced: Keyword/Semantic Matching")); + openAutoRouterAdvanced("Keyword/Semantic Matching"); + openAutoRouterAdvanced("Keyword/Semantic Matching"); await user.click(screen.getByRole("button", { name: "Add keyword rule" })); const rules = onRulesChange.mock.lastCall![0]; expect(rules[0].tier).toBe("SIMPLE"); diff --git a/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.integration.test.tsx index a03ccb11456..f243e63d387 100644 --- a/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.integration.test.tsx @@ -1,3 +1,4 @@ +import { selectAutoRouterApproach } from "../../../tests/autoRouterSetup"; import React, { useState } from "react"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import userEvent from "@testing-library/user-event"; @@ -306,7 +307,7 @@ describe("forecast classifier form", () => { ); }); - it("switches a populated standard router to Capability without saving hidden pools or their overrides", () => { + it("switches a populated standard router to Capability without saving hidden pools or their overrides", async () => { renderWithProviders( { }} />, ); - fireEvent.click(screen.getByRole("tab", { name: "Capability" })); + await selectAutoRouterApproach("Capability"); fireEvent.change(screen.getByLabelText("Solve probability threshold"), { target: { value: "0.7" } }); expect(screen.getByRole("button", { name: "Save configuration" })).toBeEnabled(); fireEvent.click(screen.getByRole("button", { name: "Save configuration" })); @@ -347,7 +348,7 @@ describe("forecast classifier form", () => { it.each(["capability", "llm_v2"] as const)( "carries non-default solver assignments when switching away from %s", - (source) => { + async (source) => { const pair = { efficient_tier: "MEDIUM", capable_tier: "COMPLEX" }; const previous: ComplexityRouterConfigValue = { ...(source === "capability" ? initial : fuseInitial), @@ -359,7 +360,7 @@ describe("forecast classifier form", () => { tier_model_params: { MEDIUM: { efficient: { max_tokens: 128 } }, COMPLEX: { capable: { speed: "fast" } } }, }; renderWithProviders(); - fireEvent.click(screen.getByRole("tab", { name: source === "capability" ? "Fuse v2" : "Capability" })); + await selectAutoRouterApproach(source === "capability" ? "Fuse v2" : "Capability"); if (source === "capability") { fireEvent.change(screen.getByLabelText("Efficient solver profile"), { target: { value: "Small solver" } }); fireEvent.change(screen.getByLabelText("Capable solver profile"), { target: { value: "Large solver" } }); @@ -402,20 +403,22 @@ describe("forecast classifier form", () => { ] as const)("restores the current rubric when switching %s through Complexity to %s", async (source, target) => { const user = userEvent.setup(); renderWithProviders(); - fireEvent.click(screen.getByRole("tab", { name: "Complexity" })); + await selectAutoRouterApproach("Complexity"); fireEvent.click(screen.getByRole("radio", { name: new RegExp(`^${target}`) })); - await user.click(screen.getByRole("combobox", { name: "Classifier Model" })); + await user.click(screen.getByRole("combobox", { name: "Judge model" })); await user.click(screen.getByRole("option", { name: "judge" })); fireEvent.click(screen.getByRole("button", { name: "Save configuration" })); const output = screen.getByRole("status", { name: "Saved configuration" }); expect(output).toHaveTextContent('"classification_rubric":"agentic"'); expect(output).toHaveTextContent('"model":"judge"'); - expect(output).toHaveTextContent('"timeout_ms":3000'); + expect(output).toHaveTextContent( + `"timeout_ms":${(source === "capability" ? initial : fuseInitial).classifier_llm_config?.timeout_ms}`, + ); expect(output).not.toHaveTextContent('"capability_classifier_config"'); expect(output).not.toHaveTextContent('"llm_v2_config"'); }); - it("saves capability threshold edits together with fitted calibration", () => { + it("saves capability threshold edits together with fitted calibration", async () => { renderWithProviders(); fireEvent.change(screen.getByLabelText("Solve probability threshold"), { target: { value: "0.6" } }); fireEvent.click(screen.getByRole("button", { name: "Classifier options" })); @@ -432,9 +435,9 @@ describe("forecast classifier form", () => { expect(screen.getByRole("button", { name: "Save configuration" })).toBeDisabled(); }); - it("switches to Fuse, requires solver context, and saves the filled fields", () => { + it("switches to Fuse, requires solver context, and saves the filled fields", async () => { renderWithProviders(); - fireEvent.click(screen.getByRole("tab", { name: "Fuse v2" })); + await selectAutoRouterApproach("Fuse v2"); expect(screen.queryByLabelText("Solve probability threshold")).not.toBeInTheDocument(); expect(screen.getByRole("button", { name: "Save configuration" })).toBeDisabled(); fireEvent.change(screen.getByLabelText("Efficient solver profile"), { diff --git a/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.tsx index c336423fe39..d759b870209 100644 --- a/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ForecastClassifierConfig.tsx @@ -36,6 +36,7 @@ interface Props { onChange: (value: ComplexityRouterConfigValue) => void; modelOptions: { value: string; label: string }[]; effortOptionsByModel: Record; + section?: "all" | "required" | "advanced"; } const NumberField = ({ @@ -188,7 +189,7 @@ const CalibrationFields = ({ const emptyCoefficients = () => ({ slope: Number.NaN, intercept: Number.NaN }); -const ForecastClassifierConfig = ({ value, onChange, modelOptions, effortOptionsByModel }: Props) => { +const ForecastClassifierConfig = ({ value, onChange, modelOptions, effortOptionsByModel, section = "all" }: Props) => { const id = React.useId(); const isCapability = value.classifier_type === "capability"; const capability = value.capability_classifier_config ?? newCapabilitySettings(); @@ -212,72 +213,195 @@ const ForecastClassifierConfig = ({ value, onChange, modelOptions, effortOptions ? "Forecasts whether the efficient solver can complete the task using the bundled capability card" : "Forecasts success for both solvers and selects efficient when the estimated quality gap is within your allowance"}

-
- - { - if (model === llm.model) return; - onChange({ ...value, classifier_llm_config: { ...llm, model: model ?? "", reasoning_effort: undefined } }); - }} - /> -
- {isCapability ? ( + {section === "all" && ( <> - updateCapability({ ...capability, base_threshold })} - /> - - ) : ( - <> - - updateFuse({ ...fuse, max_quality_gap })} - /> +
+ + { + if (model === llm.model) return; + onChange({ + ...value, + classifier_llm_config: { ...llm, model: model ?? "", reasoning_effort: undefined }, + }); + }} + /> +
)} - - - - Classifier options - - - onChange({ ...value, classifier_llm_config: { ...llm, reasoning_effort } })} - /> - onChange({ ...value, classifier_llm_config: { ...llm, timeout_ms } })} - /> - onChange({ ...value, classifier_llm_config })} - /> - onChange({ ...value, classifier_llm_config })} - /> + {section !== "advanced" && ( + <> + {isCapability ? ( + <> + updateCapability({ ...capability, base_threshold })} + /> + + ) : ( + <> + + updateFuse({ ...fuse, max_quality_gap })} + /> + + )} + + )} + {section !== "required" && ( + + {section === "all" && ( + + + Classifier options + + )} + + + onChange({ ...value, classifier_llm_config: { ...llm, reasoning_effort } }) + } + /> + onChange({ ...value, classifier_llm_config: { ...llm, timeout_ms } })} + /> + onChange({ ...value, classifier_llm_config })} + /> + onChange({ ...value, classifier_llm_config })} + /> + {isCapability && ( + updateCapability({ ...capability, threshold_step })} + /> + )} + updateTransport({ max_output_tokens })} + /> +
+ + { + if (response_format === "json_schema" || response_format === "json_object") + updateTransport({ response_format }); + }} + /> +
+
+ +

+ Optional coefficients fitted for your judge, solvers, and harness. Leave off to use raw forecasts +

+ {config.calibration && ( +
+ + setCalibrationVersion(event.target.value)} + /> +
+ )} + {isCapability && capability.calibration && ( + + updateCapability({ + ...capability, + calibration: { version: capability.calibration?.version ?? "", ...next }, + }) + } + /> + )} + {!isCapability && + fuse.calibration && + (["efficient", "capable"] as const).map((role) => ( + { + if (fuse.calibration) updateFuse({ ...fuse, calibration: { ...fuse.calibration, [role]: next } }); + }} + /> + ))} +
+
+
+ )} + {section === "all" && ( + <>
- {isCapability && ( - updateCapability({ ...capability, threshold_step })} - /> - )} - updateTransport({ max_output_tokens })} - /> -
- - { - if (response_format === "json_schema" || response_format === "json_object") - updateTransport({ response_format }); - }} - /> -
-
- -

- Optional coefficients fitted for your judge, solvers, and harness. Leave off to use raw forecasts -

- {config.calibration && ( -
- - setCalibrationVersion(event.target.value)} - /> -
- )} - {isCapability && capability.calibration && ( - - updateCapability({ - ...capability, - calibration: { version: capability.calibration?.version ?? "", ...next }, - }) - } - /> - )} - {!isCapability && - fuse.calibration && - (["efficient", "capable"] as const).map((role) => ( - { - if (fuse.calibration) updateFuse({ ...fuse, calibration: { ...fuse.calibration, [role]: next } }); - }} - /> - ))} -
-
-
+ + )}

The classifier uses its bundled prompt and always falls back to the capable solver

- {error && ( + {section !== "advanced" && error && (

{error}

diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx index 896fde3a446..7da8b12c2d7 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.integration.test.tsx @@ -98,28 +98,28 @@ describe("JEV classifier editor", () => { afterEach(() => vi.mocked(useAuthorized).mockReset()); it("uses built-in JEV without a license and preserves custom tiers and context through reload", () => { renderWithProviders(); - expect(screen.getByLabelText("Classifier Model")).toBeInTheDocument(); + expect(screen.getByLabelText("Judge model")).toBeInTheDocument(); expect(screen.getByText("Reasoning Effort")).toBeInTheDocument(); expect(screen.getByText("Classifier Prompt")).toBeInTheDocument(); expect(screen.getByRole("switch", { name: "Use images for classification" })).toBeInTheDocument(); - fireEvent.click(screen.getByRole("radio", { name: /JEV Classifier/ })); - expect(screen.getByRole("tab", { name: "Complexity" })).toHaveAttribute("aria-selected", "true"); - expect(screen.getByLabelText("JEV Model")).toHaveValue("jev-latest"); - expect(screen.getByLabelText("JEV Instructions")).toBeDisabled(); - expect(screen.queryByLabelText("Classifier Model")).not.toBeInTheDocument(); + fireEvent.click(screen.getByRole("radio", { name: /Jev Classifier/ })); + expect(screen.getByRole("radio", { name: /^Jev Classifier/ })).toBeChecked(); + expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-latest"); + expect(screen.getByLabelText("Jev Instructions")).toBeEnabled(); + expect(screen.queryByLabelText("Judge model")).not.toBeInTheDocument(); expect(screen.queryByText("Reasoning Effort")).not.toBeInTheDocument(); expect(screen.queryByText("Classifier Prompt")).not.toBeInTheDocument(); expect(screen.queryByRole("switch", { name: "Use images for classification" })).not.toBeInTheDocument(); - fireEvent.change(screen.getByLabelText("JEV Model"), { target: { value: "jev-test" } }); - fireEvent.change(screen.getByLabelText("JEV Timeout (ms)"), { target: { value: "4200" } }); + fireEvent.change(screen.getByLabelText("Jev Model"), { target: { value: "jev-test" } }); + fireEvent.change(screen.getByLabelText("Jev Timeout (ms)"), { target: { value: "4200" } }); fireEvent.change(screen.getByLabelText("Context Window Size"), { target: { value: "6" } }); fireEvent.change(screen.getByLabelText("Circuit breaker cooldown (seconds)"), { target: { value: "50" } }); fireEvent.click(screen.getByRole("switch", { name: "Classifier circuit breaker" })); fireEvent.click(screen.getByRole("button", { name: "Customize tiers" })); fireEvent.click(screen.getByRole("button", { name: "Save and reload" })); - expect(screen.getByRole("radio", { name: /JEV Classifier/ })).toBeChecked(); - expect(screen.getByLabelText("JEV Model")).toHaveValue("jev-test"); - expect(screen.getByLabelText("JEV Timeout (ms)")).toHaveValue(4200); + expect(screen.getByRole("radio", { name: /Jev Classifier/ })).toBeChecked(); + expect(screen.getByLabelText("Jev Model")).toHaveValue("jev-test"); + expect(screen.getByLabelText("Jev Timeout (ms)")).toHaveValue(4200); expect(screen.getByLabelText("Context Window Size")).toHaveValue("6"); expect(screen.getByRole("switch", { name: "Classifier circuit breaker" })).not.toBeChecked(); fireEvent.click(screen.getByRole("button", { name: "Probe current config" })); @@ -152,10 +152,10 @@ describe("JEV classifier editor", () => { return ; }; renderWithProviders(); - expect(screen.getByLabelText("JEV Instructions")).toBeEnabled(); - fireEvent.change(screen.getByLabelText("JEV Instructions"), { target: { value: "New instructions" } }); - expect(screen.getByLabelText("JEV Instructions")).toHaveValue("New instructions"); - fireEvent.click(screen.getByRole("button", { name: "Restore built-in JEV instructions" })); - expect(screen.getByLabelText("JEV Instructions")).toHaveValue(""); + expect(screen.getByLabelText("Jev Instructions")).toBeEnabled(); + fireEvent.change(screen.getByLabelText("Jev Instructions"), { target: { value: "New instructions" } }); + expect(screen.getByLabelText("Jev Instructions")).toHaveValue("New instructions"); + fireEvent.click(screen.getByRole("button", { name: "Restore built-in Jev instructions" })); + expect(screen.getByLabelText("Jev Instructions")).toHaveValue(""); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx index 25286eaef07..97609bcd8c3 100644 --- a/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/JevClassifierConfig.tsx @@ -1,10 +1,9 @@ import React, { useId } from "react"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { AutoRouterAllowanceNote } from "./AutoRouterAvailability"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; import { Textarea } from "@/components/ui/textarea"; -import { SimpleTooltip } from "@/components/ui/tooltip"; import ClassifierCircuitBreakerConfig from "./ClassifierCircuitBreakerConfig"; import type { ComplexityRouterConfigValue } from "./ComplexityRouterConfig"; import { defaultJevClassifierConfig } from "./jev_classifier_config"; @@ -17,7 +16,6 @@ export default function JevClassifierConfig({ onChange: (value: ComplexityRouterConfigValue) => void; }) { const id = useId(); - const { premiumUser } = useAuthorized(); const config = value.jev_classifier_config ?? defaultJevClassifierConfig(); const update = (patch: Partial) => onChange({ ...value, jev_classifier_config: { ...config, ...patch } }); @@ -28,11 +26,11 @@ export default function JevClassifierConfig({ Uses TypeSafe System One Choice evaluation with your configured tiers

- + update({ model: event.target.value })} />
- +
- - -
-