From 8d166258a65ba272546c5e62c3aac79cc7831ae3 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 15:49:39 -0700 Subject: [PATCH 01/51] fix(tests): stop VCR recording and replaying a test's own localhost upstream (#43346) * fix(tests): stop VCR recording and replaying a test's own localhost upstream * test(vcr): prove a localhost response an earlier run stored is never replayed * test(vcr): drive the localhost cassette checks in-process instead of through a loopback server --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- tests/_vcr_conftest_common.py | 1 + tests/llm_translation/Readme.md | 5 ++ tests/unit/test_vcr_safe_body_matcher.py | 69 ++++++++++++++++++++++++ 3 files changed, 75 insertions(+) diff --git a/tests/_vcr_conftest_common.py b/tests/_vcr_conftest_common.py index 3adc671021b..36ae70497e7 100644 --- a/tests/_vcr_conftest_common.py +++ b/tests/_vcr_conftest_common.py @@ -1090,6 +1090,7 @@ def vcr_config_dict() -> dict: "decode_compressed_response": True, "record_mode": "new_episodes", "allow_playback_repeats": True, + "ignore_localhost": True, "match_on": ( "method", "scheme", diff --git a/tests/llm_translation/Readme.md b/tests/llm_translation/Readme.md index 813c188ee7b..f0a32f6c989 100644 --- a/tests/llm_translation/Readme.md +++ b/tests/llm_translation/Readme.md @@ -16,6 +16,11 @@ The persister, header scrubbing, and 2xx-only filtering are defined in patches the same httpx transport vcrpy does) are excluded from the auto-marker — see `_RESPX_CONFLICTING_FILES` in `conftest.py`. +Requests to `localhost`, `127.0.0.1`, or `0.0.0.0` are never recorded or +replayed (`ignore_localhost` in `vcr_config_dict()`): a server the test +process starts itself on an ephemeral port is not a provider, and a cassette +entry for it would replay against whichever later test lands on that port + The same VCR cache is used by other test directories that exercise live provider APIs. The reusable conftest plumbing lives in `tests/_vcr_conftest_common.py` and is wired into: diff --git a/tests/unit/test_vcr_safe_body_matcher.py b/tests/unit/test_vcr_safe_body_matcher.py index cf4e4a1c276..71ae97e69d9 100644 --- a/tests/unit/test_vcr_safe_body_matcher.py +++ b/tests/unit/test_vcr_safe_body_matcher.py @@ -1,10 +1,15 @@ from __future__ import annotations +import json import os import sys +from pathlib import Path from types import SimpleNamespace +from typing import Final import pytest +import vcr +from vcr.request import Request _REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")) if _REPO_ROOT not in sys.path: @@ -384,3 +389,67 @@ def test_before_record_request_is_idempotent_on_the_same_request_object(): _before_record_request(req) assert req.headers[KEY_FINGERPRINT_HEADER] == fp_after_first assert fp_after_first != "no-key" + + +LOCAL_UPSTREAM: Final = "http://127.0.0.1:54321/v1/moderations" +REMOTE_UPSTREAM: Final = "https://api.openai.com/v1/moderations" + + +def _recorder_with_repo_matchers(cassette_dir: Path) -> vcr.VCR: + recorder: Final = vcr.VCR(cassette_library_dir=str(cassette_dir)) + recorder.register_matcher(SAFE_BODY_MATCHER_NAME, _safe_body_matcher) + recorder.register_matcher(KEY_FINGERPRINT_MATCHER_NAME, _key_fingerprint_matcher) + recorder.register_matcher(TOLERANT_QUERY_MATCHER_NAME, _tolerant_query_matcher) + recorder.register_matcher(TOLERANT_PATH_MATCHER_NAME, _tolerant_path_matcher) + return recorder + + +def _request_to(uri: str) -> Request: + return Request( + method="POST", + uri=uri, + body=b'{"model":"omni-moderation-latest","input":"hi"}', + headers={"content-type": "application/json"}, + ) + + +def _response_served_by(server: str) -> dict[str, object]: + payload: Final = json.dumps({"served_by": server}).encode() + return { + "status": {"code": 200, "message": "OK"}, + "headers": {"content-type": ["application/json"]}, + "body": {"string": payload}, + } + + +def _stored_uris(session: vcr.cassette.Cassette) -> list[str]: + return [request.uri for request in session.requests] + + +def test_config_never_records_a_test_owned_local_upstream(tmp_path: Path): + recorder: Final = _recorder_with_repo_matchers(tmp_path) + + with recorder.use_cassette("local_upstream.yaml", **vcr_config_dict()) as session: + session.append(_request_to(LOCAL_UPSTREAM), _response_served_by("the test's own server")) + session.append(_request_to(REMOTE_UPSTREAM), _response_served_by("a real provider")) + + assert _stored_uris(session) == [REMOTE_UPSTREAM] + assert (tmp_path / "local_upstream.yaml").exists() + + +def test_config_never_replays_a_localhost_response_an_earlier_run_stored(tmp_path: Path): + recorder: Final = _recorder_with_repo_matchers(tmp_path) + config_that_recorded_localhost: Final = vcr_config_dict() | {"ignore_localhost": False} + + with recorder.use_cassette("stored_by_an_earlier_run.yaml", **config_that_recorded_localhost) as earlier_run: + earlier_run.append(_request_to(LOCAL_UPSTREAM), _response_served_by("an earlier run's server")) + earlier_run.append(_request_to(REMOTE_UPSTREAM), _response_served_by("a real provider")) + assert _stored_uris(earlier_run) == [LOCAL_UPSTREAM, REMOTE_UPSTREAM] + + with recorder.use_cassette("stored_by_an_earlier_run.yaml", **vcr_config_dict()) as session: + replayable: Final = tuple( + bool(session.can_play_response_for(_request_to(uri))) for uri in (LOCAL_UPSTREAM, REMOTE_UPSTREAM) + ) + + assert replayable == (False, True) + assert _stored_uris(session) == [REMOTE_UPSTREAM] From f12f7b5a037ab5357643ed9e56a95cc36ba0b0b5 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 16:00:32 -0700 Subject: [PATCH 02/51] test(integration): group /v1/messages contracts under tests/integration/messages_endpoint (#43352) * test(integration): group /v1/messages contracts under tests/integration/messages Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): make ci coverage census collect nested test dirs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): nest /v1/messages contracts under messages_endpoint/providers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/scripts/assert_ci_coverage.py | 2 +- tests/integration/README.md | 2 ++ tests/integration/_support/manifest.py | 1 + .../test_anthropic_messages_fireworks_stop_wire.py | 0 .../providers/anthropic}/test_anthropic_advisor_wire.py | 0 .../anthropic}/test_anthropic_legacy_thinking_budget_wire.py | 0 .../anthropic}/test_anthropic_messages_timeout_wire.py | 0 .../test_anthropic_thinking_signature_retry_wire.py | 0 .../providers/anthropic}/test_anthropic_wire.py | 0 .../providers/anthropic}/test_websearch_interception_wire.py | 0 .../bedrock}/test_bedrock_invoke_tool_search_wire.py | 0 .../bedrock}/test_bedrock_messages_web_search_replay_wire.py | 0 .../gemini}/test_gemini_messages_cache_control_wire.py | 0 .../test_anthropic_messages_claude_code_cache_key_wire.py | 0 .../test_anthropic_messages_openai_bridge_wire.py | 0 .../test_anthropic_messages_openai_tools_wire.py | 0 .../responses_bridge}/test_responses_bridge_stream_options.py | 0 tests/integration/run.py | 4 ++-- 18 files changed, 6 insertions(+), 3 deletions(-) rename tests/integration/{providers => messages_endpoint/chat_bridge}/test_anthropic_messages_fireworks_stop_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_advisor_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_legacy_thinking_budget_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_messages_timeout_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_thinking_signature_retry_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_anthropic_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/anthropic}/test_websearch_interception_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/bedrock}/test_bedrock_invoke_tool_search_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/bedrock}/test_bedrock_messages_web_search_replay_wire.py (100%) rename tests/integration/{providers => messages_endpoint/providers/gemini}/test_gemini_messages_cache_control_wire.py (100%) rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_anthropic_messages_claude_code_cache_key_wire.py (100%) rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_anthropic_messages_openai_bridge_wire.py (100%) rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_anthropic_messages_openai_tools_wire.py (100%) rename tests/integration/{providers => messages_endpoint/responses_bridge}/test_responses_bridge_stream_options.py (100%) diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py index 01a01b1034b..a483dcec9d7 100644 --- a/.github/scripts/assert_ci_coverage.py +++ b/.github/scripts/assert_ci_coverage.py @@ -516,7 +516,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens str(path.relative_to(repo_root)) for folders in groups.values() for folder in folders - for path in (integration_root / folder).glob("test_*.py") + for path in (integration_root / folder).rglob("test_*.py") ) browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json" browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else () diff --git a/tests/integration/README.md b/tests/integration/README.md index c09904597ab..c559e7545e0 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -28,6 +28,8 @@ Provider contracts exercise actual TCP requests with synthetic credentials and l Streaming checks send real HTTP transfer chunks, including one-byte partitions, fragmented tools, incomplete transfers and a cancellation barrier. They assert meaningful text, tool arguments, final usage and persisted cost. The Redis recovery case owns a separate database and Redis process, uses the supported one-second circuit-breaker recovery setting, waits for the real subscriber and verifies response data in Redis after restart. CircleCI reuses its existing Redis image for that extra process; it never pulls an image during tests +The `messages_endpoint/` directory holds `/v1/messages` endpoint contracts: native-provider backends under `providers/` (`anthropic`, `bedrock`, `gemini`) and the translation bridges (`responses_bridge`, `chat_bridge`) at the top level. It runs in the providers shard; `run.py` selects test files recursively under each scheduled directory + The sdk shard exercises the SDK's own HTTP clients against local protocol peers with no gateway in the path, so a case here fails only when the client library or its wire behavior changes. The HTTP/2 case runs a hypercorn TLS peer offering h2 and http/1.1 over ALPN, drives the sync and async httpx handlers at it with `LITELLM_HTTP2` off and on, and asserts the version both the client and the peer observed on the wire. Put a test here only when it needs no proxy, database or Redis; a case that reaches the gateway belongs in one of the other shards The extensions shard uses the built-in generic callback and guardrail transports. It checks callback correlation and credential exclusion, guardrail rewriting and denial, retained OpenAI consumers and A2A wire versions. CircleCI runs it on parallel nodes, and each node starts its own database, Redis, upstream and proxy and runs its share of the group's files serially, split by recorded timings with `circleci tests split`. Tests keep the isolation of a serial run; they still must not assume a particular set of sibling files. `run.py --list` prints a group's files and `run.py ...` runs a subset of them diff --git a/tests/integration/_support/manifest.py b/tests/integration/_support/manifest.py index aa0b27eceda..376c2a515b7 100644 --- a/tests/integration/_support/manifest.py +++ b/tests/integration/_support/manifest.py @@ -10,6 +10,7 @@ OWNED_DIRECTORIES: Final = frozenset( "routing", "providers", "streaming", + "messages_endpoint", "configuration", "mcp", "observability", diff --git a/tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py b/tests/integration/messages_endpoint/chat_bridge/test_anthropic_messages_fireworks_stop_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py rename to tests/integration/messages_endpoint/chat_bridge/test_anthropic_messages_fireworks_stop_wire.py diff --git a/tests/integration/providers/test_anthropic_advisor_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_advisor_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_advisor_wire.py rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_advisor_wire.py diff --git a/tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_legacy_thinking_budget_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_legacy_thinking_budget_wire.py diff --git a/tests/integration/providers/test_anthropic_messages_timeout_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_timeout_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_messages_timeout_wire.py rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_messages_timeout_wire.py diff --git a/tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_thinking_signature_retry_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_thinking_signature_retry_wire.py diff --git a/tests/integration/providers/test_anthropic_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_anthropic_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_wire.py rename to tests/integration/messages_endpoint/providers/anthropic/test_anthropic_wire.py diff --git a/tests/integration/providers/test_websearch_interception_wire.py b/tests/integration/messages_endpoint/providers/anthropic/test_websearch_interception_wire.py similarity index 100% rename from tests/integration/providers/test_websearch_interception_wire.py rename to tests/integration/messages_endpoint/providers/anthropic/test_websearch_interception_wire.py diff --git a/tests/integration/providers/test_bedrock_invoke_tool_search_wire.py b/tests/integration/messages_endpoint/providers/bedrock/test_bedrock_invoke_tool_search_wire.py similarity index 100% rename from tests/integration/providers/test_bedrock_invoke_tool_search_wire.py rename to tests/integration/messages_endpoint/providers/bedrock/test_bedrock_invoke_tool_search_wire.py diff --git a/tests/integration/providers/test_bedrock_messages_web_search_replay_wire.py b/tests/integration/messages_endpoint/providers/bedrock/test_bedrock_messages_web_search_replay_wire.py similarity index 100% rename from tests/integration/providers/test_bedrock_messages_web_search_replay_wire.py rename to tests/integration/messages_endpoint/providers/bedrock/test_bedrock_messages_web_search_replay_wire.py diff --git a/tests/integration/providers/test_gemini_messages_cache_control_wire.py b/tests/integration/messages_endpoint/providers/gemini/test_gemini_messages_cache_control_wire.py similarity index 100% rename from tests/integration/providers/test_gemini_messages_cache_control_wire.py rename to tests/integration/messages_endpoint/providers/gemini/test_gemini_messages_cache_control_wire.py diff --git a/tests/integration/providers/test_anthropic_messages_claude_code_cache_key_wire.py b/tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_claude_code_cache_key_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_messages_claude_code_cache_key_wire.py rename to tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_claude_code_cache_key_wire.py diff --git a/tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py b/tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_bridge_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py rename to tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_bridge_wire.py diff --git a/tests/integration/providers/test_anthropic_messages_openai_tools_wire.py b/tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_tools_wire.py similarity index 100% rename from tests/integration/providers/test_anthropic_messages_openai_tools_wire.py rename to tests/integration/messages_endpoint/responses_bridge/test_anthropic_messages_openai_tools_wire.py diff --git a/tests/integration/providers/test_responses_bridge_stream_options.py b/tests/integration/messages_endpoint/responses_bridge/test_responses_bridge_stream_options.py similarity index 100% rename from tests/integration/providers/test_responses_bridge_stream_options.py rename to tests/integration/messages_endpoint/responses_bridge/test_responses_bridge_stream_options.py diff --git a/tests/integration/run.py b/tests/integration/run.py index 87f93873267..8f1ff1f4a92 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -14,7 +14,7 @@ GROUPS: Final = MappingProxyType( "management": ("management", "authorization", "configuration"), "accounting": ("pricing", "spend"), "database": ("database",), - "providers": ("providers", "routing", "streaming"), + "providers": ("providers", "routing", "streaming", "messages_endpoint"), "extensions": ("observability", "compatibility"), "mcp": ("mcp",), "sdk": ("sdk",), @@ -37,7 +37,7 @@ def main() -> int: group_files: Final = tuple( str(path.relative_to(root)) for folder in GROUPS[options.group] - for path in sorted((root / "tests/integration" / folder).glob("test_*.py")) + for path in sorted((root / "tests/integration" / folder).rglob("test_*.py")) ) if options.list: print("\n".join(group_files)) From e4947231058394bcfaac14f6e3f0700d4be5644c Mon Sep 17 00:00:00 2001 From: shrey-berri Date: Sat, 26 Sep 2026 16:01:20 -0700 Subject: [PATCH 03/51] fix(params): validate stream_chunk_size once, before any provider call (#43222) * fix(params): validate stream_chunk_size once and carry it as typed control options Checks stream_chunk_size at the top of completion() and acompletion(), accepts digit strings, returns a 400 naming the param unless drop_params is set, and stores the checked value under _litellm_control. Bedrock Converse and Invoke read it from litellm_params; the Bedrock-only checker and the dead Invoke pops are gone. Owned-kwarg filtering now runs through one helper everywhere. Refs LIT-8317 * test(bedrock): drop tests for the removed stream_chunk_size_from helper Refs LIT-8317 * fix(params): check stream_chunk_size before the MCP gateway branch Refs LIT-8317 * fix(params): return assert_never in the exhaustive control-options match Refs LIT-8317 * fix(params): address council review of the control options change Read all_litellm_params live so names registered after import stay LiteLLM-owned, make litellm_params a required keyword on the stream wrapper hooks, give digit strings and ints the same 18-digit range, share the default-chunking test table, test the Responses bridge through litellm.responses, and revert formatting-only churn in existing tests. Refs LIT-8317 * fix(params): address the second council review of control options Keep the Responses bridge on its original all_litellm_params forwarding, narrow _int_from_decimal_string inline so it type-checks, bound nested huge ints in the error message, store _litellm_control only when a value is set, simplify the parser to its single field, drop the one-caller wrapper, and tighten the tests. Refs LIT-8317 * fix(params): keep the 18-digit length check on stream_chunk_size strings A 19-character string with leading zeros such as 0000000000000000001 would otherwise pass as 1, although the rule and the error message say at most 18 digits. Refs LIT-8317 * test(params): tidy control options tests after council sign-off Move the Responses bridge test into the existing bridge test file, drop the rebind test that pinned an implementation detail, assert through stored_control_options instead of the storage key, and cover drop_params="true" through Bedrock streaming. Refs LIT-8317 * test(params): wrap a chunking test row that went past 120 characters Refs LIT-8317 --- litellm/caching/caching.py | 5 +- litellm/constants.py | 1 + litellm/images/main.py | 11 +- .../litellm_core_utils/get_litellm_params.py | 53 +++- litellm/llms/base_llm/chat/transformation.py | 4 + .../bedrock/chat/agentcore/transformation.py | 6 +- litellm/llms/bedrock/chat/converse_handler.py | 5 +- .../anthropic_claude3_transformation.py | 1 - .../base_invoke_transformation.py | 13 +- litellm/llms/bedrock/common_utils.py | 11 +- litellm/llms/bytez/chat/transformation.py | 5 + litellm/llms/custom_httpx/llm_http_handler.py | 2 + litellm/llms/langgraph/chat/transformation.py | 5 + litellm/llms/oci/chat/transformation.py | 6 +- litellm/llms/sagemaker/chat/transformation.py | 5 + .../vertex_ai/agent_engine/transformation.py | 5 + litellm/main.py | 37 ++- litellm/types/litellm_params.py | 28 ++- litellm/utils.py | 24 +- tests/_support/stream_chunk_size.py | 34 +-- .../providers/test_internal_params_wire.py | 9 +- tests/unit/caching/test_caching.py | 14 ++ .../test_get_litellm_params.py | 84 ++++++- .../test_base_invoke_transformation.py | 169 +++++++------ tests/unit/llms/bedrock/test_common_utils.py | 20 -- tests/unit/llms/chat/test_converse_handler.py | 135 ++++------ .../oci/chat/test_oci_chat_transformation.py | 2 + .../unit/llms/oci/test_oci_coverage_boost.py | 2 + .../test_sagemaker_chat_transformation.py | 3 + .../test_sagemaker_nova_transformation.py | 2 + .../test_responses_api_bridge_flag.py | 18 ++ tests/unit/test_filter_out_litellm_params.py | 20 ++ tests/unit/test_main.py | 233 +++++++++++++++++- tests/unit/types/test_litellm_params.py | 10 +- 34 files changed, 695 insertions(+), 287 deletions(-) delete mode 100644 tests/unit/llms/bedrock/test_common_utils.py diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index a31cad4af29..d766d1a58bc 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -25,7 +25,7 @@ from litellm._logging import verbose_logger from litellm.constants import CACHED_STREAMING_CHUNK_DELAY from litellm.litellm_core_utils.model_param_helper import ModelParamHelper from litellm.types.caching import * -from litellm.types.utils import EmbeddingResponse, all_litellm_params +from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg from .azure_blob_cache import AzureBlobCache from .base_cache import BaseCache @@ -377,7 +377,6 @@ class Cache: return preset_cache_key combined_kwargs: Final = ModelParamHelper._get_all_llm_api_params() - litellm_param_kwargs: Final = all_litellm_params is_semantic_cache: Final = self._is_semantic_cache() scope_excluded_params: Final = self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS if is_semantic_cache else frozenset() for param in kwargs: @@ -387,7 +386,7 @@ class Cache: param_value: str | None = self._get_param_value(param, kwargs) if param_value is not None: cache_key += f"{param}: {param_value}" - elif param not in litellm_param_kwargs: # check if user passed in optional param - e.g. top_k + elif not is_litellm_owned_kwarg(param): if litellm.enable_caching_on_provider_specific_optional_params is True: # feature flagged for now if kwargs[param] is None: continue # ignore None params diff --git a/litellm/constants.py b/litellm/constants.py index dac15c01fbf..a292b654778 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1611,6 +1611,7 @@ ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: Final = { # Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.) PASS_THROUGH_HEADER_PREFIX: Final = "x-pass-" INTERNAL_KWARG_PREFIX: Final = "_litellm_" +CONTROL_OPTIONS_KEY: Final = f"{INTERNAL_KWARG_PREFIX}control" AZURE_SPEECH_CUSTOM_LLM_PROVIDER: Final = "azure_speech" AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech" diff --git a/litellm/images/main.py b/litellm/images/main.py index 5ca8a726a69..7dc68dafecc 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -25,7 +25,7 @@ from litellm.llms.base_llm import BaseImageEditConfig, BaseImageGenerationConfig from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.llms.custom_llm import CustomLLM -from litellm.utils import exception_type, get_litellm_params +from litellm.utils import exception_type, filter_out_litellm_params, get_litellm_params #################### Initialize provider clients #################### llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler() @@ -52,7 +52,6 @@ from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( LITELLM_IMAGE_VARIATION_PROVIDERS, LlmProviders, - is_litellm_owned_kwarg, ) from litellm.utils import ( ImageResponse, @@ -249,9 +248,7 @@ def image_generation( "size", "style", ] - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k) - } + non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params) image_generation_config: BaseImageGenerationConfig | None = None if custom_llm_provider is not None and custom_llm_provider in LlmProviders._member_map_.values(): @@ -755,9 +752,7 @@ def image_edit( "style", "async_call", ] - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k) - } + non_default_params: Final = filter_out_litellm_params(kwargs, excluding=openai_params) litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj") litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) model_info: Final = kwargs.get("model_info", None) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index f7aaef3a51f..f28259a1b7f 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -1,9 +1,15 @@ +import reprlib from collections.abc import Mapping, MutableMapping +from dataclasses import dataclass, fields from types import MappingProxyType from typing import Final +from pydantic import TypeAdapter, ValidationError + +from litellm.constants import CONTROL_OPTIONS_KEY from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.llms.openai.data_residency import infer_openai_data_residency +from litellm.types.litellm_params import MAX_CONTROL_INT_DIGITS, ControlOptions from litellm.types.router import CustomPricingLiteLLMParams AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset( @@ -70,6 +76,51 @@ OPTIONAL_KWARGS_KEYS: Final = ( # Backward-compatible alias for existing imports/tests. _OPTIONAL_KWARGS_KEYS: Final = OPTIONAL_KWARGS_KEYS +_CONTROL_OPTIONS: Final = TypeAdapter(ControlOptions) +_CONTROL_OPTION_NAMES: Final = tuple(field.name for field in fields(ControlOptions)) +_MAX_SHOWN_INT_BITS: Final = 64 +_EXPECTED: Final = f"expected a positive integer of at most {MAX_CONTROL_INT_DIGITS} digits" + + +class _BoundedRepr(reprlib.Repr): + def repr_int(self, x: int, level: int) -> str: + if x.bit_length() > _MAX_SHOWN_INT_BITS: + return f"" + return super().repr_int(x, level) + + +_BOUNDED_REPR: Final = _BoundedRepr() + + +@dataclass(frozen=True, slots=True) +class InvalidControlOption: + param: str + message: str + + +def parse_control_options(kwargs: Mapping[str, object]) -> ControlOptions | InvalidControlOption: + given: Final = { # mutable-ok: TypeAdapter.validate_python takes a dict + name: kwargs[name] for name in _CONTROL_OPTION_NAMES if name in kwargs + } + try: + return _CONTROL_OPTIONS.validate_python(given) + except ValidationError as e: + param: Final = str(e.errors(include_url=False)[0]["loc"][0]) + return InvalidControlOption( + param=param, message=f"Invalid {param}={_BOUNDED_REPR.repr(given[param])}: {_EXPECTED}" + ) + + +def stored_control_options(litellm_params: Mapping[str, object]) -> ControlOptions: + control: Final = litellm_params.get(CONTROL_OPTIONS_KEY) + return control if isinstance(control, ControlOptions) else ControlOptions() + + +def with_control_options(litellm_params: Mapping[str, object], control: ControlOptions) -> dict[str, object]: + if control == ControlOptions(): + return dict(litellm_params) # mutable-ok: completion() hands litellm_params to provider code typed as dict + return {**litellm_params, CONTROL_OPTIONS_KEY: control} # mutable-ok: same dict contract as above + def _get_base_model_from_litellm_call_metadata( metadata: dict | None, @@ -130,7 +181,6 @@ def get_litellm_params( api_version: str | None = None, max_retries: int | None = None, litellm_request_debug: bool | None = None, - stream_chunk_size: int | None = None, **kwargs, ) -> dict: _litellm_metadata_dict: Final = litellm_metadata if isinstance(litellm_metadata, dict) else None @@ -193,7 +243,6 @@ def get_litellm_params( "max_retries": max_retries, "use_litellm_proxy": use_litellm_proxy, "litellm_request_debug": litellm_request_debug, - "stream_chunk_size": stream_chunk_size, } # Sparse extraction: only add kwargs keys that are actually present diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 948a9bc6852..f1b41a2302d 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -393,6 +393,8 @@ class BaseConfig(ABC): client: AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": raise NotImplementedError @@ -408,6 +410,8 @@ class BaseConfig(ABC): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": raise NotImplementedError diff --git a/litellm/llms/bedrock/chat/agentcore/transformation.py b/litellm/llms/bedrock/chat/agentcore/transformation.py index 29133bcfaf9..2ad23e84e9f 100644 --- a/litellm/llms/bedrock/chat/agentcore/transformation.py +++ b/litellm/llms/bedrock/chat/agentcore/transformation.py @@ -5,7 +5,7 @@ https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agentcore_InvokeAgen """ import json -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Mapping from typing import TYPE_CHECKING, Any, Final, Optional, Union from urllib.parse import quote @@ -643,6 +643,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """ Simplified sync streaming - returns a generator that yields ModelResponse chunks. @@ -862,6 +864,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): client: Optional["AsyncHTTPHandler"] = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """ Simplified async streaming - returns an async generator that yields ModelResponse chunks. diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 65a34f72167..28df1af4bc0 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -7,6 +7,7 @@ import litellm from litellm.anthropic_beta_headers_manager import ( update_headers_with_filtered_beta, ) +from litellm.litellm_core_utils.get_litellm_params import stored_control_options from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -18,7 +19,7 @@ from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, pop_aws_auth_params, run_aws_signing -from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text, stream_chunk_size_from +from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call @@ -280,7 +281,7 @@ class BedrockConverseLLM(BaseAWSLLM): ): ## SETUP ## stream: Final = optional_params.pop("stream", None) - stream_chunk_size: Final = stream_chunk_size_from(litellm_params) if stream is True else None + stream_chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size if stream is True else None unencoded_model_id: Final = optional_params.pop("model_id", None) fake_stream = optional_params.pop("fake_stream", False) json_mode: Final = optional_params.get("json_mode", False) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py index a8b94fb5703..c5abb5e9a1c 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py @@ -225,7 +225,6 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): anthropic_request.pop("model", None) anthropic_request.pop("stream", None) - anthropic_request.pop("stream_chunk_size", None) apply_bedrock_invoke_structured_output( model=model, request_body=anthropic_request, diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index 629806b58e2..9baf8110b4e 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -1,6 +1,7 @@ import copy import json import time +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, cast, get_args import httpx @@ -9,6 +10,7 @@ from pydantic import TypeAdapter, ValidationError import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import map_finish_reason +from litellm.litellm_core_utils.get_litellm_params import stored_control_options from litellm.litellm_core_utils.logging_utils import track_llm_api_timing from litellm.litellm_core_utils.prompt_templates.factory import ( cohere_message_pt, @@ -18,7 +20,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.bedrock.chat.invoke_handler import make_call, make_sync_call -from litellm.llms.bedrock.common_utils import BedrockError, stream_chunk_size_from +from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.bedrock.request_metadata import ( bedrock_request_metadata_headers, merge_bedrock_invoke_headers, @@ -180,7 +182,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): ) -> dict: ## SETUP ## stream: Final = optional_params.pop("stream", None) - optional_params.pop("stream_chunk_size", None) custom_prompt_dict: Final[dict] = litellm_params.pop("custom_prompt_dict", None) or {} hf_model_name: Final = litellm_params.get("hf_model_name", None) @@ -452,8 +453,10 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): client: AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: - chunk_size: Final = stream_chunk_size_from(logging_obj.litellm_params) + chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size completion_stream, response_headers = await make_call( client=client, api_base=api_base, @@ -489,11 +492,13 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: sync_client: Final = ( _get_httpx_client(params={}) if client is None or isinstance(client, AsyncHTTPHandler) else client ) - chunk_size: Final = stream_chunk_size_from(logging_obj.litellm_params) + chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size completion_stream, response_headers = make_sync_call( client=sync_client, api_base=api_base, diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 5f044897b2c..ccc4309fc5d 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -18,7 +18,7 @@ if TYPE_CHECKING: from litellm.types.llms.bedrock import BedrockCreateBatchRequest import httpx -from pydantic import ConfigDict, TypeAdapter, ValidationError +from pydantic import TypeAdapter, ValidationError import litellm from litellm import verbose_logger @@ -86,15 +86,6 @@ class BedrockError(BaseLLMException): _BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (*AWS_AUTH_PARAM_KEYS, "aws_region_name") -_STREAM_CHUNK_SIZE_VALIDATOR: Final[TypeAdapter[int | None]] = TypeAdapter(int | None, config=ConfigDict(strict=True)) - - -def stream_chunk_size_from(litellm_params: Mapping[str, object]) -> int | None: - raw: Final = litellm_params.get("stream_chunk_size") - try: - return _STREAM_CHUNK_SIZE_VALIDATOR.validate_python(raw) - except ValidationError as e: - raise BedrockError(status_code=400, message=f"Invalid stream_chunk_size={raw!r}. Expected int. Error: {e}") def merge_bedrock_aws_request_params( diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py index e622761dd7f..5846ba560a8 100644 --- a/litellm/llms/bytez/chat/transformation.py +++ b/litellm/llms/bytez/chat/transformation.py @@ -1,6 +1,7 @@ import json import time import traceback +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final import httpx @@ -258,6 +259,8 @@ class BytezChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -300,6 +303,8 @@ class BytezChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.BYTEZ, params={}) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 0cb1416db3f..8aa38ff3341 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -790,6 +790,7 @@ class BaseLLMHTTPHandler: messages=messages, client=client, json_mode=json_mode, + litellm_params=litellm_params, ) completion_stream, headers = self.make_sync_call( provider_config=provider_config, @@ -953,6 +954,7 @@ class BaseLLMHTTPHandler: client=client, json_mode=json_mode, signed_json_body=signed_json_body, + litellm_params=litellm_params, ) completion_stream, _response_headers = await self.make_async_call_stream_helper( diff --git a/litellm/llms/langgraph/chat/transformation.py b/litellm/llms/langgraph/chat/transformation.py index c9388ee472f..293672f1ca9 100644 --- a/litellm/llms/langgraph/chat/transformation.py +++ b/litellm/llms/langgraph/chat/transformation.py @@ -9,6 +9,7 @@ Non-streaming endpoint: POST /runs/wait """ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast import httpx @@ -285,6 +286,8 @@ class LangGraphConfig(BaseConfig): client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for synchronous streaming. @@ -344,6 +347,8 @@ class LangGraphConfig(BaseConfig): client: Optional["AsyncHTTPHandler"] = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for asynchronous streaming. diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index ecff823a18d..24f3ddd5162 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -10,7 +10,7 @@ implement the LiteLLM BaseConfig interface. Heavy-lifting lives in: """ import json -from collections.abc import AsyncIterator, Callable, Iterator +from collections.abc import AsyncIterator, Callable, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final import httpx @@ -642,6 +642,8 @@ class OCIChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "OCIStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -681,6 +683,8 @@ class OCIChatConfig(BaseConfig): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "OCIStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.OCI, params={}) diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index 04995f32d97..f99a3f9e1bc 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -7,6 +7,7 @@ LiteLLM Docs: https://docs.litellm.ai/docs/providers/aws_sagemaker#sagemaker-mes Huggingface Docs: https://huggingface.co/docs/text-generation-inference/en/messages_api """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, cast import httpx @@ -149,6 +150,8 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -191,6 +194,8 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): client: HTTPHandler | AsyncHTTPHandler | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> CustomStreamWrapper: if client is None or isinstance(client, HTTPHandler): try: diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py index b37bf473731..cf889c481a3 100644 --- a/litellm/llms/vertex_ai/agent_engine/transformation.py +++ b/litellm/llms/vertex_ai/agent_engine/transformation.py @@ -10,6 +10,7 @@ API Reference: """ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, Optional, Union, cast import httpx @@ -365,6 +366,8 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): client: Union[HTTPHandler, "AsyncHTTPHandler"] | None = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for synchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( @@ -423,6 +426,8 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): client: Optional["AsyncHTTPHandler"] = None, json_mode: bool | None = None, signed_json_body: bytes | None = None, + *, + litellm_params: Mapping[str, object], ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for asynchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( diff --git a/litellm/main.py b/litellm/main.py index 8c2afe4429a..6c85adf3ae8 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -38,7 +38,7 @@ import dotenv import httpx import openai from pydantic import BaseModel -from typing_extensions import overload +from typing_extensions import assert_never, overload import litellm @@ -48,6 +48,7 @@ from litellm import client # Other utils are imported directly to avoid circular imports from litellm.utils import ( exception_type, + filter_out_litellm_params, get_litellm_params, get_optional_params, peek_reasoning_summary_aliases, @@ -83,6 +84,9 @@ from litellm.litellm_core_utils.get_litellm_params import ( AWS_CREDENTIAL_KWARGS_KEYS, OPTIONAL_KWARGS_KEYS, PROVIDER_AFFINITY_HEADER_KWARG_KEY, + InvalidControlOption, + parse_control_options, + with_control_options, ) from litellm.litellm_core_utils.get_provider_specific_headers import ( ProviderSpecificHeaderUtils, @@ -127,7 +131,7 @@ from litellm.types.completion import ( _CompletionDispatchContext, _CompletionDispatchResult, ) -from litellm.types.litellm_params import RetryStrategy +from litellm.types.litellm_params import ControlOptions, RetryStrategy from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CustomPricingLiteLLMParams, @@ -178,7 +182,7 @@ from litellm.utils import ( from ._logging import verbose_logger from .caching.caching import disable_cache, enable_cache, update_cache -from .litellm_core_utils.core_helpers import safe_deep_copy +from .litellm_core_utils.core_helpers import normalize_drop_params, safe_deep_copy from .litellm_core_utils.fallback_utils import ( async_completion_with_fallbacks, completion_with_fallbacks, @@ -284,7 +288,6 @@ from .types.utils import ( LlmProviders, PromptTokensDetails, ProviderSpecificHeader, - is_litellm_owned_kwarg, ) ####### ENVIRONMENT VARIABLES ################### @@ -335,6 +338,21 @@ ovhcloud_transformation: Final = OVHCloudChatConfig() lemonade_transformation: Final = LemonadeChatConfig() MOCK_RESPONSE_TYPE = str | Exception | dict | ModelResponse | ModelResponseStream + + +def _resolve_control_options(kwargs: Mapping[str, object], model: str) -> ControlOptions: + control: Final = parse_control_options(kwargs) + match control: + case ControlOptions(): + return control + case InvalidControlOption(param=param, message=message): + if litellm.drop_params is True or normalize_drop_params(kwargs.get("drop_params")) is True: + return ControlOptions() + raise litellm.BadRequestError(message=message, model=model, llm_provider=None, body={"param": param}) + case _: + return assert_never(control) + + ####### COMPLETION ENDPOINTS ################ @@ -501,6 +519,7 @@ async def acompletion( loop: Final = asyncio.get_event_loop() custom_llm_provider = kwargs.get("custom_llm_provider", None) + _ = _resolve_control_options(kwargs, model) ## PROMPT MANAGEMENT HOOKS ## ######################################################### @@ -5230,6 +5249,7 @@ def completion( # Responses API config (get_provider_responses_api_config -> None). skip_responses_api_bridge: Final = kwargs.pop("_skip_responses_api_bridge", False) + control_options: Final = _resolve_control_options(kwargs, model) skip_mcp_handler: Final = kwargs.pop("_skip_mcp_handler", False) if not skip_mcp_handler and tools: from litellm.responses.mcp.chat_completions_handler import acompletion_with_mcp @@ -5370,7 +5390,6 @@ def completion( ) ######## end of unpacking kwargs ########### non_default_params: Final = get_non_default_completion_params(kwargs=kwargs) - litellm_params: dict[str, object] = {} # used to prevent unbound var errors ## PROMPT MANAGEMENT HOOKS ## from litellm.integrations.anthropic_cache_control_hook import ( @@ -5622,7 +5641,7 @@ def completion( messages = function_call_prompt(messages=messages, functions=functions_unsupported_model) # For logging - save the values of the litellm-specific params passed in - litellm_params = get_litellm_params( + requested_litellm_params: Final = get_litellm_params( acompletion=acompletion, api_key=api_key, force_timeout=force_timeout, @@ -5670,7 +5689,6 @@ def completion( max_retries=max_retries, timeout=timeout, litellm_request_debug=kwargs.get("litellm_request_debug", False), - stream_chunk_size=kwargs.get("stream_chunk_size"), tpm=kwargs.get("tpm"), rpm=kwargs.get("rpm"), use_xai_oauth=kwargs.get("use_xai_oauth", False), @@ -5683,6 +5701,7 @@ def completion( if key in kwargs }, ) + litellm_params: Final = with_control_options(requested_litellm_params, control_options) if litellm_params.get("provider_affinity_header") is not None: try: headers = add_provider_affinity_header( @@ -6352,9 +6371,7 @@ def embedding( "encoding_format", ] default_params: Final = [*openai_params, "aembedding", "extra_headers"] - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in default_params and not is_litellm_owned_kwarg(k) - } + non_default_params: Final = filter_out_litellm_params(kwargs, excluding=default_params) model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider( model=model, diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index f5ba9ebd3da..439858ea2b5 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -5,7 +5,10 @@ from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequenc from dataclasses import dataclass, field, fields, is_dataclass from itertools import chain from types import MappingProxyType -from typing import TYPE_CHECKING, Final, Literal, TypeAlias +from typing import TYPE_CHECKING, Annotated, Final, Literal, TypeAlias + +from pydantic import BeforeValidator, Field +from pydantic.dataclasses import dataclass as pydantic_dataclass if TYPE_CHECKING: import httpx @@ -234,11 +237,31 @@ class ResponseOptions: merge_reasoning_content_in_choices: bool | None = None enable_json_schema_validation: bool | None = None complete_response: bool | None = None - stream_chunk_size: int | None = None keepalive_seconds: float | None = None allow_client_keepalive_override: bool | None = None +MAX_CONTROL_INT_DIGITS: Final = 18 + + +def _int_from_decimal_string(value: object) -> object: + if isinstance(value, str) and value.isascii() and value.isdecimal() and len(value) <= MAX_CONTROL_INT_DIGITS: + return int(value) + return value + + +@pydantic_dataclass(frozen=True, slots=True, kw_only=True) +class ControlOptions: + stream_chunk_size: ( + Annotated[ + int, + BeforeValidator(_int_from_decimal_string), + Field(strict=True, gt=0, lt=10**MAX_CONTROL_INT_DIGITS), + ] + | None + ) = None + + @dataclass(frozen=True, slots=True, kw_only=True) class MockOptions: mock_response: "MockResponse | None" = None @@ -258,6 +281,7 @@ class LiteLLMOptions: guardrails: GuardrailOptions prompt: PromptOptions response: ResponseOptions + control: ControlOptions mock: MockOptions diff --git a/litellm/utils.py b/litellm/utils.py index 092fe936cf9..7ce412e818c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -288,7 +288,7 @@ except (ImportError, AttributeError, TypeError): # Convert to str (if necessary) claude_json_str = json.dumps(json_data) import importlib.metadata -from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence +from collections.abc import AsyncIterator, Callable, Collection, Iterable, Iterator, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, runtime_checkable from typing_extensions import assert_never @@ -4161,8 +4161,10 @@ def _remove_unsupported_params(non_default_params: dict, supported_openai_params return non_default_params -def filter_out_litellm_params(kwargs: Mapping[str, object]) -> dict: - return {key: value for key, value in kwargs.items() if not is_litellm_owned_kwarg(key)} +def filter_out_litellm_params( + kwargs: Mapping[str, object], excluding: Collection[str] = frozenset() +) -> dict[str, object]: + return {key: value for key, value in kwargs.items() if key not in excluding and not is_litellm_owned_kwarg(key)} def _provider_supports_vertex_params(custom_llm_provider: str) -> bool: @@ -10132,13 +10134,8 @@ def get_standard_openai_params(params: Mapping[str, object]) -> dict: return {k: v for k, v in params.items() if k in litellm.OPENAI_CHAT_COMPLETION_PARAMS and v is not None} -def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict: - openai_params: Final = litellm.OPENAI_CHAT_COMPLETION_PARAMS - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in openai_params and not is_litellm_owned_kwarg(k) - } - - return non_default_params +def get_non_default_completion_params(kwargs: Mapping[str, object]) -> dict[str, object]: + return filter_out_litellm_params(kwargs, excluding=litellm.OPENAI_CHAT_COMPLETION_PARAMS) def peek_reasoning_summary_aliases(optional_params: dict) -> object | None: @@ -10184,13 +10181,10 @@ def strip_reasoning_summary_aliases_from_optional_params( return op, rs_val -def get_non_default_transcription_params(kwargs: Mapping[str, object]) -> dict: +def get_non_default_transcription_params(kwargs: Mapping[str, object]) -> dict[str, object]: from litellm.constants import OPENAI_TRANSCRIPTION_PARAMS - non_default_params: Final = { - k: v for k, v in kwargs.items() if k not in OPENAI_TRANSCRIPTION_PARAMS and not is_litellm_owned_kwarg(k) - } - return non_default_params + return filter_out_litellm_params(kwargs, excluding=OPENAI_TRANSCRIPTION_PARAMS) def add_openai_metadata( diff --git a/tests/_support/stream_chunk_size.py b/tests/_support/stream_chunk_size.py index 051f552e282..6e6256637f0 100644 --- a/tests/_support/stream_chunk_size.py +++ b/tests/_support/stream_chunk_size.py @@ -1,26 +1,28 @@ from collections.abc import Mapping +from types import MappingProxyType from typing import Final -import litellm import pytest -from litellm.integrations.custom_logger import CustomLogger +from litellm.constants import CONTROL_OPTIONS_KEY +from litellm.types.litellm_params import ControlOptions -class LitellmParamsRecorder(CustomLogger): - def __init__(self) -> None: - super().__init__() - self.seen: tuple[Mapping[str, object], ...] = () +DEFAULT_CHUNKING_REQUESTS: Final = ( + pytest.param(MappingProxyType({}), id="unset"), + pytest.param(MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": True}), id="dropped"), + pytest.param( + MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": "true"}), id="dropped_by_string_flag" + ), + pytest.param(MappingProxyType({CONTROL_OPTIONS_KEY: ControlOptions(stream_chunk_size=1)}), id="forged_options"), + pytest.param(MappingProxyType({CONTROL_OPTIONS_KEY: {"stream_chunk_size": 1}}), id="forged_mapping"), +) - def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None: - params: Final = kwargs["litellm_params"] - assert isinstance(params, Mapping) - self.seen = (*self.seen, params) - - -def record_litellm_params(monkeypatch: pytest.MonkeyPatch) -> LitellmParamsRecorder: - recorder: Final = LitellmParamsRecorder() - monkeypatch.setattr(litellm, "input_callback", [recorder]) - return recorder +ROUTER_CHUNK_SIZE_CASES: Final = ( + pytest.param(MappingProxyType({"stream_chunk_size": 64}), 64, id="int"), + pytest.param(MappingProxyType({"stream_chunk_size": "64"}), 64, id="digit_string"), + pytest.param(MappingProxyType({}), None, id="unset"), + pytest.param(MappingProxyType({"stream_chunk_size": "sixty-four", "drop_params": True}), None, id="dropped"), +) def keys_at_every_depth(value: object) -> frozenset[str]: diff --git a/tests/integration/providers/test_internal_params_wire.py b/tests/integration/providers/test_internal_params_wire.py index 17b0fc9d815..f9f5d1e5478 100644 --- a/tests/integration/providers/test_internal_params_wire.py +++ b/tests/integration/providers/test_internal_params_wire.py @@ -8,11 +8,12 @@ from collections.abc import Callable, Mapping from pathlib import Path from typing import Final -import litellm import pytest from integration._support.upstream import INTERNAL_FIELDS from integration._support.wire import Reply, Request, wire_server -from tests._support.stream_chunk_size import keys_at_every_depth, record_litellm_params + +import litellm +from tests._support.stream_chunk_size import keys_at_every_depth TEXT: Final = "wire control" OPENAI_RESPONSE: Final = { @@ -277,13 +278,11 @@ def provider_wire_environment(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) - @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("stream", [False, True]) async def test_internal_params_never_reach_provider_body( - monkeypatch: pytest.MonkeyPatch, provider_wire_environment: None, provider: str, asynchronous: bool, stream: bool, ) -> None: - recorder: Final = record_litellm_params(monkeypatch) with wire_server(_peer(provider)) as wire: parameters: Final = { **_request_parameters(provider, wire.url), @@ -308,8 +307,6 @@ async def test_internal_params_never_reach_provider_body( assert result.choices[0].message.content == TEXT requests: Final = wire.drain() assert len(requests) == 1 - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 body: Final = json.loads(requests[0].body) keys: Final = keys_at_every_depth(body) assert "stream_chunk_size" not in keys diff --git a/tests/unit/caching/test_caching.py b/tests/unit/caching/test_caching.py index 2e4122530d8..0e0f2b7eac6 100644 --- a/tests/unit/caching/test_caching.py +++ b/tests/unit/caching/test_caching.py @@ -1,10 +1,12 @@ import asyncio import logging import re +from typing import Final from unittest.mock import MagicMock import pytest +import litellm import litellm.caching.redis_cache as redis_cache_module from litellm.caching.caching import Cache from litellm.caching.caching_handler import _PENDING_CACHE_WRITES @@ -389,3 +391,15 @@ async def test_embedding_cache_serves_base64_string_embeddings_on_repeat(monkeyp assert embedder.provider_calls == 1, "a string embedding written to the cache must be served on repeat" assert [item["embedding"] for item in second.data] == [item["embedding"] for item in first.data] == ["AACAPwAAAEA="] + + +def test_provider_specific_cache_key_ignores_litellm_owned_kwargs(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "enable_caching_on_provider_specific_optional_params", True) + cache: Final = Cache(type=LiteLLMCacheType.LOCAL) + request: Final = {"model": "gpt-4.1-mini", "messages": [{"role": "user", "content": "hi"}], "top_k": 5} + + base_key: Final = cache.get_cache_key(**request) + + assert cache.get_cache_key(**request, _litellm_control={"stream_chunk_size": 64}) == base_key + assert cache.get_cache_key(**request, litellm_trace_id="trace-1") == base_key + assert cache.get_cache_key(**{**request, "top_k": 6}) != base_key diff --git a/tests/unit/litellm_core_utils/test_get_litellm_params.py b/tests/unit/litellm_core_utils/test_get_litellm_params.py index 39bc2688ae0..9b5771092ac 100644 --- a/tests/unit/litellm_core_utils/test_get_litellm_params.py +++ b/tests/unit/litellm_core_utils/test_get_litellm_params.py @@ -7,12 +7,18 @@ Ensures backward compatibility after sparse kwargs extraction optimization. from typing import Final import pytest +from pydantic import ValidationError +from litellm.constants import CONTROL_OPTIONS_KEY from litellm.litellm_core_utils.get_litellm_params import ( _OPTIONAL_KWARGS_KEYS, + InvalidControlOption, _get_base_model_from_litellm_call_metadata, get_litellm_params, + parse_control_options, + stored_control_options, ) +from litellm.types.litellm_params import ControlOptions NAMED_PRICE_PARAMS: Final = frozenset( {"input_cost_per_token", "output_cost_per_token", "input_cost_per_second", "output_cost_per_second"} @@ -90,9 +96,8 @@ class TestGetLitellmParamsKwargsExtraction: assert "s3_endpoint_url" not in result_without_s3_kwargs assert "s3_region_name" not in result_without_s3_kwargs - def test_stream_chunk_size_is_carried_as_a_litellm_param(self) -> None: - assert get_litellm_params(stream_chunk_size=64)["stream_chunk_size"] == 64 - assert get_litellm_params()["stream_chunk_size"] is None + def test_a_caller_supplied_control_options_key_is_not_carried(self) -> None: + assert CONTROL_OPTIONS_KEY not in get_litellm_params(**{CONTROL_OPTIONS_KEY: {"stream_chunk_size": 64}}) def test_s3_credential_kwargs_are_forwarded_for_s3_signing(self): result = get_litellm_params(s3_access_key_id="s3-key", s3_secret_access_key="s3-secret") @@ -122,6 +127,79 @@ class TestGetLitellmParamsKwargsExtraction: assert result[key] == f"val_{key}" +@pytest.mark.parametrize( + "kwargs,expected", + [ + ({"stream_chunk_size": 64, "temperature": 0.2}, ControlOptions(stream_chunk_size=64)), + ({"stream_chunk_size": "64"}, ControlOptions(stream_chunk_size=64)), + ({"stream_chunk_size": None}, ControlOptions()), + ({"temperature": 0.2}, ControlOptions()), + ], +) +def test_control_options_are_read_from_the_request_kwargs(kwargs: dict[str, object], expected: ControlOptions) -> None: + assert parse_control_options(kwargs) == expected + + +@pytest.mark.parametrize( + "raw,shown", + [ + ("sixty-four", "'sixty-four'"), + (" 64", "' 64'"), + ("-1", "'-1'"), + ("\uff16\uff14", "'\uff16\uff14'"), + ("x" * 500, "'xxxxxxxxxxxx...xxxxxxxxxxxxx'"), + pytest.param(-(10**5000), "", id="huge_negative_int"), + pytest.param(-(2**64 - 1), "-18446744073709551615", id="64_bit_negative_int"), + pytest.param(-(2**64), "", id="65_bit_negative_int"), + pytest.param([-(10**5000)], "[]", id="nested_huge_int"), + pytest.param(10**18, "1000000000000000000", id="19_digit_int"), + pytest.param("1" + "0" * 18, "'1000000000000000000'", id="19_digit_string"), + pytest.param("9" * 5000, "'999999999999...9999999999999'", id="5000_digit_string"), + pytest.param("0" * 18 + "1", "'0000000000000000001'", id="19_digit_string_with_leading_zeros"), + (64.0, "64.0"), + (True, "True"), + (0, "0"), + ("0", "'0'"), + (-1, "-1"), + ], +) +def test_control_options_reject_a_stream_chunk_size_that_is_not_a_positive_int(raw: object, shown: str) -> None: + assert parse_control_options({"stream_chunk_size": raw}) == InvalidControlOption( + param="stream_chunk_size", + message=f"Invalid stream_chunk_size={shown}: expected a positive integer of at most 18 digits", + ) + + +@pytest.mark.parametrize("raw", [10**18 - 1, "9" * 18], ids=["int", "digit_string"]) +def test_control_options_accept_the_largest_18_digit_value(raw: object) -> None: + assert parse_control_options({"stream_chunk_size": raw}) == ControlOptions(stream_chunk_size=10**18 - 1) + + +def test_control_options_accept_an_18_digit_string_with_leading_zeros() -> None: + assert parse_control_options({"stream_chunk_size": "0" * 17 + "1"}) == ControlOptions(stream_chunk_size=1) + + +@pytest.mark.parametrize("raw", [0, -1, "sixty-four", 64.0, True]) +def test_control_options_enforce_their_rule_at_construction(raw: object) -> None: + with pytest.raises(ValidationError): + ControlOptions(stream_chunk_size=raw) # pyright: ignore[reportArgumentType] # the invalid type is the input + + +@pytest.mark.parametrize( + "litellm_params,expected", + [ + ({CONTROL_OPTIONS_KEY: ControlOptions(stream_chunk_size=64)}, ControlOptions(stream_chunk_size=64)), + ({}, ControlOptions()), + ({CONTROL_OPTIONS_KEY: {"stream_chunk_size": 64}}, ControlOptions()), + ({"stream_chunk_size": 64}, ControlOptions()), + ], +) +def test_stored_control_options_reads_only_the_validated_options( + litellm_params: dict[str, object], expected: ControlOptions +) -> None: + assert stored_control_options(litellm_params) == expected + + class TestGetLitellmParamsBaseModel: """Verify base_model resolution precedence.""" diff --git a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py index ed172fdfbff..d1748e1b38d 100644 --- a/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py +++ b/tests/unit/llms/bedrock/chat/invoke_transformations/test_base_invoke_transformation.py @@ -1,4 +1,6 @@ import json +from collections.abc import Mapping +from types import MappingProxyType from typing import Final from unittest.mock import AsyncMock, MagicMock @@ -6,44 +8,40 @@ import httpx import pytest import litellm -from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( - AmazonAnthropicClaudeConfig, -) from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( AmazonInvokeConfig, ) from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler -from tests._support.stream_chunk_size import ( - LitellmParamsRecorder, - keys_at_every_depth, - record_litellm_params, -) +from tests._support.stream_chunk_size import DEFAULT_CHUNKING_REQUESTS, ROUTER_CHUNK_SIZE_CASES, keys_at_every_depth @pytest.mark.parametrize( - "config,model", + "model", [ - (AmazonInvokeConfig, "anthropic.claude-3-sonnet-20240229-v1:0"), - (AmazonInvokeConfig, "amazon.titan-text-express-v1"), - (AmazonInvokeConfig, "mistral.mistral-7b-instruct-v0:2"), - (AmazonAnthropicClaudeConfig, "anthropic.claude-sonnet-4-6"), + "anthropic.claude-sonnet-4-6", + "amazon.titan-text-express-v1", + "mistral.mistral-7b-instruct-v0:2", ], ) -def test_transform_request_drops_stream_chunk_size(config, model): - """stream_chunk_size is a LiteLLM-internal knob for re-chunking the HTTP - response stream. Leaking it into the provider request body makes Bedrock - reject the whole request: ValidationException 'stream_chunk_size: Extra - inputs are not permitted'.""" - request_body = config().transform_request( - model=model, +def test_completion_keeps_stream_chunk_size_out_of_invoke_bodies(model: str) -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) + + litellm.completion( + model=f"bedrock/invoke/{model}", messages=[{"role": "user", "content": "hi"}], - optional_params={"stream": True, "stream_chunk_size": 2048, "max_tokens": 10}, - litellm_params={}, - headers={}, + stream=True, + max_tokens=10, + client=client, + aws_access_key_id="fake", + aws_secret_access_key="fake", + aws_region_name="us-east-1", + stream_chunk_size=2048, ) - assert "stream_chunk_size" not in json.dumps(request_body) + request: Final = send.call_args.args[0] + assert "stream_chunk_size" not in keys_at_every_depth(json.loads(request.content)), request.content def test_validate_environment_maps_guardrail_config_to_invoke_headers(): @@ -243,10 +241,7 @@ def test_transform_response_hands_json_mode_to_nova(): assert json.loads(result.choices[0].message.content) == {"city": "Paris", "temperature": 21} -def _stream_invoke_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, **kwargs -) -> tuple[MagicMock, MagicMock, LitellmParamsRecorder]: - recorder: Final = record_litellm_params(monkeypatch) +def _stream_invoke_completion_with_spied_client(**kwargs: object) -> tuple[MagicMock, MagicMock]: mock_response = MagicMock() mock_response.status_code = 200 mock_response.iter_bytes = MagicMock(return_value=iter([])) @@ -263,39 +258,33 @@ def _stream_invoke_completion_with_spied_client( aws_region_name="us-east-1", **kwargs, ) - return mock_response.iter_bytes, client.post, recorder + return mock_response.iter_bytes, client.post -def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_invoke_body( - monkeypatch: pytest.MonkeyPatch, -): - iter_bytes_spy, post_spy, recorder = _stream_invoke_completion_with_spied_client(monkeypatch, stream_chunk_size=64) +def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_invoke_body() -> None: + iter_bytes_spy, post_spy = _stream_invoke_completion_with_spied_client(stream_chunk_size=64) iter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 -def test_completion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch): - iter_bytes_spy, _, recorder = _stream_invoke_completion_with_spied_client(monkeypatch) +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +def test_completion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], +) -> None: + iter_bytes_spy, _ = _stream_invoke_completion_with_spied_client(**request_kwargs) iter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -async def _astream_invoke_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, **kwargs -) -> tuple[MagicMock, AsyncMock, LitellmParamsRecorder]: +async def _astream_invoke_completion_with_spied_client(**kwargs: object) -> tuple[MagicMock, AsyncMock]: async def _no_bytes(): return yield b"" mock_response = MagicMock() mock_response.status_code = 200 - recorder: Final = record_litellm_params(monkeypatch) mock_response.aiter_bytes = MagicMock(return_value=_no_bytes()) aiter_bytes_spy = mock_response.aiter_bytes client = AsyncHTTPHandler() @@ -311,57 +300,49 @@ async def _astream_invoke_completion_with_spied_client( aws_region_name="us-east-1", **kwargs, ) - return aiter_bytes_spy, client.post, recorder + return aiter_bytes_spy, client.post @pytest.mark.asyncio -async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_invoke_body( - monkeypatch: pytest.MonkeyPatch, -): - aiter_bytes_spy, post_spy, recorder = await _astream_invoke_completion_with_spied_client( - monkeypatch, stream_chunk_size=64 - ) +async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_invoke_body() -> None: + aiter_bytes_spy, post_spy = await _astream_invoke_completion_with_spied_client(stream_chunk_size=64) aiter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 @pytest.mark.asyncio -async def test_acompletion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch): - aiter_bytes_spy, _, recorder = await _astream_invoke_completion_with_spied_client(monkeypatch) +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +async def test_acompletion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], +) -> None: + aiter_bytes_spy, _ = await _astream_invoke_completion_with_spied_client(**request_kwargs) aiter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -@pytest.mark.parametrize("stream_chunk_size,expected_chunk_size", [(64, 64), (None, None)]) -def test_router_deployment_stream_chunk_size_reaches_iter_bytes( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size, expected_chunk_size -): - recorder: Final = record_litellm_params(monkeypatch) - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.iter_bytes = MagicMock(return_value=iter([])) - client = HTTPHandler() - client.post = MagicMock(return_value=mock_response) - deployment_params = { +INVOKE_DEPLOYMENT: Final = MappingProxyType( + { "model": "bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", "aws_access_key_id": "fake", "aws_secret_access_key": "fake", "aws_region_name": "us-east-1", } - router = litellm.Router( - model_list=[ - { - "model_name": "invoke-chunked", - "litellm_params": deployment_params - | ({} if stream_chunk_size is None else {"stream_chunk_size": stream_chunk_size}), - } - ] +) + + +@pytest.mark.parametrize("deployment_extras,expected_chunk_size", ROUTER_CHUNK_SIZE_CASES) +def test_router_deployment_stream_chunk_size_reaches_iter_bytes( + deployment_extras: Mapping[str, object], expected_chunk_size: int | None +) -> None: + mock_response: Final = MagicMock() + mock_response.status_code = 200 + mock_response.iter_bytes = MagicMock(return_value=iter([])) + client: Final = HTTPHandler() + client.post = MagicMock(return_value=mock_response) + router: Final = litellm.Router( + model_list=[{"model_name": "invoke-chunked", "litellm_params": {**INVOKE_DEPLOYMENT, **deployment_extras}}] ) router.completion( @@ -374,17 +355,11 @@ def test_router_deployment_stream_chunk_size_reaches_iter_bytes( mock_response.iter_bytes.assert_called_once_with(chunk_size=expected_chunk_size) data: Final = client.post.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == stream_chunk_size -def test_stream_wrapper_rejects_non_int_stream_chunk_size(monkeypatch: pytest.MonkeyPatch): - record_litellm_params(monkeypatch) - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.iter_bytes = MagicMock(return_value=iter([])) - client = HTTPHandler() - client.post = MagicMock(return_value=mock_response) +def test_invoke_stream_rejects_non_int_stream_chunk_size_before_calling_bedrock() -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) with pytest.raises(litellm.BadRequestError): litellm.completion( @@ -398,4 +373,28 @@ def test_stream_wrapper_rejects_non_int_stream_chunk_size(monkeypatch: pytest.Mo stream_chunk_size="sixty-four", ) - client.post.assert_not_called() + send.assert_not_called() + + +def test_router_deployment_with_a_non_numeric_stream_chunk_size_gets_a_400_before_calling_bedrock() -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "invoke-chunked", + "litellm_params": {**INVOKE_DEPLOYMENT, "stream_chunk_size": "sixty-four"}, + } + ] + ) + + with pytest.raises(litellm.BadRequestError) as exc_info: + router.completion( + model="invoke-chunked", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + ) + + assert exc_info.value.status_code == 400 + send.assert_not_called() diff --git a/tests/unit/llms/bedrock/test_common_utils.py b/tests/unit/llms/bedrock/test_common_utils.py deleted file mode 100644 index cfcc15f186b..00000000000 --- a/tests/unit/llms/bedrock/test_common_utils.py +++ /dev/null @@ -1,20 +0,0 @@ -import pytest - -from litellm.llms.bedrock.common_utils import BedrockError, stream_chunk_size_from - - -def test_stream_chunk_size_from_absent_is_none(): - assert stream_chunk_size_from({}) is None - - -def test_stream_chunk_size_from_int_is_returned(): - assert stream_chunk_size_from({"stream_chunk_size": 64}) == 64 - - -@pytest.mark.parametrize("bad_value", ["64", 6.4, True]) -def test_stream_chunk_size_from_rejects_non_int_with_400(bad_value): - with pytest.raises(BedrockError) as excinfo: - stream_chunk_size_from({"stream_chunk_size": bad_value}) - - assert excinfo.value.status_code == 400 - assert repr(bad_value) in excinfo.value.message diff --git a/tests/unit/llms/chat/test_converse_handler.py b/tests/unit/llms/chat/test_converse_handler.py index cbb8e3acf78..57bb9ab771f 100644 --- a/tests/unit/llms/chat/test_converse_handler.py +++ b/tests/unit/llms/chat/test_converse_handler.py @@ -1,5 +1,6 @@ import json -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Mapping +from types import MappingProxyType from typing import Final from unittest.mock import AsyncMock, MagicMock @@ -11,11 +12,7 @@ from litellm.llms.bedrock.chat import BedrockConverseLLM from litellm.llms.bedrock.chat.converse_handler import make_sync_call from litellm.llms.bedrock.common_utils import _get_all_bedrock_regions from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler -from tests._support.stream_chunk_size import ( - LitellmParamsRecorder, - keys_at_every_depth, - record_litellm_params, -) +from tests._support.stream_chunk_size import DEFAULT_CHUNKING_REQUESTS, ROUTER_CHUNK_SIZE_CASES, keys_at_every_depth def test_encode_model_id_with_inference_profile(): @@ -319,10 +316,7 @@ def test_completion_plumbs_stream_chunk_size_through_converse() -> None: iter_bytes_spy.assert_called_once_with(chunk_size=2048) -def _stream_converse_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None = None -) -> tuple[MagicMock, MagicMock, LitellmParamsRecorder]: - recorder: Final = record_litellm_params(monkeypatch) +def _stream_converse_completion_with_spied_client(**request: object) -> tuple[MagicMock, MagicMock]: mock_response: Final = MagicMock() mock_response.status_code = 200 mock_response.iter_bytes = MagicMock(return_value=iter([])) @@ -337,43 +331,35 @@ def _stream_converse_completion_with_spied_client( aws_access_key_id="fake", aws_secret_access_key="fake", aws_region_name="us-east-1", - stream_chunk_size=stream_chunk_size, + **request, ) - return mock_response.iter_bytes, client.post, recorder + return mock_response.iter_bytes, client.post -def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_converse_body( - monkeypatch: pytest.MonkeyPatch, -) -> None: - iter_bytes_spy, post_spy, recorder = _stream_converse_completion_with_spied_client( - monkeypatch, stream_chunk_size=64 - ) +def test_completion_stream_chunk_size_reaches_iter_bytes_but_not_converse_body() -> None: + iter_bytes_spy, post_spy = _stream_converse_completion_with_spied_client(stream_chunk_size=64) iter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 -def test_completion_without_stream_chunk_size_uses_default_chunking(monkeypatch: pytest.MonkeyPatch) -> None: - iter_bytes_spy, _, recorder = _stream_converse_completion_with_spied_client(monkeypatch) +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +def test_completion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], +) -> None: + iter_bytes_spy, _ = _stream_converse_completion_with_spied_client(**request_kwargs) iter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -async def _astream_converse_completion_with_spied_client( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None = None -) -> tuple[MagicMock, AsyncMock, LitellmParamsRecorder]: +async def _astream_converse_completion_with_spied_client(**request: object) -> tuple[MagicMock, AsyncMock]: async def _no_bytes(chunk_size: int | None = None) -> AsyncIterator[bytes]: return yield b"" mock_response: Final = MagicMock() mock_response.status_code = 200 - recorder: Final = record_litellm_params(monkeypatch) mock_response.aiter_bytes = MagicMock(return_value=_no_bytes()) aiter_bytes_spy: Final = mock_response.aiter_bytes client: Final = AsyncHTTPHandler() @@ -387,61 +373,51 @@ async def _astream_converse_completion_with_spied_client( aws_access_key_id="fake", aws_secret_access_key="fake", aws_region_name="us-east-1", - stream_chunk_size=stream_chunk_size, + **request, ) - return aiter_bytes_spy, client.post, recorder + return aiter_bytes_spy, client.post @pytest.mark.asyncio -async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_converse_body( - monkeypatch: pytest.MonkeyPatch, -) -> None: - aiter_bytes_spy, post_spy, recorder = await _astream_converse_completion_with_spied_client( - monkeypatch, stream_chunk_size=64 - ) +async def test_acompletion_stream_chunk_size_reaches_aiter_bytes_but_not_converse_body() -> None: + aiter_bytes_spy, post_spy = await _astream_converse_completion_with_spied_client(stream_chunk_size=64) aiter_bytes_spy.assert_called_once_with(chunk_size=64) data: Final = post_spy.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == 64 @pytest.mark.asyncio -async def test_acompletion_without_stream_chunk_size_uses_default_chunking( - monkeypatch: pytest.MonkeyPatch, +@pytest.mark.parametrize("request_kwargs", DEFAULT_CHUNKING_REQUESTS) +async def test_acompletion_uses_default_chunking_unless_a_valid_size_is_requested( + request_kwargs: Mapping[str, object], ) -> None: - aiter_bytes_spy, _, recorder = await _astream_converse_completion_with_spied_client(monkeypatch) + aiter_bytes_spy, _ = await _astream_converse_completion_with_spied_client(**request_kwargs) aiter_bytes_spy.assert_called_once_with(chunk_size=None) - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] is None -@pytest.mark.parametrize("stream_chunk_size,expected_chunk_size", [(64, 64), (None, None)]) -def test_router_deployment_stream_chunk_size_reaches_iter_bytes( - monkeypatch: pytest.MonkeyPatch, stream_chunk_size: int | None, expected_chunk_size: int | None -) -> None: - recorder: Final = record_litellm_params(monkeypatch) - mock_response: Final = MagicMock() - mock_response.status_code = 200 - mock_response.iter_bytes = MagicMock(return_value=iter([])) - client: Final = HTTPHandler() - client.post = MagicMock(return_value=mock_response) - deployment_params: Final = { +CONVERSE_DEPLOYMENT: Final = MappingProxyType( + { "model": "bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0", "aws_access_key_id": "fake", "aws_secret_access_key": "fake", "aws_region_name": "us-east-1", } +) + + +@pytest.mark.parametrize("deployment_extras,expected_chunk_size", ROUTER_CHUNK_SIZE_CASES) +def test_router_deployment_stream_chunk_size_reaches_iter_bytes( + deployment_extras: Mapping[str, object], expected_chunk_size: int | None +) -> None: + mock_response: Final = MagicMock() + mock_response.status_code = 200 + mock_response.iter_bytes = MagicMock(return_value=iter([])) + client: Final = HTTPHandler() + client.post = MagicMock(return_value=mock_response) router: Final = litellm.Router( - model_list=[ - { - "model_name": "converse-chunked", - "litellm_params": deployment_params - | ({} if stream_chunk_size is None else {"stream_chunk_size": stream_chunk_size}), - } - ] + model_list=[{"model_name": "converse-chunked", "litellm_params": {**CONVERSE_DEPLOYMENT, **deployment_extras}}] ) router.completion( @@ -454,20 +430,18 @@ def test_router_deployment_stream_chunk_size_reaches_iter_bytes( mock_response.iter_bytes.assert_called_once_with(chunk_size=expected_chunk_size) data: Final = client.post.call_args.kwargs["data"] assert "stream_chunk_size" not in keys_at_every_depth(json.loads(data)), data - assert len(recorder.seen) == 1 - assert recorder.seen[0]["stream_chunk_size"] == stream_chunk_size -def test_converse_stream_rejects_non_int_stream_chunk_size_before_calling_bedrock(monkeypatch: pytest.MonkeyPatch): - record_litellm_params(monkeypatch) - client = HTTPHandler() - client.post = MagicMock() +@pytest.mark.parametrize("stream", [True, False], ids=["stream", "non_stream"]) +def test_converse_rejects_non_int_stream_chunk_size_before_calling_bedrock(stream: bool) -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(send))) with pytest.raises(litellm.BadRequestError): litellm.completion( model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0", messages=[{"role": "user", "content": "hi"}], - stream=True, + stream=stream, client=client, aws_access_key_id="fake", aws_secret_access_key="fake", @@ -475,30 +449,7 @@ def test_converse_stream_rejects_non_int_stream_chunk_size_before_calling_bedroc stream_chunk_size="sixty-four", ) - client.post.assert_not_called() - - -def test_converse_non_stream_ignores_invalid_stream_chunk_size(): - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.json = MagicMock(return_value=_converse_response_body()) - mock_response.text = json.dumps(_converse_response_body()) - mock_response.headers = httpx.Headers() - client = HTTPHandler() - client.post = MagicMock(return_value=mock_response) - - response = litellm.completion( - model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0", - messages=[{"role": "user", "content": "hi"}], - client=client, - aws_access_key_id="fake", - aws_secret_access_key="fake", - aws_region_name="us-east-1", - stream_chunk_size="64", - ) - - assert response.choices[0].message.content == "hi" - client.post.assert_called_once() + send.assert_not_called() def _bedrock_error_response(status_code: int, request_id: str) -> httpx.Response: diff --git a/tests/unit/llms/oci/chat/test_oci_chat_transformation.py b/tests/unit/llms/oci/chat/test_oci_chat_transformation.py index 708187b8ae1..462b1d6ea72 100644 --- a/tests/unit/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/unit/llms/oci/chat/test_oci_chat_transformation.py @@ -1247,6 +1247,7 @@ class TestOCIStreamingSignedBody: mock_logging = MagicMock() config.get_sync_custom_stream_wrapper( + litellm_params={}, api_base="https://example.com", headers={}, data={"key": "value"}, @@ -1286,6 +1287,7 @@ class TestOCIStreamingSignedBody: payload = {"key": "value"} config.get_sync_custom_stream_wrapper( + litellm_params={}, api_base="https://example.com", headers={}, data=payload, diff --git a/tests/unit/llms/oci/test_oci_coverage_boost.py b/tests/unit/llms/oci/test_oci_coverage_boost.py index 7c91ece70b5..8f7588c5de7 100644 --- a/tests/unit/llms/oci/test_oci_coverage_boost.py +++ b/tests/unit/llms/oci/test_oci_coverage_boost.py @@ -1111,6 +1111,7 @@ def test_get_sync_custom_stream_wrapper_returns_wrapper(): mock_client.post.return_value = mock_response wrapper = config.get_sync_custom_stream_wrapper( + litellm_params={}, model=_GENERIC_MODEL, custom_llm_provider="oci", logging_obj=MagicMock(), @@ -1143,6 +1144,7 @@ async def test_get_async_custom_stream_wrapper_returns_wrapper(): mock_client.post = AsyncMock(return_value=mock_response) wrapper = await config.get_async_custom_stream_wrapper( + litellm_params={}, model=_GENERIC_MODEL, custom_llm_provider="oci", logging_obj=MagicMock(), diff --git a/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py index 697f5a7ff59..cb20b3390bb 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_chat_transformation.py @@ -125,6 +125,7 @@ def test_sync_first_event_emitted_after_a_single_frame(): response = httpx.Response(200, stream=stream) wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper( + litellm_params={}, model="phi-4", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), @@ -147,6 +148,7 @@ def test_sync_events_emitted_incrementally_without_bursting(): response = httpx.Response(200, stream=stream) wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper( + litellm_params={}, model="phi-4", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), @@ -171,6 +173,7 @@ async def test_async_first_event_emitted_after_a_single_frame(): response = httpx.Response(200, stream=stream) wrapper = await SagemakerChatConfig().get_async_custom_stream_wrapper( + litellm_params={}, model="phi-4", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), diff --git a/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py b/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py index 5cc414819e3..d878bc70a09 100644 --- a/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py +++ b/tests/unit/llms/sagemaker/test_sagemaker_nova_transformation.py @@ -309,6 +309,7 @@ class TestSagemakerChatBackwardsCompatibility: ) as mock_csw: mock_csw.return_value = MagicMock() self.config.get_sync_custom_stream_wrapper( + litellm_params={}, model="my-hf-endpoint", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), @@ -348,6 +349,7 @@ class TestSagemakerChatBackwardsCompatibility: mock_csw.return_value = MagicMock() asyncio.run( self.config.get_async_custom_stream_wrapper( + litellm_params={}, model="my-hf-endpoint", custom_llm_provider="sagemaker_chat", logging_obj=MagicMock(), diff --git a/tests/unit/responses/test_responses_api_bridge_flag.py b/tests/unit/responses/test_responses_api_bridge_flag.py index 642495fab86..fb1361c1f49 100644 --- a/tests/unit/responses/test_responses_api_bridge_flag.py +++ b/tests/unit/responses/test_responses_api_bridge_flag.py @@ -12,6 +12,7 @@ from typing import Final from unittest.mock import MagicMock, patch import httpx +import openai import pytest import respx @@ -592,3 +593,20 @@ class TestUseResponsesApiBridgeFlag: mock_native_handler.assert_called_once() assert result is not None + + def test_bridge_still_rejects_an_invalid_stream_chunk_size(self) -> None: + send: Final = MagicMock(return_value=httpx.Response(200)) + client: Final = openai.OpenAI(api_key="fake-key", http_client=httpx.Client(transport=httpx.MockTransport(send))) + + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.responses( + model="openai/gpt-4.1-mini", + input="hi", + use_chat_completions_api=True, + stream_chunk_size="sixty-four", + client=client, + num_retries=0, + ) + + assert exc_info.value.param == "stream_chunk_size" + send.assert_not_called() diff --git a/tests/unit/test_filter_out_litellm_params.py b/tests/unit/test_filter_out_litellm_params.py index 72f8f5f1478..342e251d611 100644 --- a/tests/unit/test_filter_out_litellm_params.py +++ b/tests/unit/test_filter_out_litellm_params.py @@ -2,6 +2,10 @@ Test filter_out_litellm_params helper function. """ +from typing import Final + + +import litellm from litellm.utils import filter_out_litellm_params @@ -34,3 +38,19 @@ def test_filter_out_litellm_params(): assert "litellm_trace_id" not in filtered assert "proxy_server_request" not in filtered assert "secret_fields" not in filtered + + +def test_filter_out_litellm_params_also_drops_the_excluded_names(): + kwargs = {"temperature": 0.2, "top_k": 5, "litellm_trace_id": "trace-1", "_litellm_control": object()} + + assert filter_out_litellm_params(kwargs, excluding=("temperature",)) == {"top_k": 5} + + +def test_filter_out_litellm_params_sees_a_name_appended_to_the_public_list_after_import(): + litellm.all_litellm_params.append("registered_later") + try: + filtered: Final = filter_out_litellm_params({"registered_later": 1, "top_k": 2}) + finally: + litellm.all_litellm_params.remove("registered_later") + + assert filtered == {"top_k": 2} diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index 7bef35d8559..e0e1fcfe105 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -22,10 +22,16 @@ from unittest.mock import MagicMock, patch import litellm from litellm import main as litellm_main +from litellm.constants import CONTROL_OPTIONS_KEY from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.custom_prompt_management import CustomPromptManagement from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs +from litellm.litellm_core_utils.get_litellm_params import stored_control_options from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage +from litellm.types.litellm_params import ControlOptions +from litellm.types.llms.openai import AllMessageValues +from litellm.types.prompts.init_prompts import PromptSpec +from litellm.types.utils import Delta, ModelResponseStream, StandardCallbackDynamicParams, StreamingChoices, Usage @pytest.fixture(autouse=True) @@ -4273,3 +4279,228 @@ def test_completion_rejects_untranslatable_tool_choice_with_a_400(tool_choice): ) assert exc_info.value.status_code == 400 assert f"tool_choice={tool_choice}" in str(exc_info.value) + + +@pytest.mark.parametrize("raw", ["sixty-four", 0, -1]) +def test_completion_rejects_an_invalid_stream_chunk_size_with_a_400_naming_the_param(raw: object) -> None: + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size=raw, + mock_response="unused", + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.param == "stream_chunk_size" + assert f"Invalid stream_chunk_size={raw!r}: expected a positive integer of at most 18 digits" in str(exc_info.value) + + +class _PromptHookRecorder(CustomPromptManagement): + def __init__(self, on_prompt: MagicMock) -> None: + super().__init__() + self.on_prompt: Final = on_prompt + + def get_chat_completion_prompt( + self, + model: str, + messages: list[AllMessageValues], + non_default_params: dict, + prompt_id: str | None, + prompt_variables: dict | None, + dynamic_callback_params: StandardCallbackDynamicParams, + prompt_spec: PromptSpec | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: + self.on_prompt("sync") + return model, messages, non_default_params + + async def async_get_chat_completion_prompt( + self, + model: str, + messages: list[AllMessageValues], + non_default_params: dict, + prompt_id: str | None, + prompt_variables: dict | None, + dynamic_callback_params: StandardCallbackDynamicParams, + litellm_logging_obj: LiteLLMLogging, + prompt_spec: PromptSpec | None = None, + tools: list[dict] | None = None, + prompt_label: str | None = None, + prompt_version: int | None = None, + ignore_prompt_manager_model: bool | None = False, + ignore_prompt_manager_optional_params: bool | None = False, + ) -> tuple[str, list[AllMessageValues], dict]: + self.on_prompt("async") + return model, messages, non_default_params + + +async def _call_completion(is_async: bool, **kwargs: object) -> None: + if is_async: + await litellm.acompletion(**kwargs) + else: + litellm.completion(**kwargs) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async,hook", [(False, "sync"), (True, "async")], ids=["completion", "acompletion"]) +async def test_the_prompt_hook_runs_when_stream_chunk_size_is_valid( + monkeypatch: pytest.MonkeyPatch, is_async: bool, hook: str +) -> None: + on_prompt: Final = MagicMock() + monkeypatch.setattr(litellm, "callbacks", [_PromptHookRecorder(on_prompt)]) + + await _call_completion( + is_async, + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + prompt_id="greeting", + stream_chunk_size=64, + mock_response="hi", + ) + + on_prompt.assert_any_call(hook) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", [False, True], ids=["completion", "acompletion"]) +async def test_an_invalid_stream_chunk_size_is_rejected_before_any_prompt_hook_runs( + monkeypatch: pytest.MonkeyPatch, is_async: bool +) -> None: + on_prompt: Final = MagicMock() + monkeypatch.setattr(litellm, "callbacks", [_PromptHookRecorder(on_prompt)]) + + with pytest.raises(litellm.BadRequestError): + await _call_completion( + is_async, + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + prompt_id="greeting", + stream_chunk_size="sixty-four", + mock_response="hi", + ) + + on_prompt.assert_not_called() + + +def _completion_logging_obj(call_id: str) -> LiteLLMLogging: + return LiteLLMLogging( + model="gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="completion", + start_time=datetime(2026, 1, 1), + litellm_call_id=call_id, + function_id=f"{call_id}-function", + ) + + +def test_completion_carries_the_control_options_into_the_logged_litellm_params() -> None: + logging_obj: Final = _completion_logging_obj("control-params") + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size=64, + mock_response="hi", + litellm_logging_obj=logging_obj, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions(stream_chunk_size=64) + + +def test_completion_ignores_a_caller_supplied_control_options_key() -> None: + logging_obj: Final = _completion_logging_obj("control-params-injection") + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + mock_response="hi", + litellm_logging_obj=logging_obj, + **{CONTROL_OPTIONS_KEY: {"stream_chunk_size": 1}}, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions() + + +@pytest.mark.parametrize("drop_params", [True, "true"]) +def test_drop_params_drops_an_invalid_stream_chunk_size_instead_of_rejecting_it(drop_params: object) -> None: + logging_obj: Final = _completion_logging_obj(f"drop-params-{drop_params}") + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size="sixty-four", + drop_params=drop_params, + mock_response="hi", + litellm_logging_obj=logging_obj, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions() + + +def test_drop_params_keeps_a_dropped_stream_chunk_size_out_of_the_provider_request( + respx_mock: respx.MockRouter, +) -> None: + api_base: Final = "http://localhost:12346/v1" + mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/chat/completions.*").mock( + return_value=httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-drop", + "object": "chat.completion", + "created": 1712697600, + "model": "gpt-4.1-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + ) + ) + + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + api_base=api_base, + api_key="fake_openai_api_key", + stream_chunk_size="sixty-four", + drop_params=True, + ) + + assert mock_route.called + sent: Final = json.loads(respx_mock.calls[0].request.content) + assert "stream_chunk_size" not in sent, sent + assert sent["model"] == "gpt-4.1-mini" + + +@pytest.mark.asyncio +async def test_global_drop_params_drops_an_invalid_stream_chunk_size_on_acompletion( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "drop_params", True) + logging_obj: Final = _completion_logging_obj("global-drop-params") + await litellm.acompletion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size=0, + mock_response="hi", + litellm_logging_obj=logging_obj, + ) + assert stored_control_options(logging_obj.litellm_params) == ControlOptions() + + +def test_completion_rejects_an_invalid_stream_chunk_size_before_the_mcp_gateway() -> None: + with pytest.raises(litellm.BadRequestError) as exc_info: + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + tools=[{"type": "mcp", "server_label": "gateway", "server_url": "litellm_proxy"}], + stream_chunk_size="sixty-four", + ) + assert exc_info.value.param == "stream_chunk_size" + + +def test_drop_params_false_still_rejects_an_invalid_stream_chunk_size() -> None: + with pytest.raises(litellm.BadRequestError): + litellm.completion( + model="openai/gpt-4.1-mini", + messages=[{"role": "user", "content": "hi"}], + stream_chunk_size="sixty-four", + drop_params=False, + mock_response="hi", + ) diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index 33467a78aa8..a2d944fcf39 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -504,7 +504,8 @@ LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2}, litellm_params.GuardrailOptions: {"guardrails": ("default",)}, litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}}, - litellm_params.ResponseOptions: {"stream_chunk_size": 64}, + litellm_params.ResponseOptions: {"keepalive_seconds": 1.5}, + litellm_params.ControlOptions: {"stream_chunk_size": 64}, litellm_params.MockOptions: {"mock_timeout": True}, litellm_params.CallState: { "completion_call_id": "call", @@ -533,7 +534,8 @@ LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"}, litellm_params.GuardrailOptions: {"guardrails": (1,)}, litellm_params.PromptOptions: {"prompt_id": 1}, - litellm_params.ResponseOptions: {"stream_chunk_size": "64"}, + litellm_params.ResponseOptions: {"keepalive_seconds": "1.5"}, + litellm_params.ControlOptions: {"stream_chunk_size": "sixty-four"}, litellm_params.MockOptions: {"mock_timeout": "true"}, litellm_params.CallState: {"completion_call_id": 1}, litellm_params.AgenticLoopState: {"depth": "1"}, @@ -583,10 +585,8 @@ def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, samp @pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id) def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None: - instance: Final = _leaf_instance(leaf, sample) - with pytest.raises(ValidationError): - _strict_leaf_validation(leaf, instance) + _strict_leaf_validation(leaf, _leaf_instance(leaf, sample)) @pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id) From b248b1c7dc12c194c4e1e176b9432391ef5ae5ab Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 16:09:09 -0700 Subject: [PATCH 04/51] fix(openai): exclude fine-tuned and custom gpt-5-chat aliases from gpt-5 reasoning path (#43185) * fix(openai): exclude fine-tuned and custom gpt-5-chat aliases from gpt-5 reasoning path Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(openai): keep gpt-5-chat alias regression test diff minimal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(openai): cover temperature pass-through for gpt-5-chat aliases Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(openai): annotate locals and wrap long lines in gpt-5-chat alias test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../llms/openai/chat/gpt_5_transformation.py | 2 +- .../llms/openai/test_is_model_gpt_5_model.py | 34 +++++++++++++++++-- 2 files changed, 32 insertions(+), 4 deletions(-) diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index d0e5ff01e71..bf6b52225f2 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -69,7 +69,7 @@ GPT_REASONING_SERIES_MARKERS: Final = ("gpt-5", "gpt-6") def is_gpt_reasoning_series_name(model: str) -> bool: normalized: Final = model.split("/")[-1] - return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and not normalized.startswith("gpt-5-chat") + return any(marker in model for marker in GPT_REASONING_SERIES_MARKERS) and "gpt-5-chat" not in normalized class OpenAIGPT5Config(OpenAIGPTConfig): diff --git a/tests/unit/llms/openai/test_is_model_gpt_5_model.py b/tests/unit/llms/openai/test_is_model_gpt_5_model.py index 0bb8425d95e..f6fef92fc6a 100644 --- a/tests/unit/llms/openai/test_is_model_gpt_5_model.py +++ b/tests/unit/llms/openai/test_is_model_gpt_5_model.py @@ -26,14 +26,18 @@ There are two distinct families: ``gpt-5.3-chat``, …) — ARE GPT-5 reasoning models and must stay on the GPT-5 path. -The fix uses a prefix check (``startswith("gpt-5-chat")``) on the normalised model -name instead of a substring check, which correctly distinguishes the two families. +The fix uses a substring check for ``gpt-5-chat`` on the normalised model +name (not a prefix check), which correctly distinguishes the two families. """ +from typing import Final + import pytest -from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config +import litellm from litellm.llms.azure.chat.gpt_5_transformation import AzureOpenAIGPT5Config +from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig # --------------------------------------------------------------------------- # Parametrized fixtures @@ -73,6 +77,9 @@ NON_GPT5_MODELS = [ "gpt-5-chat", # gpt-5-chat family — regular chat path "gpt-5-chat-latest", # gpt-5-chat family with alias suffix "gpt-5-chat-2025-08-07", # gpt-5-chat family with date suffix + "ft:gpt-5-chat-latest:org:abc", + "my-custom-gpt-5-chat", + "openai/ft:gpt-5-chat-latest:org:abc", "gpt-4", "gpt-4o", "gpt-4-turbo", @@ -117,6 +124,27 @@ class TestOpenAIGPT5ConfigIsModelGpt5Model: model ), f"Expected '{model}' (gpt-5-chat family) NOT to be on the GPT-5 path" + def test_responses_api_gpt5_chat_aliases_are_not_gpt5(self): + for model in ["ft:gpt-5-chat-latest:org:abc", "openai/my-custom-gpt-5-chat"]: + assert not OpenAIResponsesAPIConfig._is_gpt_5_model( + model + ), f"Expected Responses API '{model}' NOT to be on the GPT-5 path" + + @pytest.mark.parametrize("model", ["ft:gpt-5-chat-latest:org:abc", "my-custom-gpt-5-chat"]) + def test_gpt5_chat_aliases_keep_non_default_temperature(self, model: str): + chat_params: Final = litellm.get_optional_params( + model=model, custom_llm_provider="openai", temperature=0.7 + ) + responses_params: Final = OpenAIResponsesAPIConfig().map_openai_params( + response_api_optional_params={"temperature": 0.7}, model=model, drop_params=False + ) + assert chat_params["temperature"] == 0.7, ( + f"chat completions dropped or rejected temperature for '{model}'" + ) + assert responses_params["temperature"] == 0.7, ( + f"responses dropped or rejected temperature for '{model}'" + ) + # Models that are gpt-5.4 or newer. main.py gates the automatic switch to the # /v1/responses bridge (when reasoning_effort is set and tools are passed) on From 849f3037b4f43c4e4f60233ae6a5f62905d8e192 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 16:31:52 -0700 Subject: [PATCH 05/51] fix(langtrace): deliver spans to app.langtrace.ai/api/trace with x-api-key (#43322) * test(langtrace): integration test for the built-in callback wire (path, x-api-key) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(langtrace): deliver spans to app.langtrace.ai/api/trace with x-api-key The built-in langtrace callback posted to the dead host langtrace.ai, sent the key as api_key instead of x-api-key, and let the OTLP endpoint normalizer append /v1/traces to the complete /api/trace path, so every export returned 404. Default the host to https://app.langtrace.ai, honor LANGTRACE_API_HOST for self-hosted servers, pass the key as an exporter header instead of a process-wide env var, and keep the /api/trace path unchanged for traces on the langtrace callback only Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(langtrace): keep LANGTRACE_API_HOST that already ends in /api/trace Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(langtrace): deterministic audit inventory for the built-in callback (surfaces, failures, endpoints, chaos) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(langtrace): build the repeated-request body once Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(langtrace): outage cell asserts at-most-once delivery, not span loss The OTLP HTTP exporter reposts once on ConnectionError and the batch processor may still be flushing the previous burst when the sink closes, so whether the outage burst is lost or delivered after revival depends on timing. The invariant is no duplicate and recovery on the same port Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(langtrace): delivered spans must carry the prompt in the gen_ai.content.prompt event The upstream echoes the marker into the completion, so a whole-span match alone would still pass if the prompt event disappeared Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(langtrace): disable model info refresh so the scripted upstream only sees completion requests 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/integrations/langtrace.py | 8 + litellm/integrations/opentelemetry.py | 4 + litellm/litellm_core_utils/litellm_logging.py | 5 +- .../observability/test_langtrace_delivery.py | 684 ++++++++++++++++++ tests/unit/integrations/test_opentelemetry.py | 38 + .../test_litellm_logging.py | 42 ++ 6 files changed, 779 insertions(+), 2 deletions(-) create mode 100644 tests/integration/observability/test_langtrace_delivery.py diff --git a/litellm/integrations/langtrace.py b/litellm/integrations/langtrace.py index 0b4e1393ee6..53f5d2a0318 100644 --- a/litellm/integrations/langtrace.py +++ b/litellm/integrations/langtrace.py @@ -10,6 +10,14 @@ if TYPE_CHECKING: else: Span = Any +LANGTRACE_DEFAULT_HOST: Final = "https://app.langtrace.ai" +LANGTRACE_TRACE_PATH: Final = "/api/trace" + + +def langtrace_trace_endpoint(api_host: str | None) -> str: + host: Final = (api_host or LANGTRACE_DEFAULT_HOST).rstrip("/") + return host if host.endswith(LANGTRACE_TRACE_PATH) else host + LANGTRACE_TRACE_PATH + class LangtraceAttributes: """ diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 948e3113337..8d588896b2f 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -15,6 +15,7 @@ from litellm.integrations._types.open_inference import ( SpanAttributes, ) from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.langtrace import LANGTRACE_TRACE_PATH from litellm.integrations.opentelemetry_utils.gen_ai_semconv import ( OTEL_SEMCONV_STABILITY_OPT_IN_ENV, OTELGenAISemconvMixin, @@ -3334,6 +3335,9 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): if signal_type == "traces" and "/v2/trace/otlp" in endpoint: return endpoint + if signal_type == "traces" and self.callback_name == "langtrace" and endpoint.endswith(LANGTRACE_TRACE_PATH): + return endpoint + # Check if endpoint already ends with the correct signal path target_path: Final = f"/v1/{signal_type}" if endpoint.endswith(target_path): diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index f6211869913..9ee7a7b0a7a 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -60,6 +60,7 @@ from litellm.integrations.arize.arize import ArizeLogger from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.deepeval.deepeval import DeepEvalLogger +from litellm.integrations.langtrace import langtrace_trace_endpoint from litellm.integrations.mlflow import MlflowLogger from litellm.integrations.sqs import SQSLogger from litellm.litellm_core_utils.classifier_logging import ( @@ -4920,9 +4921,9 @@ def _init_custom_logger_compatible_class( otel_config = OpenTelemetryConfig( exporter="otlp_http", - endpoint="https://langtrace.ai/api/trace", + endpoint=langtrace_trace_endpoint(os.getenv("LANGTRACE_API_HOST")), + headers=f"x-api-key={os.environ['LANGTRACE_API_KEY']}", ) - os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = f"api_key={os.getenv('LANGTRACE_API_KEY')}" for callback in _in_memory_loggers: if isinstance(callback, OpenTelemetry) and callback.callback_name == "langtrace": return callback diff --git a/tests/integration/observability/test_langtrace_delivery.py b/tests/integration/observability/test_langtrace_delivery.py new file mode 100644 index 00000000000..84c9ed5bec0 --- /dev/null +++ b/tests/integration/observability/test_langtrace_delivery.py @@ -0,0 +1,684 @@ +import asyncio +import json +import re +import signal +import time +import uuid +from collections.abc import Callable, Iterator, Sequence +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass, field +from itertools import repeat +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import psutil +import pytest +import yaml +from anthropic import Anthropic +from integration._support.client import Gateway, eventually +from integration._support.process import OwnedProxy, owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from openai import AsyncOpenAI, OpenAI +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest +from opentelemetry.proto.trace.v1.trace_pb2 import Span, Status +from pydantic import TypeAdapter + +TRACE_PATH: Final = "/api/trace" +STOCK_CONFIG: Final = Path("tests/integration/proxy_config.yaml") +_PROXY_CONFIG: Final = TypeAdapter(dict[str, object]) +_SETTINGS: Final = TypeAdapter(dict[str, object]) +_MARKER: Final = re.compile(rb"lt[0-9a-f]{32}") +_USAGE: Final = {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15} + + +def _marker() -> str: + return "lt" + uuid.uuid4().hex + + +def _sse(events: Sequence[object]) -> tuple[bytes, ...]: + return tuple(b"data: " + json.dumps(event).encode() + b"\n\n" for event in events) + (b"data: [DONE]\n\n",) + + +def _chat_reply(marker: str, stream: bool) -> Reply: + identity: Final = "chatcmpl-" + marker + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "echo " + marker}, + "finish_reason": "stop", + } + ], + "usage": _USAGE, + } + ).encode() + ) + head: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + return Reply( + content_type="text/event-stream", + chunks=_sse( + ( + {**head, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "echo "}}]}, + {**head, "choices": [{"index": 0, "delta": {"content": marker}}]}, + {**head, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + {**head, "choices": [], "usage": _USAGE}, + ) + ), + ) + + +def _responses_reply(marker: str, stream: bool) -> Reply: + completed: Final = { + "id": "resp_" + marker, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msg_" + marker, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "echo " + marker, "annotations": []}], + } + ], + "parallel_tool_calls": False, + "tool_choice": "auto", + "tools": [], + "usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15}, + } + if not stream: + return Reply(body=json.dumps(completed).encode()) + events: Final = ( + {"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}}, + { + "type": "response.output_text.delta", + "item_id": "msg_" + marker, + "output_index": 0, + "content_index": 0, + "delta": "echo " + marker, + }, + {"type": "response.completed", "response": completed}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _upstream(request: Request) -> Reply: + match: Final = _MARKER.search(request.body) + assert match is not None, request.body[:300] + marker: Final = match.group().decode() + if b'"fail"' in request.body: + return Reply(status=401, body=json.dumps({"error": {"message": "bad provider key " + marker}}).encode()) + stream: Final = json.loads(request.body).get("stream") is True + if request.target.endswith("/responses"): + return _responses_reply(marker, stream) + return _chat_reply(marker, stream) + + +def _config(tmp_path: Path, **litellm_settings: object) -> Path: + config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text())) + settings: Final = {**_SETTINGS.validate_python(config["litellm_settings"]), **litellm_settings} + general: Final = {**_SETTINGS.validate_python(config["general_settings"]), "disable_model_info_refresh": True} + path: Final = tmp_path / "langtrace.yaml" + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings, "general_settings": general})) + return path + + +def _spans(batches: Sequence[Request]) -> tuple[Span, ...]: + return tuple( + span + for batch in batches + for resource_spans in ExportTraceServiceRequest.FromString(batch.body).resource_spans + for scope_spans in resource_spans.scope_spans + for span in scope_spans.spans + ) + + +def _prompt_events(span: Span) -> tuple[str, ...]: + return tuple( + attribute.value.string_value + for event in span.events + if event.name == "gen_ai.content.prompt" + for attribute in event.attributes + if attribute.key == "gen_ai.prompt" + ) + + +def _assert_prompted_with(span: Span, marker: str) -> Span: + prompts: Final = _prompt_events(span) + assert any(marker in prompt for prompt in prompts), (span.name, prompts) + return span + + +def _spans_carrying(batches: Sequence[Request], marker: str, name: str | None = "litellm_request") -> tuple[Span, ...]: + return tuple( + span for span in _spans(batches) if name in (None, span.name) and marker.encode() in span.SerializeToString() + ) + + +def _streamed_text(sse: str, key: str) -> str: + def strings(node: object) -> Iterator[str]: + if isinstance(node, dict): + for field_name, value in node.items(): + if field_name == key and isinstance(value, str): + yield value + else: + yield from strings(value) + if isinstance(node, list): + for item in node: + yield from strings(item) + + events: Final = tuple( + json.loads(line.removeprefix("data: ")) + for line in sse.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + return "".join(text for event in events for text in strings(event)) + + +def _accepted(request: Request) -> Reply: + return Reply(body=b'{"message":"Traces added successfully"}') + + +@dataclass(frozen=True, slots=True) +class _Sink: + wire: Wire + api_key: str + # mutable-ok: drain() consumes, so batches accumulate across polls + received: list[Request] = field(default_factory=list) + + def collect(self) -> tuple[Request, ...]: + self.received.extend(self.wire.drain()) + return tuple(self.received) + + def spans_for(self, marker: str) -> tuple[Span, ...]: + return _spans_carrying(self.collect(), marker) + + def assert_wire_contract(self, batches: Sequence[Request], target: str = TRACE_PATH) -> None: + for request in batches: + assert (request.method, request.target) == ("POST", target), (request.method, request.target) + assert request.headers.get("x-api-key") == self.api_key, request.headers + assert "api_key" not in request.headers, request.headers + assert request.headers.get("content-type") == "application/x-protobuf", request.headers + assert self.api_key.encode() not in b"".join(batch.body for batch in batches) + + def delivered_once(self, marker: str, seconds: float = 20) -> Span: + batches: Final = eventually( + self.collect, lambda value: len(_spans_carrying(value, marker)) >= 1, seconds=seconds + ) + self.assert_wire_contract(batches) + settled: Final = eventually( + self.collect, lambda value: len(_spans_carrying(value, marker)) >= 2, seconds=1, return_last_on_timeout=True + ) + spans: Final = _spans_carrying(settled, marker) + assert len(spans) == 1, [span.span_id for span in spans] + return _assert_prompted_with(spans[0], marker) + + +@dataclass(frozen=True, slots=True) +class _Rig: + proxy: Gateway + model: str + provider: Wire + sink: _Sink + + def provider_hits(self, marker: str) -> int: + return sum(marker.encode() in request.body for request in self.provider.drain()) + + +@contextmanager +def _langtrace_rig( + gateway: Gateway, + tmp_path: Path, + *, + mode: str = "callbacks", + host: Callable[[str], str] = lambda url: url, + api_key: str | None = None, + workers: int = 1, + respond: Callable[[Request], Reply] = _accepted, + sink_port: int = 0, +) -> Iterator[_Rig]: + key: Final = "synthetic-langtrace-key-" + uuid.uuid4().hex if api_key is None else api_key + with wire_server(_upstream) as provider, wire_server(respond, port=sink_port) as sink: + overrides: Final = { + "LANGTRACE_API_KEY": key, + "LANGTRACE_API_HOST": host(sink.url), + "OTEL_BSP_SCHEDULE_DELAY": "300", + } + config: Final = _config(tmp_path, **{mode: ["langtrace"]}) + with ( + owned_proxy(gateway, tmp_path, overrides, config=config, workers=workers) as proxy, + proxy.scenario() as scenario, + ): + yield _Rig(proxy, scenario.model(api_base=provider.url + "/v1"), provider, _Sink(sink, key)) + + +def _chat_httpx(rig: _Rig, marker: str, stream: bool) -> str: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": rig.model, "messages": [{"role": "user", "content": marker}], "stream": stream}, + ) + assert response.status_code == 200, response.text + return _streamed_text(response.text, "content") if stream else response.text + + +def _chat_openai_sync_stream(rig: _Rig, marker: str, stream: bool) -> str: + with OpenAI(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client: + chunks: Final = client.chat.completions.create( + model=rig.model, messages=[{"role": "user", "content": marker}], stream=True + ) + return "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + + +def _chat_openai_async(rig: _Rig, marker: str, stream: bool) -> str: + async def call() -> str: + async with AsyncOpenAI(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client: + completion: Final = await client.chat.completions.create( + model=rig.model, messages=[{"role": "user", "content": marker}] + ) + return completion.model_dump_json() + + return asyncio.run(call()) + + +def _messages_anthropic(rig: _Rig, marker: str, stream: bool) -> str: + with Anthropic(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client: + message: Final = client.messages.create( + model=rig.model, max_tokens=64, messages=[{"role": "user", "content": marker}] + ) + return message.model_dump_json() + + +def _messages_httpx(rig: _Rig, marker: str, stream: bool) -> str: + response: Final = rig.proxy.request( + "POST", + "/v1/messages", + {"model": rig.model, "max_tokens": 64, "messages": [{"role": "user", "content": marker}], "stream": stream}, + ) + assert response.status_code == 200, response.text + return _streamed_text(response.text, "text") if stream else response.text + + +def _responses_openai(rig: _Rig, marker: str, stream: bool) -> str: + with OpenAI(base_url=str(rig.proxy.client.base_url), api_key=rig.proxy.key, max_retries=0) as client: + return client.responses.create(model=rig.model, input=marker).model_dump_json() + + +def _responses_httpx(rig: _Rig, marker: str, stream: bool) -> str: + response: Final = rig.proxy.request( + "POST", "/v1/responses", {"model": rig.model, "input": marker, "stream": stream} + ) + assert response.status_code == 200, response.text + return _streamed_text(response.text, "delta") if stream else response.text + + +@dataclass(frozen=True, slots=True) +class _Surface: + call: Callable[[_Rig, str, bool], str] + stream: bool + + +_SURFACES: Final = ( + pytest.param(_Surface(_chat_httpx, False), id="chat-httpx"), + pytest.param(_Surface(_chat_openai_sync_stream, True), id="chat-openai-sync-stream"), + pytest.param(_Surface(_chat_openai_async, False), id="chat-openai-async"), + pytest.param(_Surface(_messages_anthropic, False), id="messages-anthropic"), + pytest.param(_Surface(_messages_httpx, True), id="messages-httpx-stream"), + pytest.param(_Surface(_responses_openai, False), id="responses-openai"), + pytest.param(_Surface(_responses_httpx, True), id="responses-httpx-stream"), +) + + +def _assert_delivered(rig: _Rig, surface: _Surface, marker: str) -> Span: + text: Final = surface.call(rig, marker, surface.stream) + assert "echo " + marker in text, text + assert rig.provider_hits(marker) == 1 + return rig.sink.delivered_once(marker) + + +@pytest.mark.parametrize("surface", _SURFACES) +def test_langtrace_span_reaches_api_trace_with_x_api_key(gateway: Gateway, tmp_path: Path, surface: _Surface) -> None: + with _langtrace_rig(gateway, tmp_path) as rig: + _assert_delivered(rig, surface, _marker()) + + +def test_langtrace_exports_cache_hit_twin_as_its_own_span(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = _marker() + with _langtrace_rig(gateway, tmp_path) as rig: + first: Final = rig.proxy.request( + "POST", "/v1/chat/completions", {"model": rig.model, "messages": [{"role": "user", "content": marker}]} + ) + assert first.status_code == 200, first.text + rig.sink.delivered_once(marker) + second: Final = rig.proxy.request( + "POST", "/v1/chat/completions", {"model": rig.model, "messages": [{"role": "user", "content": marker}]} + ) + assert second.status_code == 200 and second.headers.get("x-litellm-cache-key"), second.headers + assert second.json()["id"] == first.json()["id"], second.text + assert rig.provider_hits(marker) == 1 + batches: Final = eventually( + rig.sink.collect, lambda value: len(_spans_carrying(value, marker)) >= 2, seconds=20 + ) + rig.sink.assert_wire_contract(batches) + assert len(_spans_carrying(batches, marker)) == 2 + + +def test_langtrace_success_callback_mode_delivers(gateway: Gateway, tmp_path: Path) -> None: + with _langtrace_rig(gateway, tmp_path, mode="success_callback") as rig: + _assert_delivered(rig, _Surface(_chat_httpx, False), _marker()) + + +def test_langtrace_failure_callback_mode_exports_provider_error_span(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = _marker() + with _langtrace_rig(gateway, tmp_path, mode="failure_callback") as rig: + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": rig.model, "messages": [{"role": "user", "content": marker + " fail"}], "user": "fail"}, + ) + assert response.status_code == 401, response.text + assert "bad provider key " + marker in response.text, response.text + assert rig.provider_hits(marker) == 1 + span: Final = rig.sink.delivered_once(marker) + assert span.status.code == Status.STATUS_CODE_ERROR, span.status + + +@pytest.mark.parametrize("status", (403, 404), ids=("forbidden", "not-found")) +def test_langtrace_rejecting_sink_leaves_callers_and_later_exports_intact( + gateway: Gateway, tmp_path: Path, status: int +) -> None: + scripted: Final[SimpleQueue[int]] = SimpleQueue() + + def respond(request: Request) -> Reply: + return Reply(status=scripted.get_nowait()) if not scripted.empty() else _accepted(request) + + rejected: Final = _marker() + accepted: Final = _marker() + with _langtrace_rig(gateway, tmp_path, respond=respond) as rig: + scripted.put(status) + assert "echo " + rejected in _chat_httpx(rig, rejected, False) + batches: Final = eventually( + rig.sink.collect, lambda value: len(_spans_carrying(value, rejected)) >= 1, seconds=20 + ) + rig.sink.assert_wire_contract(batches) + assert scripted.empty() + assert "echo " + accepted in _chat_httpx(rig, accepted, False) + rig.sink.delivered_once(accepted) + assert rig.proxy.request("GET", "/health/liveliness").status_code == 200 + + +def test_langtrace_missing_api_key_logs_startup_error_and_exports_nothing(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = _marker() + with wire_server(_upstream) as provider, wire_server(_accepted) as sink: + overrides: Final = {"LANGTRACE_API_HOST": sink.url, "OTEL_BSP_SCHEDULE_DELAY": "300"} + with ( + owned_proxy_process( + gateway, + tmp_path, + overrides, + config=_config(tmp_path, callbacks=["langtrace"]), + remove_environment=("LANGTRACE_API_KEY",), + ) as owned, + owned.gateway.scenario() as scenario, + ): + assert "LANGTRACE_API_KEY not found in environment variables" in owned.log.read_text() + rig: Final = _Rig(owned.gateway, scenario.model(api_base=provider.url + "/v1"), provider, _Sink(sink, "")) + assert "echo " + marker in _chat_httpx(rig, marker, False) + assert rig.provider_hits(marker) == 1 + batches: Final = eventually( + rig.sink.collect, lambda value: len(value) >= 1, seconds=2, return_last_on_timeout=True + ) + assert batches == (), batches + + +def test_langtrace_empty_api_key_still_posts_to_api_trace(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = _marker() + with _langtrace_rig(gateway, tmp_path, api_key="") as rig: + assert "echo " + marker in _chat_httpx(rig, marker, False) + batches: Final = eventually( + rig.sink.collect, lambda value: len(_spans_carrying(value, marker)) >= 1, seconds=20 + ) + for request in batches: + assert (request.method, request.target) == ("POST", TRACE_PATH), (request.method, request.target) + assert "api_key" not in request.headers, request.headers + assert request.headers.get("x-api-key", "") == "", request.headers + + +@pytest.mark.parametrize( + "host", + (lambda url: url + "/", lambda url: url + TRACE_PATH, lambda url: url + TRACE_PATH + "/"), + ids=("trailing-slash", "already-suffixed", "suffixed-trailing-slash"), +) +def test_langtrace_api_host_variants_append_api_trace_exactly_once( + gateway: Gateway, tmp_path: Path, host: Callable[[str], str] +) -> None: + with _langtrace_rig(gateway, tmp_path, host=host) as rig: + _assert_delivered(rig, _Surface(_chat_httpx, False), _marker()) + + +def test_langtrace_logs_repeated_identical_requests_once_each(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = _marker() + with _langtrace_rig(gateway, tmp_path) as rig: + body: Final = { + "model": rig.model, + "messages": [{"role": "user", "content": marker}], + "cache": {"no-cache": True}, + } + responses: Final = tuple(rig.proxy.request("POST", "/v1/chat/completions", body) for _ in range(2)) + assert [response.status_code for response in responses] == [200, 200], [r.text for r in responses] + assert rig.provider_hits(marker) == 2 + batches: Final = eventually( + rig.sink.collect, lambda value: len(_spans_carrying(value, marker)) >= 2, seconds=20 + ) + rig.sink.assert_wire_contract(batches) + settled: Final = eventually( + rig.sink.collect, + lambda value: len(_spans_carrying(value, marker)) >= 3, + seconds=1, + return_last_on_timeout=True, + ) + assert len(_spans_carrying(settled, marker)) == 2 + + +def test_generic_otel_callback_keeps_v1_traces_suffix_on_api_trace_endpoint(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = _marker() + with wire_server(_upstream) as provider, wire_server(_accepted) as sink: + overrides: Final = { + "OTEL_EXPORTER": "otlp_http", + "OTEL_ENDPOINT": sink.url + TRACE_PATH, + "OTEL_HEADERS": "x-api-key=generic-otel-key", + "OTEL_BSP_SCHEDULE_DELAY": "300", + } + with ( + owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["otel"])) as proxy, + proxy.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1") + response: Final = proxy.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": marker}]} + ) + assert response.status_code == 200, response.text + collector: Final = _Sink(sink, "generic-otel-key") + batches: Final = eventually( + collector.collect, lambda value: len(_spans_carrying(value, marker)) >= 1, seconds=20 + ) + collector.assert_wire_contract(batches, target=TRACE_PATH + "/v1/traces") + + +def test_langtrace_otel_v2_route_still_targets_collector_v1_traces(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = _marker() + with wire_server(_upstream) as provider, wire_server(_accepted) as collector: + overrides: Final = { + "LITELLM_OTEL_V2": "true", + "LANGTRACE_API_KEY": "unused-by-the-collector-route", + "OTEL_EXPORTER_OTLP_ENDPOINT": collector.url, + "OTEL_EXPORTER_OTLP_PROTOCOL": "http/protobuf", + "OTEL_EXPORTER_OTLP_HEADERS": "x-api-key=collector-key", + "OTEL_BSP_SCHEDULE_DELAY": "300", + } + with ( + owned_proxy(gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["langtrace"])) as proxy, + proxy.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1") + response: Final = proxy.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": marker}]} + ) + assert response.status_code == 200, response.text + sink: Final = _Sink(collector, "collector-key") + batches: Final = eventually( + sink.collect, lambda value: len(_spans_carrying(value, marker, name=None)) >= 1, seconds=20 + ) + sink.assert_wire_contract(batches, target="/v1/traces") + + +_BURST: Final = ( + (_chat_httpx, False), + (_chat_httpx, True), + (_messages_httpx, False), + (_messages_httpx, True), + (_responses_httpx, False), + (_responses_httpx, True), +) + + +def _burst_call(rig: _Rig, index: int, marker: str) -> str: + call, stream = _BURST[index % len(_BURST)] + return call(rig, marker, stream) + + +def _burst(rig: _Rig, size: int) -> tuple[str, ...]: + markers: Final = tuple(_marker() for _ in range(size)) + with ThreadPoolExecutor(max_workers=size) as pool: + texts: Final = tuple(pool.map(_burst_call, repeat(rig), range(size), markers)) + for marker, text in zip(markers, texts, strict=True): + assert "echo " + marker in text, text + return markers + + +def _assert_each_once(sink: _Sink, markers: Sequence[str], seconds: float = 30) -> None: + batches: Final = eventually( + sink.collect, lambda value: all(_spans_carrying(value, marker) for marker in markers), seconds=seconds + ) + sink.assert_wire_contract(batches) + settled: Final = eventually( + sink.collect, + lambda value: any(len(_spans_carrying(value, marker)) > 1 for marker in markers), + seconds=1, + return_last_on_timeout=True, + ) + counts: Final = {marker: len(_spans_carrying(settled, marker)) for marker in markers} + assert all(count == 1 for count in counts.values()), counts + for marker in markers: + _assert_prompted_with(_spans_carrying(settled, marker)[0], marker) + + +def test_langtrace_two_workers_deliver_every_burst_span_exactly_once(gateway: Gateway, tmp_path: Path) -> None: + with _langtrace_rig(gateway, tmp_path, workers=2) as rig: + markers: Final = _burst(rig, 24) + _assert_each_once(rig.sink, markers) + + +def test_langtrace_sink_outage_mid_burst_recovers_on_the_same_port(gateway: Gateway, tmp_path: Path) -> None: + key: Final = "synthetic-langtrace-key-" + uuid.uuid4().hex + with wire_server(_accepted) as probe: + port: Final = int(probe.url.rsplit(":", 1)[1]) + host: Final = f"http://127.0.0.1:{port}" + with wire_server(_upstream) as provider: + overrides: Final = {"LANGTRACE_API_KEY": key, "LANGTRACE_API_HOST": host, "OTEL_BSP_SCHEDULE_DELAY": "300"} + with ( + owned_proxy_process( + gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["langtrace"]), workers=2 + ) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1") + with wire_server(_accepted, port=port) as sink: + rig: Final = _Rig(owned.gateway, model, provider, _Sink(sink, key)) + _assert_each_once(rig.sink, _burst(rig, 6)) + log: Final = owned.log + failures_before: Final = log.read_text().count("Exception while exporting Span batch") + outage: Final = _burst(rig, 12) + eventually( + lambda: log.read_text().count("Exception while exporting Span batch"), + lambda value: value > failures_before, + seconds=20, + ) + assert owned.gateway.request("GET", "/health/liveliness").status_code == 200 + with wire_server(_accepted, port=port) as revived: + recovered: Final = _Rig(owned.gateway, model, provider, _Sink(revived, key)) + _assert_each_once(recovered.sink, _burst(recovered, 6)) + counts: Final = {marker: len(recovered.sink.spans_for(marker)) for marker in outage} + assert all(count <= 1 for count in counts.values()), counts + + +def test_langtrace_slow_sink_does_not_delay_callers_or_duplicate_spans(gateway: Gateway, tmp_path: Path) -> None: + def slow(request: Request) -> Reply: + time.sleep(1) + return _accepted(request) + + with _langtrace_rig(gateway, tmp_path, respond=slow) as rig: + started: Final = time.monotonic() + markers: Final = _burst(rig, 6) + assert time.monotonic() - started < 5 + _assert_each_once(rig.sink, markers, seconds=40) + + +def test_langtrace_survives_a_killed_worker(gateway: Gateway, tmp_path: Path) -> None: + key: Final = "synthetic-langtrace-key-" + uuid.uuid4().hex + with wire_server(_upstream) as provider, wire_server(_accepted) as sink: + overrides: Final = {"LANGTRACE_API_KEY": key, "LANGTRACE_API_HOST": sink.url, "OTEL_BSP_SCHEDULE_DELAY": "300"} + with ( + owned_proxy_process( + gateway, tmp_path, overrides, config=_config(tmp_path, callbacks=["langtrace"]), workers=2 + ) as owned, + httpx.Client( + base_url=owned.gateway.client.base_url, + timeout=15, + trust_env=False, + limits=httpx.Limits(max_keepalive_connections=0), + ) as fresh_connections, + ): + proxy: Final = Gateway(fresh_connections, owned.gateway.key, owned.gateway.upstream_url) + with proxy.scenario() as scenario: + rig: Final = _Rig(proxy, scenario.model(api_base=provider.url + "/v1"), provider, _Sink(sink, key)) + _assert_kill_and_recovery(owned, rig) + + +def _cmdline(process: psutil.Process) -> str: + try: + return " ".join(process.cmdline()) + except psutil.Error: + return "" + + +def _assert_kill_and_recovery(owned: OwnedProxy, rig: _Rig) -> None: + _assert_delivered(rig, _Surface(_chat_httpx, False), _marker()) + + def uvicorn_workers() -> tuple[psutil.Process, ...]: + return tuple(child for child in psutil.Process(owned.process.pid).children() if "spawn_main" in _cmdline(child)) + + workers: Final = uvicorn_workers() + assert len(workers) == 2, workers + workers[0].send_signal(signal.SIGKILL) + eventually( + uvicorn_workers, + lambda value: len(value) == 2 and workers[0].pid not in {child.pid for child in value}, + seconds=30, + ) + _assert_each_once(rig.sink, _burst(rig, 6)) diff --git a/tests/unit/integrations/test_opentelemetry.py b/tests/unit/integrations/test_opentelemetry.py index 52eeec31e71..175bd95c263 100644 --- a/tests/unit/integrations/test_opentelemetry.py +++ b/tests/unit/integrations/test_opentelemetry.py @@ -1949,6 +1949,44 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase): expected, ) + @parameterized.expand( + [ + ("https://app.langtrace.ai/api/trace", "https://app.langtrace.ai/api/trace"), + ("https://app.langtrace.ai/api/trace/", "https://app.langtrace.ai/api/trace"), + ("http://localhost:3000/api/trace", "http://localhost:3000/api/trace"), + ] + ) + def test_langtrace_callback_keeps_api_trace_endpoint_unchanged(self, input_url: str, expected: str) -> None: + """Langtrace ingests OTLP at the complete /api/trace path, so no /v1/traces is appended.""" + otel = OpenTelemetry(callback_name="langtrace") + self.assertEqual(otel._normalize_otel_endpoint(input_url, "traces"), expected) + + @parameterized.expand( + [ + (None, "https://app.langtrace.ai/api/trace", "https://app.langtrace.ai/api/trace/v1/traces"), + ("otel", "https://app.langtrace.ai/api/trace", "https://app.langtrace.ai/api/trace/v1/traces"), + ("otel", "https://collector.example.com/api/trace", "https://collector.example.com/api/trace/v1/traces"), + ("langtrace", "https://app.langtrace.ai", "https://app.langtrace.ai/v1/traces"), + ] + ) + def test_api_trace_exemption_is_scoped_to_langtrace_callback( + self, callback_name: str | None, input_url: str, expected: str + ) -> None: + """Any other callback, or a Langtrace host without the /api/trace path, keeps OTLP normalization.""" + otel = OpenTelemetry(callback_name=callback_name) + self.assertEqual(otel._normalize_otel_endpoint(input_url, "traces"), expected) + + def test_langtrace_callback_still_normalizes_logs_and_metrics(self) -> None: + otel = OpenTelemetry(callback_name="langtrace") + self.assertEqual( + otel._normalize_otel_endpoint("https://app.langtrace.ai/api/trace", "logs"), + "https://app.langtrace.ai/api/trace/v1/logs", + ) + self.assertEqual( + otel._normalize_otel_endpoint("https://app.langtrace.ai/api/trace", "metrics"), + "https://app.langtrace.ai/api/trace/v1/metrics", + ) + def test_normalize_endpoint_none(self): """Test that None endpoint returns None""" otel = OpenTelemetry() diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index f211505d06d..c8b02ebc790 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -1616,6 +1616,48 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch): logging_module._in_memory_loggers.clear() +@pytest.mark.parametrize( + ("api_host", "expected_endpoint"), + [ + (None, "https://app.langtrace.ai/api/trace"), + ("http://langtrace.internal:3000/", "http://langtrace.internal:3000/api/trace"), + ("http://langtrace.internal:3000/api/trace", "http://langtrace.internal:3000/api/trace"), + ], +) +def test_langtrace_callback_exports_to_api_trace_with_x_api_key( + monkeypatch: pytest.MonkeyPatch, api_host: str | None, expected_endpoint: str +) -> None: + """The exporter must post to Langtrace's complete /api/trace path with the key in x-api-key, + without leaking it into the process-wide OTEL_EXPORTER_OTLP_TRACES_HEADERS.""" + from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter + + from litellm.integrations.opentelemetry import OpenTelemetry + from litellm.litellm_core_utils import litellm_logging as logging_module + + api_key: Final = "synthetic-langtrace-key" + monkeypatch.setenv("LANGTRACE_API_KEY", api_key) + monkeypatch.delenv("LANGTRACE_API_HOST", raising=False) + monkeypatch.delenv("OTEL_EXPORTER_OTLP_TRACES_HEADERS", raising=False) + if api_host is not None: + monkeypatch.setenv("LANGTRACE_API_HOST", api_host) + logging_module._in_memory_loggers.clear() + try: + logger: Final = logging_module._init_custom_logger_compatible_class( + logging_integration="langtrace", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert type(logger) is OpenTelemetry and logger.callback_name == "langtrace" + exporter: Final = logger._get_span_processor().span_exporter + assert isinstance(exporter, OTLPSpanExporter) + assert exporter._endpoint == expected_endpoint + assert exporter._headers == {"x-api-key": api_key} + assert "OTEL_EXPORTER_OTLP_TRACES_HEADERS" not in os.environ + finally: + logging_module._in_memory_loggers.clear() + + @pytest.mark.asyncio async def test_logging_result_for_bridge_calls(logging_obj): """ From b2e82cf3beeb1cf87645ef2b91d214664895e0ea Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 26 Sep 2026 16:44:30 -0700 Subject: [PATCH 06/51] fix(caching): stand default cache points down when extra_body hides a direct client mark (#43341) * fix(caching): stand default cache points down when extra_body hides a direct client mark On native /v1/messages the extra_body envelope is dropped, so a client tool mark or root cache_control reaches Anthropic even when extra_body overrides it. The stand-down check only counted the envelope-merged view and injected two default marks on top of the client's. * fix(caching): keep chat completions on the envelope-merged mark count for the default stand-down Chat completions merge extra_body over the request, so a direct tool mark that extra_body replaces never reaches the provider there. Only /v1/messages, where the native transforms drop the envelope, needs to count marks on both sides. --- .../anthropic_cache_control_hook.py | 19 +++++++--- .../test_anthropic_cache_control_hook.py | 36 +++++++++++++++++++ 2 files changed, 50 insertions(+), 5 deletions(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 0d6cbc2232e..1db144b5fdc 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -695,6 +695,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): tools: list | None = None, cache_control: object = None, request_kwargs: object = None, + on_messages_route: bool = False, ) -> bool: """Return True if the request already carries any client-supplied cache_control. @@ -704,10 +705,14 @@ class AnthropicCacheControlHook(CustomPromptManagement): envelope. Configured injection points are an explicit instruction and are applied alongside the client's marks, bounded by the provider cap. """ - return ( - AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system) - + AnthropicCacheControlHook.count_external_cache_breakpoints(tools, cache_control, request_kwargs) - ) > 0 + external_breakpoints: Final = ( + AnthropicCacheControlHook.count_external_cache_breakpoints_on_messages_route( + tools, cache_control, request_kwargs + ) + if on_messages_route + else AnthropicCacheControlHook.count_external_cache_breakpoints(tools, cache_control, request_kwargs) + ) + return AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system) + external_breakpoints > 0 @staticmethod def get_default_injection_points( @@ -719,6 +724,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): enable_prompt_caching: bool | None = None, cache_control: object = None, request_kwargs: object = None, + on_messages_route: bool = False, ) -> list[CacheControlInjectionPoint]: """Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on. @@ -739,7 +745,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): if not supports_anthropic_cache_control(model, custom_llm_provider): return [] - if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools, cache_control, request_kwargs): + if AnthropicCacheControlHook._request_has_cache_control( + messages, system, tools, cache_control, request_kwargs, on_messages_route + ): return [] if is_claude_code_one_shot_subagent_request( @@ -968,6 +976,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): enable_prompt_caching=enable_prompt_caching, cache_control=cache_control, request_kwargs=kwargs, + on_messages_route=True, ) if model is not None else () diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index f787d370f04..1d70a21af7b 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -2633,6 +2633,42 @@ class TestConfiguredInjectionPointsSurviveClientMarks: assert kwargs["cache_control"] is root_cache_control assert "litellm_gateway_injected_cache" not in kwargs["litellm_metadata"] + @pytest.mark.parametrize( + "tools,kwargs,injected", + [ + ([MARKED_V1_TOOL], {"extra_body": {"tools": [UNMARKED_V1_TOOL]}}, False), + (None, {"cache_control": EPHEMERAL, "extra_body": {"cache_control": None}}, False), + ([UNMARKED_V1_TOOL], {"extra_body": {"tools": [UNMARKED_V1_TOOL]}}, True), + ], + ids=["extra_body_unmarks_direct_tool", "extra_body_nulls_root_cache_control", "no_client_mark_anywhere"], + ) + def test_v1_messages_automatic_defaults_stand_down_for_a_direct_mark_extra_body_hides( + self, monkeypatch, tools, kwargs, injected + ): + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + request_kwargs = {**copy.deepcopy(kwargs), "litellm_metadata": {}} + + result_messages, result_system = self._inject( + copy.deepcopy(self.V1_MESSAGES), request_kwargs, tools=copy.deepcopy(tools) + ) + + assert AnthropicCacheControlHook.count_request_cache_breakpoints(result_messages, result_system) == ( + 2 if injected else 0 + ) + assert ("litellm_gateway_injected_cache" in request_kwargs["litellm_metadata"]) is injected + + def test_chat_automatic_defaults_apply_when_extra_body_drops_the_only_client_mark(self, monkeypatch): + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + params = {"extra_body": {"tools": [self.UNMARKED_TOOL]}} + + self._seed(params, copy.deepcopy(self.CLEAN_MESSAGES), tools=[self.MARKED_TOOL_TOP_LEVEL]) + affinity = AnthropicCacheControlHook.messages_with_default_injections( + copy.deepcopy(self.CLEAN_MESSAGES), ["claude-sonnet-4-5"], tools=[self.MARKED_TOOL_TOP_LEVEL], request_kwargs=params + ) + + assert [p["index"] for p in params["cache_control_injection_points"]] == [None, -1] + assert AnthropicCacheControlHook.count_request_cache_breakpoints(affinity) == 2 + @pytest.mark.parametrize( "marked_turns,expected_system", [(2, [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]), (3, "sys")], From 7244040908658ce94fbb072bb77243a6ed123efd Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 26 Sep 2026 16:44:49 -0700 Subject: [PATCH 07/51] fix(mcp): report reachability without stored credentials (#43240) --- litellm/models/mcp_server.py | 4 +- .../mcp_server/mcp_server_manager.py | 62 ++- litellm/proxy/_lazy_openapi_snapshot.json | 30 +- .../mcp_management_endpoints.py | 32 +- .../mcp_server/test_mcp_env_vars.py | 40 +- .../mcp_server/test_mcp_server_manager.py | 361 +++++++++++++----- .../test_mcp_management_endpoints.py | 169 +++++++- .../_components/MCPServerCard.test.tsx | 14 + .../mcp-servers/_components/MCPServerCard.tsx | 7 +- .../_components/mcp_servers.test.tsx | 2 + .../mcp-servers/_components/mcp_servers.tsx | 5 +- .../AIHub/MCPHubTableColumns.test.tsx | 13 +- .../components/AIHub/MCPHubTableColumns.tsx | 8 +- .../src/components/mcp_tools/types.tsx | 4 +- .../src/components/networking.test.ts | 24 ++ .../src/components/networking.tsx | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 11 +- 17 files changed, 644 insertions(+), 143 deletions(-) diff --git a/litellm/models/mcp_server.py b/litellm/models/mcp_server.py index 9125d708e79..efc8574932f 100644 --- a/litellm/models/mcp_server.py +++ b/litellm/models/mcp_server.py @@ -73,9 +73,9 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): mcp_info: MCPInfo | None = None static_headers: dict[str, str] | None = None env_vars: list[MCPEnvVar] | None = None - status: Literal["healthy", "unhealthy", "unknown"] | None = Field( + status: Literal["healthy", "reachable", "unhealthy", "unknown"] | None = Field( default="unknown", - description="Health status: 'healthy', 'unhealthy', 'unknown'", + description="Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked)", ) last_health_check: datetime | None = None health_check_error: str | None = None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8befc99cad4..d0d9100971d 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -897,6 +897,41 @@ def _sanitized_error_text(exc: Exception) -> str: return re.sub(r"https?://\S+", "", str(exc))[:200] +async def _mcp_server_reachability( + server: MCPServer, *, timeout: float +) -> tuple[Literal["reachable", "unhealthy", "unknown"], str | None]: + if server.transport not in (MCPTransport.http, MCPTransport.sse) or not server.url: + return "unknown", "Server reachability requires an HTTP or SSE URL" + try: + url: Final = httpx.URL(server.url) + except (httpx.InvalidURL, ValueError): + return "unknown", "Server reachability requires an HTTP URL without embedded credentials" + if url.scheme not in ("http", "https") or not url.host or url.userinfo: + return "unknown", "Server reachability requires an HTTP URL without embedded credentials" + + async def probe() -> None: + handler: Final = get_async_httpx_client(llm_provider="mcp_reachability") + async with handler.client.stream( + "GET", + url, + headers={"Accept": "text/event-stream, application/json"}, + auth=None, + follow_redirects=False, + timeout=timeout, + ): + pass + + try: + await asyncio.wait_for(probe(), timeout=timeout) + except (asyncio.TimeoutError, httpx.TimeoutException): + return "unhealthy", f"Reachability check timed out after {timeout} seconds" + except asyncio.CancelledError: + return "unknown", "Reachability check was cancelled" + except Exception as exc: + return "unhealthy", f"Reachability check failed ({type(exc).__name__})" + return "reachable", None + + async def _openapi_spec_health( spec_path: str, *, timeout: float ) -> tuple[Literal["healthy", "unhealthy", "unknown"], str | None]: @@ -6986,13 +7021,9 @@ class MCPServerManager: ) ) - status: Literal["healthy", "unhealthy", "unknown"] = "unknown" + status: Literal["healthy", "reachable", "unhealthy", "unknown"] = "unknown" health_check_error = None - # Check if we should skip health check based on auth configuration - should_skip_health_check = False - - # Skip if server requires per-user authentication (OAuth2 or passthrough auth) if ( server.requires_per_user_auth or ( @@ -7003,9 +7034,8 @@ class MCPServerManager: ) or self._references_per_user_env_var(server) ): - should_skip_health_check = True - - if not should_skip_health_check: + status, health_check_error = await _mcp_server_reachability(server, timeout=MCP_HEALTH_CHECK_TIMEOUT) + else: try: resolved_static_headers: Final = await self._resolve_static_headers_with_env_vars( server=server, @@ -7081,6 +7111,8 @@ class MCPServerManager: self, user_api_key_auth: UserAPIKeyAuth | None = None, server_ids: list[str] | None = None, + *, + checked_server_ids: frozenset[str] = frozenset(), ) -> list[LiteLLM_MCPServerTable]: """ Get all MCP servers that the user has access to, with health status and team information. @@ -7105,7 +7137,7 @@ class MCPServerManager: # Check all accessible servers target_server_ids = allowed_server_ids - return await self._run_health_checks(target_server_ids) + return await self._run_health_checks([sid for sid in target_server_ids if sid not in checked_server_ids]) async def get_all_allowed_mcp_servers( self, @@ -7236,9 +7268,15 @@ class MCPServerManager: if not target_server_ids: return [] - tasks: Final = [self.health_check_server(server_id) for server_id in target_server_ids] - results: Final = await asyncio.gather(*tasks) - return [server for server in results if server is not None] + unique_server_ids: Final = tuple(dict.fromkeys(target_server_ids)) + batch_size: Final = 10 + batches: Final = [ + await asyncio.gather( + *(self.health_check_server(server_id) for server_id in unique_server_ids[offset : offset + batch_size]) + ) + for offset in range(0, len(unique_server_ids), batch_size) + ] + return [server for batch in batches for server in batch if server is not None] global_mcp_server_manager: Final[MCPServerManager] = MCPServerManager() diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3d44315341b..ded4db6d2aa 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -32168,6 +32168,7 @@ { "enum": [ "healthy", + "reachable", "unhealthy", "unknown" ], @@ -32178,7 +32179,7 @@ } ], "default": "unknown", - "description": "Health status: 'healthy', 'unhealthy', 'unknown'", + "description": "Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked)", "title": "Status" }, "subject_token_type": { @@ -35224,6 +35225,7 @@ { "enum": [ "healthy", + "reachable", "unhealthy", "unknown" ], @@ -35234,7 +35236,7 @@ } ], "default": "unknown", - "description": "Health status: 'healthy', 'unhealthy', 'unknown'", + "description": "Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked)", "title": "Status" }, "subject_token_type": { @@ -38095,6 +38097,18 @@ "description": "Server IDs to check. If not provided, checks all accessible servers.", "title": "Server Ids" } + }, + { + "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.", + "in": "query", + "name": "include_reachability", + "required": false, + "schema": { + "default": false, + "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.", + "title": "Include Reachability", + "type": "boolean" + } } ], "responses": { @@ -38389,6 +38403,18 @@ "title": "Server Id", "type": "string" } + }, + { + "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.", + "in": "query", + "name": "include_reachability", + "required": false, + "schema": { + "default": false, + "description": "Allow the 'reachable' status for responding servers whose authentication is unchecked.", + "title": "Include Reachability", + "type": "boolean" + } } ], "responses": { diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index d92443b104b..346b1a75ac6 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1296,6 +1296,12 @@ if MCP_AVAILABLE: return redacted_mcp_servers + def _mcp_health_status_for_response( + health_status: Literal["healthy", "reachable", "unhealthy", "unknown"] | None, + include_reachability: bool, + ) -> Literal["healthy", "reachable", "unhealthy", "unknown"] | None: + return "unknown" if health_status == "reachable" and not include_reachability else health_status + @router.get( "/server/health", description="Health check for MCP servers", @@ -1307,6 +1313,10 @@ if MCP_AVAILABLE: description="Server IDs to check. If not provided, checks all accessible servers.", ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + include_reachability: Annotated[ + bool, + Query(description="Allow the 'reachable' status for responding servers whose authentication is unchecked."), + ] = False, ): """ Perform health checks on one or more MCP servers. @@ -1331,21 +1341,31 @@ if MCP_AVAILABLE: if user_mcp_management_mode == "view_all" and not _is_restricted_virtual_key_request(user_api_key_dict): servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_unfiltered(server_ids=server_ids) - return [{"server_id": server.server_id, "status": server.status} for server in servers] + return [ + { + "server_id": server.server_id, + "status": _mcp_health_status_for_response(server.status, include_reachability), + } + for server in servers + ] auth_contexts: Final = await build_effective_auth_contexts(user_api_key_dict) - server_status_map: Final[dict[str, Literal["healthy", "unhealthy", "unknown"] | None]] = {} + server_status_map: Final[dict[str, Literal["healthy", "reachable", "unhealthy", "unknown"] | None]] = {} for auth_context in auth_contexts: servers = await global_mcp_server_manager.get_all_mcp_servers_with_health_and_teams( user_api_key_auth=auth_context, server_ids=server_ids, + checked_server_ids=frozenset(server_status_map), ) for server in servers: if server.server_id not in server_status_map: server_status_map[server.server_id] = server.status - return [{"server_id": server_id, "status": status} for server_id, status in server_status_map.items()] + return [ + {"server_id": server_id, "status": _mcp_health_status_for_response(status, include_reachability)} + for server_id, status in server_status_map.items() + ] @router.post( "/server/register", @@ -1615,6 +1635,10 @@ if MCP_AVAILABLE: request: Request, server_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + include_reachability: Annotated[ + bool, + Query(description="Allow the 'reachable' status for responding servers whose authentication is unchecked."), + ] = False, ): """ Get the info on the mcp server specified by the `server_id` @@ -1672,7 +1696,7 @@ if MCP_AVAILABLE: try: health_result: Final = await global_mcp_server_manager.health_check_server(server_id) # Update the server object with health check results - mcp_server.status = health_result.status if health_result.status else "unknown" + mcp_server.status = _mcp_health_status_for_response(health_result.status, include_reachability) or "unknown" mcp_server.last_health_check = health_result.last_health_check mcp_server.health_check_error = health_result.health_check_error except Exception as e: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 93b894f7645..fff4221f243 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -6,7 +6,13 @@ connection. The DB-backed per-user flow is exercised in higher-level tests in tests/mcp_tests. """ +from typing import Final +from unittest.mock import AsyncMock + import pytest +from respx import MockRouter + +from litellm.types.mcp_server.mcp_server_manager import MCPServer # Look up these names lazily on every access. Tests in this directory call # ``importlib.reload`` on the utils module to exercise registration logic, @@ -568,7 +574,7 @@ async def test_resolve_static_headers_user_value_wins_over_empty_global( assert headers == {"Authorization": "Bearer user-secret"} -# ── health-check skip for per-user-env-var-backed headers ────────────────── +# ── health-check reachability for per-user-env-var-backed headers ─────────── @pytest.mark.parametrize( @@ -615,32 +621,26 @@ def test_references_per_user_env_var(static_headers, env_vars, expected): @pytest.mark.asyncio -async def test_health_check_skips_servers_referencing_per_user_env_var( - mock_server, monkeypatch -): - """A userless health probe cannot fill per-user ${NAME} placeholders, so a - server whose static_headers reference one must report 'unknown' without - connecting. Otherwise it forwards the literal placeholder upstream, gets a - 401, and flips to 'unhealthy' even though real user calls succeed.""" +async def test_health_check_reaches_servers_without_forwarding_per_user_env_vars( + mock_server: MCPServer, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter +) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, ) - manager = MCPServerManager() + manager: Final = MCPServerManager() manager.registry[mock_server.server_id] = mock_server + create_client: Final = AsyncMock() + monkeypatch.setattr(manager, "_create_mcp_client", create_client) + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + route: Final = respx_mock.get(mock_server.url).respond(401) - created = [] + result: Final = await manager.health_check_server(mock_server.server_id) - async def fake_create_client(*args, **kwargs): - created.append((args, kwargs)) - raise RuntimeError("upstream rejected literal ${NAME}") - - monkeypatch.setattr(manager, "_create_mcp_client", fake_create_client) - - result = await manager.health_check_server(mock_server.server_id) - - assert created == [] - assert result.status == "unknown" + create_client.assert_not_called() + assert route.call_count == 1 + assert not {"x-db-url", "x-other"}.intersection(route.calls[0].request.headers) + assert result.status == "reachable" assert result.health_check_error is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 16bffa1a356..db476e86043 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -5,6 +5,7 @@ import json import logging import os import sys +from collections.abc import AsyncIterator from datetime import datetime from pathlib import Path from typing import Any, Dict, Final, Literal, Optional @@ -4894,69 +4895,258 @@ class TestMCPServerManager: assert result.last_health_check is not None @pytest.mark.asyncio - async def test_health_check_server_oauth2_skips_check(self): - """Test that health check is skipped for OAuth2 servers and returns unknown status""" - manager = MCPServerManager() - - # Mock OAuth2 server - server = MCPServer( + @pytest.mark.parametrize("oauth2_flow", [None, "authorization_code", "client_credentials"]) + async def test_health_check_server_oauth2_reports_reachability( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, oauth2_flow: Literal["authorization_code", "client_credentials"] | None + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + server: Final = MCPServer( server_id="oauth2-server", name="oauth2-server", transport=MCPTransport.http, auth_type=MCPAuth.oauth2, url="http://oauth2-server.com", + oauth2_flow=oauth2_flow, + client_id="client-id", + client_secret="stored-client-secret", + static_headers={"Authorization": "Bearer static-secret", "X-API-Key": "key-secret", "Cookie": "secret"}, ) - - manager.get_mcp_server_by_id = MagicMock(return_value=server) - - # _create_mcp_client should not be called for OAuth2 servers + manager.registry[server.server_id] = server manager._create_mcp_client = AsyncMock() + route: Final = respx_mock.get(server.url).respond(401) - # Perform health check - result = await manager.health_check_server("oauth2-server") + result: Final = await manager.health_check_server(server.server_id, mcp_auth_header="caller-secret") - # Verify that client was not created (health check was skipped) manager._create_mcp_client.assert_not_called() + assert result.status == "reachable" + assert result.health_check_error is None + assert result.last_health_check is not None + assert route.call_count == 1 + assert not {"authorization", "x-api-key", "cookie"}.intersection(route.calls[0].request.headers) - # Verify results - assert isinstance(result, LiteLLM_MCPServerTable) - assert result.server_id == "oauth2-server" - assert result.status == "unknown" + @pytest.mark.asyncio + @pytest.mark.parametrize("auth_type", [ + MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.basic, MCPAuth.authorization, MCPAuth.token, + MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag, MCPAuth.true_passthrough, MCPAuth.oauth_delegate, + ]) + @pytest.mark.parametrize("transport", [MCPTransport.http, MCPTransport.sse]) + @pytest.mark.parametrize("response_code", [200, 204, 302, 401, 403, 405, 503]) + async def test_health_check_without_credentials_accepts_any_http_response( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, auth_type: MCPAuthType, transport: Literal[MCPTransport.http, MCPTransport.sse], + response_code: int, + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="no-token-server", + name="no-token-server", + transport=transport, + auth_type=auth_type, + authentication_token=None, + url="http://no-token-server.com", + ) + manager.registry[server.server_id] = server + manager._create_mcp_client = AsyncMock() + route: Final = respx_mock.get(server.url).respond(response_code) + + result: Final = await manager.health_check_server(server.server_id) + + manager._create_mcp_client.assert_not_called() + assert route.call_count == 1 + assert result.status == "reachable" assert result.health_check_error is None assert result.last_health_check is not None @pytest.mark.asyncio - async def test_health_check_server_no_token_skips_check(self): - """Test that health check is skipped when auth_type is set but authentication_token is missing""" - manager = MCPServerManager() + @pytest.mark.parametrize("response_code", [200, 302]) + async def test_health_reachability_closes_sse_without_body_redirect_or_cookie_reuse( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, response_code: int + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + class UnreadBody(httpx.AsyncByteStream): + def __init__(self) -> None: + self.read = False + self.closed = False - # Mock server with auth_type but no authentication_token - server = MCPServer( - server_id="no-token-server", - name="no-token-server", - transport=MCPTransport.http, - auth_type=MCPAuth.bearer_token, - authentication_token=None, # No token - url="http://no-token-server.com", + async def __aiter__(self) -> AsyncIterator[bytes]: + self.read = True + yield b"secret SSE body" + + async def aclose(self) -> None: + self.closed = True + + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="streaming-health", name="streaming-health", transport=MCPTransport.sse, + auth_type=MCPAuth.oauth2, url="https://mcp.example.test/events", + ) + manager.registry[server.server_id] = server + bodies: Final = (UnreadBody(), UnreadBody()) + route: Final = respx_mock.get(server.url).mock(side_effect=[ + httpx.Response(response_code, stream=body, headers={ + "Content-Type": "text/event-stream", "Set-Cookie": "health=secret; Path=/", + "Location": "http://127.0.0.1/private", + }) for body in bodies + ]) + + first: Final = await manager.health_check_server(server.server_id) + second: Final = await manager.health_check_server(server.server_id) + + assert (first.status, second.status) == ("reachable", "reachable") + assert route.call_count == len(respx_mock.calls) == 2 + assert all(body.closed and not body.read for body in bodies) + assert all("cookie" not in call.request.headers for call in route.calls) + + @pytest.mark.asyncio + @pytest.mark.parametrize(("transport", "url"), [ + (MCPTransport.stdio, "https://mcp.example.test"), + (MCPTransport.http, None), (MCPTransport.http, ""), (MCPTransport.http, "not-a-url"), + (MCPTransport.http, "ftp://mcp.example.test"), + (MCPTransport.http, "https://user:secret@mcp.example.test"), + (MCPTransport.http, "https://mcp.example.test:bad/mcp"), + ]) + async def test_health_reachability_rejects_unprobeable_urls_without_requests( + self, respx_mock: MockRouter, transport: Literal[MCPTransport.http, MCPTransport.stdio], url: str | None + ) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="unprobeable", name="unprobeable", transport=transport, auth_type=MCPAuth.oauth2, url=url, + ) + manager.registry[server.server_id] = server + + result: Final = await manager.health_check_server(server.server_id) + + assert result.status == "unknown" + assert result.health_check_error and "secret" not in result.health_check_error + assert not respx_mock.calls + + @pytest.mark.asyncio + @pytest.mark.parametrize("failure", [ + httpx.ConnectError("TLS/connection failure with secret details"), + httpx.ReadTimeout("secret timeout details"), + httpx.RemoteProtocolError("secret malformed response"), + ]) + async def test_health_reachability_reports_no_response_without_secret_details( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, failure: httpx.RequestError + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="failed-health", name="failed-health", transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, is_byok=True, url="https://mcp.example.test/secret?token=secret", + ) + manager.registry[server.server_id] = server + route: Final = respx_mock.get(server.url).mock(side_effect=failure) + + result: Final = await manager.health_check_server(server.server_id) + + assert result.status == "unhealthy" + assert result.health_check_error and "secret" not in result.health_check_error + assert route.call_count == 1 + + @pytest.mark.asyncio + async def test_health_reachability_contains_ssl_setup_errors(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("SSL_SECURITY_LEVEL", "invalid-secret-cipher") + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="bad-tls", name="bad-tls", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, url="https://mcp.example.test", + ) + manager.registry[server.server_id] = server + + result: Final = await manager.health_check_server(server.server_id) + + assert result.status == "unhealthy" + assert result.health_check_error == "Reachability check failed (SSLError)" + + @pytest.mark.asyncio + @pytest.mark.parametrize("cancel", [False, True]) + async def test_health_reachability_timeout_and_cancellation_clean_up( + self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, cancel: bool + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + monkeypatch.setattr("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCP_HEALTH_CHECK_TIMEOUT", 0.1) + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="slow-health", name="slow-health", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, url="https://mcp.example.test/slow", + ) + manager.registry[server.server_id] = server + started: Final = asyncio.Event() + stopped: Final = asyncio.Event() + + async def slow_response(request: httpx.Request) -> httpx.Response: + started.set() + try: + await asyncio.Event().wait() + return httpx.Response(200) + finally: + stopped.set() + + respx_mock.get(server.url).mock(side_effect=slow_response) + task: Final = asyncio.create_task(manager.health_check_server(server.server_id)) + await asyncio.wait_for(started.wait(), timeout=1) + if cancel: + task.cancel() + result: Final = await task + + assert result.status == ("unknown" if cancel else "unhealthy") + assert result.health_check_error == ( + "Reachability check was cancelled" if cancel else "Reachability check timed out after 0.1 seconds" + ) + assert stopped.is_set() + + @pytest.mark.asyncio + @pytest.mark.parametrize("server_count", [0, 1, 10, 11, 25]) + @pytest.mark.parametrize("filtered", [False, True]) + async def test_bulk_health_checks_deduplicate_and_bound_upstream_requests( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter, server_count: int, filtered: bool + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + + class Probe: + def __init__(self) -> None: + self.active = 0 + self.peak = 0 + + async def respond(self, request: httpx.Request) -> httpx.Response: + self.active += 1 + self.peak = max(self.peak, self.active) + try: + await asyncio.sleep(0) + return httpx.Response(401) + finally: + self.active -= 1 + + manager: Final = MCPServerManager() + server_ids: Final = [f"health-{index}" for index in range(server_count)] + manager.registry = { + server_id: MCPServer( + server_id=server_id, name=server_id, transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, url=f"https://health.example.test/{server_id}", + ) + for server_id in server_ids + } + probe: Final = Probe() + route: Final = respx_mock.get(host="health.example.test").mock(side_effect=probe.respond) + requested_ids: Final = [*server_ids, *reversed(server_ids), *server_ids, "not-registered"] + + results: Final = ( + await manager.get_all_mcp_servers_with_health_and_teams( + user_api_key_auth=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + server_ids=requested_ids, + ) + if filtered + else await manager.get_all_mcp_servers_with_health_unfiltered(server_ids=requested_ids) ) - manager.get_mcp_server_by_id = MagicMock(return_value=server) - - # _create_mcp_client should not be called - manager._create_mcp_client = AsyncMock() - - # Perform health check - result = await manager.health_check_server("no-token-server") - - # Verify that client was not created (health check was skipped) - manager._create_mcp_client.assert_not_called() - - # Verify results - assert isinstance(result, LiteLLM_MCPServerTable) - assert result.server_id == "no-token-server" - assert result.status == "unknown" - assert result.health_check_error is None - assert result.last_health_check is not None + assert [(server.server_id, server.status) for server in results] == [ + (server_id, "reachable") for server_id in server_ids + ] + assert route.call_count == server_count + assert probe.peak == min(server_count, 10) + assert probe.active == 0 @pytest.mark.asyncio async def test_health_check_server_with_static_headers(self): @@ -5003,70 +5193,58 @@ class TestMCPServerManager: assert result.health_check_error is None @pytest.mark.asyncio - async def test_health_check_skips_passthrough_auth_with_authorization_header(self): - """Test that health check is skipped for servers with passthrough Authorization header""" - manager = MCPServerManager() - - # Mock server with auth_type=none and Authorization in extra_headers (passthrough auth) - server = MCPServer( + async def test_health_check_reaches_passthrough_auth_with_authorization_header( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + server: Final = MCPServer( server_id="github-server", name="github-server", transport=MCPTransport.http, auth_type=MCPAuth.none, authentication_token=None, url="http://github-server.com", - extra_headers=["Authorization"], # Passthrough auth configured + extra_headers=["Authorization"], ) - - manager.get_mcp_server_by_id = MagicMock(return_value=server) - - # _create_mcp_client should not be called (health check should be skipped) + manager.registry[server.server_id] = server manager._create_mcp_client = AsyncMock() + route: Final = respx_mock.get(server.url).respond(401) - # Perform health check - result = await manager.health_check_server("github-server") + result: Final = await manager.health_check_server(server.server_id) - # Verify that client was not created (health check was skipped) manager._create_mcp_client.assert_not_called() - - # Verify results - assert isinstance(result, LiteLLM_MCPServerTable) - assert result.server_id == "github-server" - assert result.status == "unknown" + assert route.call_count == 1 + assert "authorization" not in route.calls[0].request.headers + assert result.status == "reachable" assert result.health_check_error is None assert result.last_health_check is not None @pytest.mark.asyncio - async def test_health_check_skips_passthrough_auth_with_api_key_header(self): - """Test that health check is skipped for servers with passthrough x-api-key header""" - manager = MCPServerManager() - - # Mock server with auth_type=none and x-api-key in extra_headers - server = MCPServer( + async def test_health_check_reaches_passthrough_auth_with_api_key_header( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + server: Final = MCPServer( server_id="sourcegraph-server", name="sourcegraph-server", transport=MCPTransport.http, auth_type=MCPAuth.none, authentication_token=None, url="http://sourcegraph-server.com", - extra_headers=["x-api-key"], # Passthrough auth configured + extra_headers=["x-api-key"], ) - - manager.get_mcp_server_by_id = MagicMock(return_value=server) - - # _create_mcp_client should not be called + manager.registry[server.server_id] = server manager._create_mcp_client = AsyncMock() + route: Final = respx_mock.get(server.url).respond(403) - # Perform health check - result = await manager.health_check_server("sourcegraph-server") + result: Final = await manager.health_check_server(server.server_id) - # Verify that client was not created (health check was skipped) manager._create_mcp_client.assert_not_called() - - # Verify results - assert isinstance(result, LiteLLM_MCPServerTable) - assert result.server_id == "sourcegraph-server" - assert result.status == "unknown" + assert route.call_count == 1 + assert "x-api-key" not in route.calls[0].request.headers + assert result.status == "reachable" assert result.health_check_error is None assert result.last_health_check is not None @@ -9239,16 +9417,19 @@ class TestRegistryTableConversionPreservesEnvVars: self._assert_env_vars_round_tripped(table) @pytest.mark.asyncio - async def test_health_check_server_preserves_env_vars(self): - # OAuth2 without client credentials needs a per-user token, so the - # health check is skipped (no network) and we exercise the table - # construction path directly. - manager = MCPServerManager() - server = self._server_with_env_vars() + async def test_health_check_server_preserves_env_vars( + self, monkeypatch: pytest.MonkeyPatch, respx_mock: MockRouter + ) -> None: + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = MCPServerManager() + server: Final = self._server_with_env_vars() assert server.requires_per_user_auth is True manager.registry[server.server_id] = server - table = await manager.health_check_server(server.server_id) + route: Final = respx_mock.get(server.url).respond(401) + table: Final = await manager.health_check_server(server.server_id) self._assert_env_vars_round_tripped(table) + assert route.call_count == 1 + assert "x-db-url" not in route.calls[0].request.headers class TestHealthCheckInterpolatesGlobalEnvVars: diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index a9ec575e99b..8aa16b817fc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -9,11 +9,12 @@ from contextlib import ExitStack, contextmanager from dataclasses import dataclass, field from datetime import datetime, timedelta from types import SimpleNamespace -from typing import Final, List, Optional, cast +from typing import Final, List, Literal, Optional, cast from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter, ValidationError from respx import MockRouter from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient @@ -4311,6 +4312,170 @@ async def test_health_discovery_respects_route_restricted_key_grants( assert all(row["status"] == expected_status for row in result) +@pytest.mark.asyncio +@pytest.mark.respx(assert_all_called=False) +@pytest.mark.parametrize("include_reachability", [False, True]) +@pytest.mark.parametrize( + ("requested", "expected"), + [ + (None, ("shared", "first", "second")), + ((), ("shared", "first", "second")), + (("shared", "shared", "denied"), ("shared",)), + (("second", "first"), ("first", "second")), + (("denied",), ()), + ], +) +async def test_health_checks_probe_shared_servers_once_across_auth_contexts( + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + requested: tuple[str, ...] | None, + expected: tuple[str, ...], + include_reachability: bool, +) -> None: + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = mcp_server_manager.MCPServerManager() + manager.registry = { + server_id: MCPServer( + server_id=server_id, + name=server_id, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url=f"https://mcp.example.test/{server_id}", + ) + for server_id in ("shared", "first", "second", "denied") + } + routes: Final = { + server_id: respx_mock.get(server.url).respond(401) + for server_id, server in manager.registry.items() + } + contexts: Final = [ + UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key=f"test-health-{index}", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id=f"health-{index}", mcp_servers=list(grants) + ), + ) + for index, grants in enumerate((("shared", "first"), ("shared", "second"))) + ] + with ( + patch.object( + mgmt_endpoints, "global_mcp_server_manager", manager + ), + patch.object( + mcp_server_manager, "global_mcp_server_manager", manager + ), + patch.object( + mgmt_endpoints, "build_effective_auth_contexts", AsyncMock(return_value=contexts) + ), + patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "restricted"}), + ): + result: Final = await mgmt_endpoints.health_check_servers( + server_ids=list(requested) if requested is not None else None, + user_api_key_dict=contexts[0], + include_reachability=include_reachability, + ) + + expected_status: Final = "reachable" if include_reachability else "unknown" + assert sorted(result, key=lambda row: row["server_id"]) == [ + {"server_id": server_id, "status": expected_status} for server_id in sorted(expected) + ] + if requested: + assert [row["server_id"] for row in result] == list(expected) + assert {server_id: route.call_count for server_id, route in routes.items()} == { + server_id: int(server_id in expected) for server_id in routes + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["restricted", "view_all"]) +@pytest.mark.parametrize("detail", [False, True]) +@pytest.mark.parametrize("flag", [None, "false", "true"]) +async def test_health_reachability_requires_explicit_api_opt_in( + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + mode: str, + detail: bool, + flag: str | None, +) -> None: + from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + class HealthResponse(BaseModel): + server_id: str + status: str | None + + class LegacyHealthResponse(BaseModel): + server_id: str + status: Literal["healthy", "unhealthy", "unknown"] | None + + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + manager: Final = mcp_server_manager.MCPServerManager() + server: Final = MCPServer( + server_id="health-compatibility", + name="health-compatibility", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + url="https://mcp.example.test/mcp", + ) + manager.registry[server.server_id] = server + route: Final = respx_mock.get(server.url).respond(401) + caller: Final = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="test-health-compatibility", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="health-compatibility", mcp_servers=[server.server_id] + ), + ) + + def authenticated_caller() -> UserAPIKeyAuth: + return caller + + app: Final = FastAPI() + app.include_router(mgmt_endpoints.router) + app.dependency_overrides[mgmt_endpoints.user_api_key_auth] = authenticated_caller + suffix: Final = server.server_id if detail else "health" + query: Final = {} if flag is None else {"include_reachability": flag} + with ( + patch.object( # test-quality-ok: TQ008 inject the real registry into the legacy route binding + mgmt_endpoints, "global_mcp_server_manager", manager + ), + patch.object( # test-quality-ok: TQ008 permission resolution uses the shared registry + mcp_server_manager, "global_mcp_server_manager", manager + ), + patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": mode}), + patch.object( # test-quality-ok: TQ008 select the config-backed detail path without a database + mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock() + ), + patch.object( # test-quality-ok: TQ008 a missing database row falls back to the real registry + mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=None) + ), + ): + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://gateway") as client: + response: Final = await client.get(f"/v1/mcp/server/{suffix}", params=query) + + assert response.status_code == 200, response.text + rows: Final = ( + [HealthResponse.model_validate_json(response.content)] + if detail else TypeAdapter(list[HealthResponse]).validate_json(response.content) + ) + expected_status: Final = "reachable" if flag == "true" else "unknown" + assert [row.model_dump() for row in rows] == [{"server_id": server.server_id, "status": expected_status}] + assert route.call_count == 1 + legacy_parser: Final = ( + LegacyHealthResponse.model_validate_json + if detail else TypeAdapter(list[LegacyHealthResponse]).validate_json + ) + if flag == "true": + with pytest.raises(ValidationError, match="literal_error"): + legacy_parser(response.content) + else: + legacy_parser(response.content) + + class TestMCPRegistryEndpoint: def test_registry_returns_404_when_flag_missing(self): client = create_mcp_router_test_client() diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index 71c2e107774..100b0ea93d3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -1,5 +1,6 @@ import React from "react"; import { fireEvent, render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { describe, it, expect, vi, afterEach } from "vitest"; import MCPServerCard from "./MCPServerCard"; import type { MCPServer } from "@/components/mcp_tools/types"; @@ -18,6 +19,19 @@ function renderCard(overrides: Partial) { render(); } +describe("MCPServerCard health", () => { + it("explains that reachable does not verify authentication or tools", async () => { + const user = userEvent.setup(); + renderCard({ status: "reachable", oauth2_flow: "authorization_code" }); + + await user.hover(screen.getByText("Reachable")); + + expect(await screen.findByText("Server responded. Authentication and tools were not checked")).toBeInTheDocument(); + expect(screen.queryByText("No health data")).not.toBeInTheDocument(); + expect(screen.queryByText("Healthy")).not.toBeInTheDocument(); + }); +}); + describe("MCPServerCard OAuth flow indicator", () => { it("shows the 'OAuth flow not set' badge for an oauth2 server with no oauth2_flow", () => { renderCard({ auth_type: "oauth2", oauth2_flow: null }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index 42fb95d5951..775809e3670 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -11,7 +11,7 @@ import { } from "@/components/ui/dropdown-menu"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { cn } from "@/lib/cva.config"; -import { AUTH_TYPE, type MCPServer } from "@/components/mcp_tools/types"; +import { AUTH_TYPE, MCP_REACHABLE_DESCRIPTION, type MCPServer } from "@/components/mcp_tools/types"; import { Logo } from "@/components/molecules/logo/Logo"; import { getMaskedAndFullUrl } from "./utils"; @@ -33,6 +33,7 @@ interface MCPServerCardProps { const HEALTH_TONE: Record = { healthy: { dot: "bg-success" }, + reachable: { dot: "bg-info" }, unhealthy: { dot: "bg-destructive" }, unknown: { dot: "bg-border" }, }; @@ -332,6 +333,7 @@ const HealthChip: FC = ({ ); } + const hasHealthData = Boolean(lastCheck || error || status === "reachable"); return ( = ({ />
Health: {status}
+ {status === "reachable" &&
{MCP_REACHABLE_DESCRIPTION}
} {lastCheck &&
Last check: {new Date(lastCheck).toLocaleString()}
} {error && (
@@ -362,7 +365,7 @@ const HealthChip: FC = ({
{error}
)} - {!lastCheck && !error &&
No health data
} + {!hasHealthData &&
No health data
} {onRecheck &&
Click to recheck
}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx index 1217d878489..9a3ff0cc6cb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx @@ -113,12 +113,14 @@ describe("compareServers", () => { it("sorts health before recency and display name", () => { const servers: MCPServer[] = [ { ...server("healthy", "aaa", "2026-03-01T00:00:00Z"), status: "healthy" }, + { ...server("reachable", "aaa", "2026-04-01T00:00:00Z"), status: "reachable" }, { ...server("unknown", "bbb", "2026-02-01T00:00:00Z"), status: "unknown" }, { ...server("unhealthy", "zzz", "2026-01-01T00:00:00Z"), status: "unhealthy" }, ]; expect(servers.sort((a, b) => compareServers(a, b, "health")).map((s) => s.server_id)).toEqual([ "unhealthy", "unknown", + "reachable", "healthy", ]); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index b4b7ab6b3c8..56a18e0ca4c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -61,7 +61,8 @@ const SORT_OPTIONS: { value: SortKey; label: string }[] = [ const HEALTH_RANK: Record = { unhealthy: 0, unknown: 1, - healthy: 2, + reachable: 2, + healthy: 3, }; const compareByName = (a: MCPServer, b: MCPServer): number => { @@ -191,7 +192,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i const healthStatus = healthMap.get(server.server_id); return { ...server, - status: healthStatus ? (healthStatus as "healthy" | "unhealthy" | "unknown") : server.status, + status: healthStatus ? (healthStatus as MCPServer["status"]) : server.status, }; }); }, [mcpServers, healthStatuses]); diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx index e32c861f13a..1032b03a3ce 100644 --- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx @@ -28,10 +28,10 @@ const mockServer: MCPServerData = { env: {}, }; -function renderTable(onServerClick = vi.fn()) { +function renderTable(onServerClick = vi.fn(), servers = [mockServer]) { render( server.server_id} sortingMode="client" @@ -42,6 +42,15 @@ function renderTable(onServerClick = vi.fn()) { } describe("getMCPHubTableColumns", () => { + it("explains the limited check for a reachable server", async () => { + const user = userEvent.setup(); + renderTable(vi.fn(), [{ ...mockServer, status: "reachable" }]); + + await user.hover(screen.getByText("reachable")); + + expect(await screen.findByText("Server responded. Authentication and tools were not checked")).toBeInTheDocument(); + }); + it("renders the server row", () => { renderTable(); expect(screen.getByText("exa_test")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx index 6a1ede11201..20a14bcb476 100644 --- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx @@ -4,6 +4,7 @@ import { ColumnDef } from "@tanstack/react-table"; import { Copy, Info, MoreHorizontal } from "lucide-react"; import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { MCP_REACHABLE_DESCRIPTION } from "@/components/mcp_tools/types"; import { IdentityCell, StatusBadge, type StatusTone } from "@/components/shared/table_cells"; import { Badge } from "@/components/ui/badge"; import { buttonVariants } from "@/components/ui/button"; @@ -49,6 +50,7 @@ const STATUS_TONES: Record = { inactive: "error", unknown: "neutral", healthy: "success", + reachable: "info", unhealthy: "error", }; @@ -150,7 +152,11 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps) enableSorting: true, sortingFn: "alphanumeric", cell: ({ row }) => ( - + ), }, { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 47df369fb8a..be7d39616ca 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -405,6 +405,8 @@ export interface MCPToolsViewerProps { extraHeaders?: string[] | null; } +export const MCP_REACHABLE_DESCRIPTION = "Server responded. Authentication and tools were not checked"; + export interface MCPServer { server_id: string; is_config?: boolean; @@ -435,7 +437,7 @@ export interface MCPServer { updated_by: string; extra_headers?: string[] | null; static_headers?: Record | null; - status?: "healthy" | "unhealthy" | "unknown"; + status?: "healthy" | "reachable" | "unhealthy" | "unknown"; last_health_check?: string | null; health_check_error?: string | null; teams?: Team[]; diff --git a/ui/litellm-dashboard/src/components/networking.test.ts b/ui/litellm-dashboard/src/components/networking.test.ts index e14f1939ee1..b964231804e 100644 --- a/ui/litellm-dashboard/src/components/networking.test.ts +++ b/ui/litellm-dashboard/src/components/networking.test.ts @@ -706,6 +706,30 @@ describe("testMCPToolsListRequest auth headers", () => { }); }); +describe("fetchMCPServerHealth", () => { + const originalFetch = global.fetch; + + afterEach(() => { + global.fetch = originalFetch; + }); + + it.each([{ serverIds: undefined }, { serverIds: [] }, { serverIds: ["server one", "server&two"] }])( + "opts into reachability while preserving requested servers: $serverIds", + async ({ serverIds }) => { + const mockFetch = vi.fn().mockResolvedValue(new Response("[]", { status: 200 })); + global.fetch = mockFetch; + + await Networking.fetchMCPServerHealth("test-token", serverIds); + + expect(mockFetch).toHaveBeenCalledOnce(); + const url = new URL(String(mockFetch.mock.calls[0][0]), "http://localhost"); + expect(url.pathname).toMatch(/\/v1\/mcp\/server\/health$/); + expect(url.searchParams.get("include_reachability")).toBe("true"); + expect(url.searchParams.getAll("server_ids")).toEqual(serverIds ?? []); + }, + ); +}); + describe("getAutoRouterClassifierDefaultPromptCall", () => { const originalFetch = global.fetch; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index c4008d512c3..e1271b9151f 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -4968,6 +4968,7 @@ export const fetchMCPServerHealth = async (accessToken: string, serverIds?: stri return await apiClient.get(`/v1/mcp/server/health`, { accessToken, query: { + include_reachability: true, server_ids: serverIds && serverIds.length > 0 ? serverIds : undefined, }, }); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 0500aeb95c8..b5d515bd1fb 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -32504,10 +32504,10 @@ export interface components { } | null; /** * Status - * @description Health status: 'healthy', 'unhealthy', 'unknown' + * @description Health status: 'healthy', 'unhealthy', 'unknown', or 'reachable' (requires include_reachability=true; authentication and tools unchecked) * @default unknown */ - status: ("healthy" | "unhealthy" | "unknown") | null; + status: ("healthy" | "reachable" | "unhealthy" | "unknown") | null; /** Subject Token Type */ subject_token_type?: string | null; /** Submitted At */ @@ -72212,6 +72212,8 @@ export interface operations { query?: { /** @description Server IDs to check. If not provided, checks all accessible servers. */ server_ids?: string[] | null; + /** @description Allow the 'reachable' status for responding servers whose authentication is unchecked. */ + include_reachability?: boolean; }; header?: never; path?: never; @@ -72363,7 +72365,10 @@ export interface operations { }; fetch_mcp_server_v1_mcp_server__server_id__get: { parameters: { - query?: never; + query?: { + /** @description Allow the 'reachable' status for responding servers whose authentication is unchecked. */ + include_reachability?: boolean; + }; header?: never; path: { server_id: string; From 5d777c16d9e59c690886d978f7546b130bd2432d Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 26 Sep 2026 16:53:02 -0700 Subject: [PATCH 08/51] fix(mcp): align hub publication status and controls (#43241) * fix(mcp): align hub publication status and controls * refactor(mcp): keep hub visibility guard outside table rendering --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> --- cookbook/litellm_proxy_server/mcp/README.md | 37 ++++ .../mcp_server/mcp_server_manager.py | 24 +-- .../mcp_management_endpoints.py | 35 ++-- .../public_endpoints/public_endpoints.py | 14 +- .../mcp_server/test_mcp_server_manager.py | 54 ++++++ .../test_mcp_management_endpoints.py | 156 ++++++++++++++++ .../public_endpoints/test_public_endpoints.py | 66 +++++-- .../_components/MCPPermissionManagement.tsx | 4 +- .../_components/MCPServerCard.test.tsx | 12 ++ .../mcp-servers/_components/MCPServerCard.tsx | 19 +- .../_components/mcp_server_view.test.tsx | 11 +- .../_components/mcp_server_view.tsx | 21 +-- .../mcp-servers/_components/utils.test.tsx | 29 +++ .../mcp-servers/_components/utils.tsx | 35 +++- .../AIHub/MCPHubTableColumns.test.tsx | 19 +- .../components/AIHub/MCPHubTableColumns.tsx | 6 +- .../components/AIHub/ModelHubTable.test.tsx | 30 ++- .../src/components/AIHub/ModelHubTable.tsx | 15 +- .../AIHub/forms/MakeMCPPublicForm.test.tsx | 174 +++++++++++++----- .../AIHub/forms/MakeMCPPublicForm.tsx | 121 ++++++++---- .../src/components/mcp_tools/types.tsx | 2 + 21 files changed, 720 insertions(+), 164 deletions(-) create mode 100644 cookbook/litellm_proxy_server/mcp/README.md diff --git a/cookbook/litellm_proxy_server/mcp/README.md b/cookbook/litellm_proxy_server/mcp/README.md new file mode 100644 index 00000000000..aeee0719019 --- /dev/null +++ b/cookbook/litellm_proxy_server/mcp/README.md @@ -0,0 +1,37 @@ +# Publish MCP servers in the AI Hub + +Set `litellm_settings.public_mcp_servers` to the concrete IDs of the servers you want listed in the public AI Hub. Pin `server_id` in each configuration entry so the publication list stays stable across deployments + +```yaml +mcp_servers: + documentation: + server_id: documentation-mcp + url: https://mcp.example.com/mcp + transport: http + available_on_public_internet: true + +litellm_settings: + public_mcp_hub_strict_whitelist: true + public_mcp_servers: + - documentation-mcp +``` + +Use `documentation-mcp`, the `server_id`, in the publication list. The configuration key `documentation`, display names, and aliases are not publication IDs. Database-created servers use the ID returned by `/v1/mcp/server` + +The dashboard's **AI Hub > MCP Hub > Manage MCP Hub Visibility** dialog edits this same list. Its YAML example includes the selected server IDs. With database-backed configuration (`store_model_in_db: true`), a value declared in YAML is owned by that file: edit the file and reload, or remove that key from YAML to let the dashboard manage it in the database. File-backed deployments can save the list directly to their configuration file + +To remove all explicit entries, save an empty selection in the dialog or configure: + +```yaml +litellm_settings: + public_mcp_hub_strict_whitelist: true + public_mcp_servers: [] +``` + +## Hub listing and network access + +The **Hub listing** column in AI Hub identifies servers that appear in `/public/mcp_hub`. The dashboard derives this status from the current registry and publication settings. Setting `mcp_info.is_public` on a server does not publish it; that response field is derived metadata. `mcp_info.is_public_explicit` identifies registered servers included in the explicit publication list + +Gateway cards and server details show **All Networks** when `available_on_public_internet` is enabled or the server is explicitly published in `public_mcp_servers`. They show **Internal Only** when both are false. The per-server flag defaults to `true`; explicit publication overrides a disabled flag for compatibility. Older proxies that omit the metadata needed to determine access show **Unknown**. These labels describe allowed client IPs; authentication and tool permissions still apply + +The default `public_mcp_hub_strict_whitelist: true` lists only registered servers in `public_mcp_servers`. Legacy mode (`false`) additionally lists registered servers with `available_on_public_internet: true`. In legacy mode, clearing the explicit publication list leaves these automatically listed servers visible. Enable strict mode when the publication list should fully determine hub visibility diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d0d9100971d..31896d9ddc5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -6799,6 +6799,16 @@ class MCPServerManager: return server return None + @staticmethod + def _is_public_mcp_server(server: MCPServer, public_ids: Container[str]) -> bool: + return server.server_id in public_ids or ( + not litellm.public_mcp_hub_strict_whitelist and server.available_on_public_internet + ) + + def is_mcp_server_public(self, server_id: str) -> bool: + server: Final = self.registry.get(server_id) or self.config_mcp_servers.get(server_id) + return server is not None and self._is_public_mcp_server(server, litellm.public_mcp_servers or ()) + def get_public_mcp_servers(self) -> list[MCPServer]: """ Return the MCP servers published to the AI Hub via /v1/mcp/make_public. @@ -6816,18 +6826,8 @@ class MCPServerManager: deployments that relied on the OR-with-default semantics; will be removed in a future release. """ - if litellm.public_mcp_hub_strict_whitelist: - if litellm.public_mcp_servers is None: - return [] - public_ids = set(litellm.public_mcp_servers) - return [server for server in self.get_registry().values() if server.server_id in public_ids] - - public_ids = set(litellm.public_mcp_servers or []) - return [ - server - for server in self.get_registry().values() - if server.available_on_public_internet or server.server_id in public_ids - ] + public_ids: Final = frozenset(litellm.public_mcp_servers or ()) + return [server for server in self.get_registry().values() if self._is_public_mcp_server(server, public_ids)] def expand_permission_list(self, identifiers: list[str]) -> list[str]: """ diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 346b1a75ac6..deb0e00ff9b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -650,7 +650,16 @@ if MCP_AVAILABLE: if hasattr(redacted_server, "credentials"): setattr(redacted_server, "credentials", _preserved_admin_config_credentials(redacted_server.credentials)) - return redacted_server + is_public: Final = global_mcp_server_manager.is_mcp_server_public(redacted_server.server_id) + return redacted_server.model_copy( + update={ + "mcp_info": { + **(redacted_server.mcp_info or {}), + "is_public": is_public, + "is_public_explicit": is_public and redacted_server.server_id in (litellm.public_mcp_servers or ()), + } + } + ) def _preserved_admin_config_credentials( credentials: "MCPCredentials | str | None", @@ -832,10 +841,10 @@ if MCP_AVAILABLE: sanitized.updated_at = None # `mcp_info` is arbitrary metadata; keep only an explicit safe subset. - is_public = False - if isinstance(sanitized.mcp_info, dict): - is_public = bool(sanitized.mcp_info.get("is_public")) - sanitized.mcp_info = {"is_public": True} if is_public else None + sanitized.mcp_info = { + "is_public": (sanitized.mcp_info or {}).get("is_public") is True, + "is_public_explicit": (sanitized.mcp_info or {}).get("is_public_explicit") is True, + } return sanitized @@ -1260,14 +1269,6 @@ if MCP_AVAILABLE: for server in redacted_mcp_servers: server.connected_app_reachable = server.server_id in reachable_ids - # augment the mcp servers with public status - if litellm.public_mcp_servers is not None: - for server in redacted_mcp_servers: - if server.server_id in litellm.public_mcp_servers: - if server.mcp_info is None: - server.mcp_info = {} - server.mcp_info["is_public"] = True - # Annotate has_user_credential for BYOK servers (single batched query) from litellm.proxy.proxy_server import prisma_client as _byok_prisma_client @@ -3041,9 +3042,6 @@ if MCP_AVAILABLE: }, ) - if litellm.public_mcp_servers is None: - litellm.public_mcp_servers = [] - for server_id in request.mcp_server_ids: server = global_mcp_server_manager.get_mcp_server_by_id(server_id=server_id) if server is None: @@ -3052,16 +3050,15 @@ if MCP_AVAILABLE: detail=f"MCP Server with ID {server_id} not found", ) - litellm.public_mcp_servers = request.mcp_server_ids - # Update config with new settings if "litellm_settings" not in config or config["litellm_settings"] is None: config["litellm_settings"] = {} - config["litellm_settings"]["public_mcp_servers"] = litellm.public_mcp_servers + config["litellm_settings"]["public_mcp_servers"] = request.mcp_server_ids # Save the updated config await proxy_config.save_config(new_config=config) + litellm.public_mcp_servers = request.mcp_server_ids verbose_proxy_logger.debug( "Updated public mcp servers to: %s by user: %s", litellm.public_mcp_servers, user_api_key_dict.user_id diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index 26a5c44fce1..bba5ef681d0 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -300,7 +300,19 @@ async def get_mcp_servers(): ) public_mcp_servers: Final = global_mcp_server_manager.get_public_mcp_servers() - return [MCPPublicServer.model_validate(server.model_dump()) for server in public_mcp_servers] + return [ + MCPPublicServer.model_validate( + { + **server.model_dump(), + "mcp_info": { + **(server.mcp_info or {}), + "is_public": True, + "is_public_explicit": server.server_id in (litellm.public_mcp_servers or ()), + }, + } + ) + for server in public_mcp_servers + ] @router.get( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index db476e86043..70ef4312f4c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -9566,6 +9566,60 @@ class TestGetPublicMCPServers: manager.config_mcp_servers[s.server_id] = s return manager + @pytest.mark.parametrize("registered_in", ("config", "database", "both", "neither")) + @pytest.mark.parametrize("public_ids", (None, [], ["server-id"], ["server-alias"], ["Server Name"])) + @pytest.mark.parametrize( + "strict,network_access,implicitly_public", + ((True, True, False), (True, False, False), (False, True, True), (False, False, False)), + ) + def test_public_status_agrees_with_hub_membership( + self, + registered_in: Literal["config", "database", "both", "neither"], + public_ids: list[str] | None, + strict: bool, + network_access: bool, + implicitly_public: bool, + ) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="server-id", + name="server-alias", + alias="server-alias", + server_name="Server Name", + transport=MCPTransport.http, + available_on_public_internet=network_access, + mcp_info={"is_public": True, "description": "Preserve custom metadata"}, + ) + config_server: Final = ( + server.model_copy(update={"available_on_public_internet": not network_access}) + if registered_in == "both" + else server + ) + manager.config_mcp_servers = ( + {server.server_id: config_server} if registered_in in ("config", "both") else {} + ) + manager.registry = {server.server_id: server} if registered_in in ("database", "both") else {} + original_server: Final = server.model_dump() + original_config_server: Final = config_server.model_dump() + expected_public: Final = registered_in != "neither" and ( + public_ids == [server.server_id] or implicitly_public + ) + + with ( + patch("litellm.public_mcp_servers", public_ids), + patch("litellm.public_mcp_hub_strict_whitelist", strict), + ): + public_servers: Final = manager.get_public_mcp_servers() + assert manager.is_mcp_server_public(server.server_id) is expected_public + assert [item.server_id for item in public_servers] == ( + [server.server_id] if expected_public else [] + ) + assert manager.is_mcp_server_public("server-alias") is False + assert manager.is_mcp_server_public("missing-server") is False + + assert server.model_dump() == original_server + assert config_server.model_dump() == original_config_server + @patch("litellm.public_mcp_servers", None) def test_returns_empty_when_whitelist_is_none(self): """No /make_public call yet → hub returns nothing, regardless of diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 8aa16b817fc..11b3dcf54bc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, LiteLLM_MCPServerTable, LitellmUserRoles, + MakeMCPServersPublicRequest, MCPTransport, MCPUserCredentialResponse, NewMCPServerRequest, @@ -154,6 +155,161 @@ def patch_proxy_general_settings(settings: dict): ) +@pytest.mark.asyncio +@pytest.mark.parametrize("from_db", (False, True)) +@pytest.mark.parametrize( + "strict,explicit,expected_public", + ((True, True, True), (True, False, False), (False, False, True)), +) +async def test_mcp_publication_list_and_detail_derive_current_status( + from_db: bool, strict: bool, explicit: bool, expected_public: bool +) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="publication-server", + name="publication-server", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + available_on_public_internet=True, + mcp_info={ + "is_public": not expected_public, + "is_public_explicit": not explicit, + "description": "Keep this description", + }, + ) + manager.registry = {server.server_id: server} if from_db else {} + manager.config_mcp_servers = {} if from_db else {server.server_id: server} + record: Final = manager._build_mcp_server_table(server) + original_metadata: Final = dict(server.mcp_info or {}) + admin: Final = generate_mock_user_api_key_auth() + + with ( + patch("litellm.public_mcp_servers", [server.server_id] if explicit else []), + patch("litellm.public_mcp_hub_strict_whitelist", strict), + patch.object(mgmt_endpoints, "global_mcp_server_manager", manager), + patch.object(mgmt_endpoints, "get_prisma_client_or_throw", return_value=MagicMock()), + patch.object(mgmt_endpoints, "get_mcp_server", AsyncMock(return_value=record if from_db else None)), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.general_settings", {"user_mcp_management_mode": "view_all"}), + ): + listing: Final = await mgmt_endpoints.fetch_all_mcp_servers( + user_api_key_dict=admin, team_id=None, connected_app_view=False + ) + detail: Final = await mgmt_endpoints.fetch_mcp_server( + request=_make_mock_request(), server_id=server.server_id, user_api_key_dict=admin + ) + assert len(listing) == 1 + for projected in (listing[0], detail): + assert projected.mcp_info == { + "is_public": expected_public, + "is_public_explicit": explicit, + "description": "Keep this description", + } + assert bool(manager.get_public_mcp_servers()) is expected_public + + assert server.mcp_info == original_metadata + assert record.mcp_info == original_metadata + + +@pytest.mark.parametrize("approval_status", ("pending_review", "rejected", "draft", "active")) +@pytest.mark.parametrize("strict", (False, True)) +def test_mcp_publication_projection_excludes_unregistered_lifecycle_records( + approval_status: str, strict: bool +) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + record: Final = LiteLLM_MCPServerTable( + server_id="unregistered-server", + transport=MCPTransport.http, + approval_status=approval_status, + credentials={"auth_value": "test-secret"}, + available_on_public_internet=True, + mcp_info={"is_public": True, "is_public_explicit": True}, + ) + original: Final = record.model_dump() + with ( + patch("litellm.public_mcp_servers", [record.server_id]), + patch("litellm.public_mcp_hub_strict_whitelist", strict), + patch.object(mgmt_endpoints, "global_mcp_server_manager", MCPServerManager()), + ): + for project in ( + mgmt_endpoints._redact_mcp_credentials, + mgmt_endpoints._sanitize_mcp_server_for_non_admin, + mgmt_endpoints._sanitize_mcp_server_for_virtual_key, + ): + projected: Final = project(record) + assert projected.mcp_info == {"is_public": False, "is_public_explicit": False} + assert projected.credentials is None + assert record.model_dump() == original + + +@pytest.mark.asyncio +@pytest.mark.parametrize("previous_ids", (None, ["old-server"])) +@pytest.mark.parametrize( + "selected_ids,save_error,role,error_status", + ( + (["new-server"], None, LitellmUserRoles.PROXY_ADMIN, None), + ([], None, LitellmUserRoles.PROXY_ADMIN, None), + (["new-server"], HTTPException(400, "Owned by config file"), LitellmUserRoles.PROXY_ADMIN, 400), + (["new-server"], RuntimeError("Database write failed"), LitellmUserRoles.PROXY_ADMIN, 500), + (["missing-server"], None, LitellmUserRoles.PROXY_ADMIN, 404), + (["new-server"], None, LitellmUserRoles.INTERNAL_USER, 403), + ), +) +async def test_mcp_publication_updates_runtime_only_after_successful_save( + previous_ids: list[str] | None, + selected_ids: list[str], + save_error: HTTPException | RuntimeError | None, + role: LitellmUserRoles, + error_status: int | None, +) -> None: + import litellm + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager: Final = MCPServerManager() + server: Final = generate_mock_mcp_server_config_record(server_id="new-server") + manager.config_mcp_servers = {server.server_id: server} + expected_config: Final = {"litellm_settings": {"drop_params": True, "public_mcp_servers": selected_ids}} + + async def save_config(new_config: Mapping[str, object]) -> None: + assert litellm.public_mcp_servers is previous_ids + assert new_config == expected_config + if save_error is not None: + raise save_error + + save: Final = AsyncMock(side_effect=save_config) + proxy_config: Final = SimpleNamespace( + get_config=AsyncMock(return_value={"litellm_settings": {"drop_params": True}}), + save_config=save, + ) + request: Final = MakeMCPServersPublicRequest(mcp_server_ids=selected_ids) + caller: Final = generate_mock_user_api_key_auth(user_role=role) + with ( + patch("litellm.public_mcp_servers", previous_ids), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + if error_status is None: + response: Final = await mgmt_endpoints.make_mcp_servers_public(request, caller) + assert response["public_mcp_servers"] == selected_ids + assert litellm.public_mcp_servers == selected_ids + else: + with pytest.raises(HTTPException) as error: + await mgmt_endpoints.make_mcp_servers_public(request, caller) + assert error.value.status_code == error_status + assert litellm.public_mcp_servers is previous_ids + + if error_status in (403, 404): + save.assert_not_awaited() + else: + save.assert_awaited_once_with(new_config=expected_config) + + class TestMCPCredentialsTokenExchangeProfile: """token_exchange_profile must be a declared MCPCredentials field so the management API can persist the entra_obo profile. An undeclared key is silently stripped by pydantic when the 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 0dec44af402..18839a65d62 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -1086,43 +1086,73 @@ def test_clean_display_name_passthrough_when_no_suffix(): assert _clean_display_name("") == "" -def test_public_mcp_hub_returns_only_whitelisted_servers(): - """Regression: /public/mcp_hub must gate strictly on - litellm.public_mcp_servers, mirroring /public/model_hub and - /public/agent_hub. Servers with available_on_public_internet=True that - are not on the whitelist must not leak.""" +@pytest.mark.parametrize( + "strict,explicit,expected_listed", + ((True, True, True), (True, False, False), (False, True, True), (False, False, True)), +) +@pytest.mark.parametrize("stored_public", (None, False, True)) +def test_public_mcp_hub_derives_publication_metadata_without_mutating_registry( + strict: bool, + explicit: bool, + expected_listed: bool, + stored_public: bool | None, +) -> None: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport - app = FastAPI() + app: Final = FastAPI() app.include_router(router) - app.dependency_overrides[user_api_key_auth] = lambda: MagicMock() - client = TestClient(app) + client: Final = TestClient(app) - listed = MCPServer( + server: Final = MCPServer( server_id="listed", name="listed", server_name="listed", transport=MCPTransport.http, available_on_public_internet=True, + mcp_info=( + { + "is_public": stored_public, + "is_public_explicit": not explicit, + "description": "Preserve custom metadata", + } + if stored_public is not None + else None + ), ) - - mock_manager = MagicMock() - mock_manager.get_public_mcp_servers.return_value = [listed] + unlisted: Final = MCPServer( + server_id="unlisted", + name="unlisted", + transport=MCPTransport.http, + available_on_public_internet=False, + mcp_info={"is_public": True, "is_public_explicit": True}, + ) + manager: Final = MCPServerManager() + manager.config_mcp_servers = {server.server_id: server} + manager.registry = {unlisted.server_id: unlisted} + original_registry: Final = {key: value.model_dump() for key, value in manager.get_registry().items()} with ( - patch("litellm.public_mcp_servers", ["listed"]), + patch("litellm.public_mcp_servers", [server.server_id] if explicit else []), + patch("litellm.public_mcp_hub_strict_whitelist", strict), patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", - mock_manager, + manager, ), ): - response = client.get("/public/mcp_hub") + response: Final = client.get("/public/mcp_hub") assert response.status_code == 200 - data = response.json() - assert [item["server_id"] for item in data] == ["listed"] - app.dependency_overrides.clear() + data: Final = response.json() + assert [item["server_id"] for item in data] == ([server.server_id] if expected_listed else []) + if expected_listed: + assert data[0]["mcp_info"] == { + **({"description": "Preserve custom metadata"} if stored_public is not None else {}), + "is_public": True, + "is_public_explicit": explicit, + } + assert {key: value.model_dump() for key, value in manager.get_registry().items()} == original_registry def test_public_mcp_hub_returns_empty_when_whitelist_unset(): diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx index cb423b435ae..a48c991bc7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx @@ -217,12 +217,12 @@ const MCPPermissionManagement: React.FC = ({
Internal network only - +

- Turn on to restrict access to callers within your internal network only. + Turn on to restrict public IPs. Explicitly published server IDs remain accessible from public IPs.

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index 100b0ea93d3..d298d9d8145 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -126,3 +126,15 @@ describe("MCPServerCard per-user credentials", () => { expect(screen.queryByRole("button", { name: "Set" })).not.toBeInTheDocument(); }); }); + +describe("MCPServerCard network access", () => { + it("shows effective network access without a hub listing badge", () => { + renderCard({ + available_on_public_internet: false, + mcp_info: { server_name: "demo_server", is_public: true, is_public_explicit: true }, + }); + + expect(screen.getByText("All Networks")).toBeInTheDocument(); + expect(screen.queryByText(/^Hub:/)).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index 775809e3670..bb153f94665 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -13,7 +13,7 @@ import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/comp import { cn } from "@/lib/cva.config"; import { AUTH_TYPE, MCP_REACHABLE_DESCRIPTION, type MCPServer } from "@/components/mcp_tools/types"; import { Logo } from "@/components/molecules/logo/Logo"; -import { getMaskedAndFullUrl } from "./utils"; +import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils"; interface MCPServerCardProps { server: MCPServer; @@ -70,7 +70,7 @@ const MCPServerCard: FC = ({ server.auth_type === AUTH_TYPE.OAUTH2 && !server.oauth2_flow && !server.delegate_auth_to_upstream; const status = server.status || "unknown"; const healthTone = HEALTH_TONE[status] ?? HEALTH_TONE.unknown; - const isPublic = server.available_on_public_internet; + const networkAccess = getMCPNetworkAccess(server); const accessGroups = (server.mcp_access_groups ?? []).filter((g): g is string => typeof g === "string"); const missing = missingUserFields ?? []; @@ -236,10 +236,17 @@ const MCPServerCard: FC = ({ )} - - - {isPublic ? "Public" : "Internal"} - + + + + {networkAccess.label} + + } + /> + {networkAccess.description} + {accessGroups.slice(0, 2).map((g) => ( { }); it("shows the read-only settings summary before editing", async () => { - renderView({ allow_all_keys: true, available_on_public_internet: false }); + renderView({ + allow_all_keys: true, + available_on_public_internet: false, + mcp_info: { server_name: "demo server", is_public: true, is_public_explicit: true }, + }); await userEvent.click(screen.getByRole("tab", { name: "Settings" })); expect(await screen.findByText("MCP Server Settings")).toBeInTheDocument(); expect(screen.getByText("Allow All Keys")).toBeInTheDocument(); expect(screen.getByText("Enabled")).toBeInTheDocument(); - expect(screen.getByText("Internal only")).toBeInTheDocument(); + expect(screen.getByText("Network access")).toBeInTheDocument(); + expect(screen.getByText("All Networks")).toBeInTheDocument(); + expect(screen.queryByText("MCP Hub")).not.toBeInTheDocument(); + expect(screen.queryByText("Listed")).not.toBeInTheDocument(); expect(screen.queryByText("edit form")).not.toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx index a7ff34301a0..c97596ce0f6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx @@ -13,7 +13,7 @@ import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel"; import { getSecureItem } from "@/utils/secureStorage"; import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; import MCPServerCostDisplay from "./mcp_server_cost_display"; -import { getMaskedAndFullUrl } from "./utils"; +import { getMaskedAndFullUrl, getMCPNetworkAccess } from "./utils"; import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; import { CheckIcon, CopyIcon } from "lucide-react"; @@ -68,6 +68,7 @@ export const MCPServerView: React.FC = ({ const returningFromEditOAuth = isReturningFromEditOAuth(canEdit, mcpServer.server_id); const [editing, setEditing] = useState(isEditing || returningFromEditOAuth); const [showFullUrl, setShowFullUrl] = useState(false); + const networkAccess = getMCPNetworkAccess(mcpServer); const [copiedStates, setCopiedStates] = useState>({}); const [selectedTabIndex, setSelectedTabIndex] = useState(returningFromEditOAuth ? 2 : initialTabIndex); const canViewUserCredentials = userRole !== null && isProxyAdminTierRole(userRole); @@ -318,19 +319,13 @@ export const MCPServerView: React.FC = ({
-

Network Access

+

Network access

- {mcpServer.available_on_public_internet ? ( - - - Public - - ) : ( - - - Internal only - - )} + + + {networkAccess.label} + +

{networkAccess.description}

{handleAuth(mcpServer.auth_type) === "oauth2" && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx index 3b4fda400c2..bf30821d73a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.test.tsx @@ -3,11 +3,40 @@ import { extractMCPToken, maskUrl, getMaskedAndFullUrl, + getMCPNetworkAccess, validateMCPServerUrl, validateMCPServerName, normalizeToolOverrideMap, } from "./utils"; +describe("getMCPNetworkAccess", () => { + it.each([ + { publicIp: true, explicit: false, label: "All Networks" }, + { publicIp: false, explicit: true, label: "All Networks" }, + { publicIp: true, explicit: true, label: "All Networks" }, + { publicIp: false, explicit: false, label: "Internal Only" }, + { publicIp: true, explicit: undefined, label: "All Networks" }, + { publicIp: false, explicit: undefined, label: "Unknown" }, + { publicIp: undefined, explicit: false, label: "Unknown" }, + ])("reports $label for network=$publicIp and publication=$explicit", ({ publicIp, explicit, label }) => { + expect( + getMCPNetworkAccess({ + available_on_public_internet: publicIp, + mcp_info: { server_name: "demo", is_public: true, is_public_explicit: explicit }, + }).label, + ).toBe(label); + }); + + it("explains when hub publication permits public IPs", () => { + expect( + getMCPNetworkAccess({ + available_on_public_internet: false, + mcp_info: { server_name: "demo", is_public_explicit: true }, + }).description, + ).toContain("because this server is published in MCP Hub"); + }); +}); + describe("extractMCPToken", () => { it("should extract token after /mcp/", () => { const result = extractMCPToken("https://example.com/mcp/abc123"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx index 4738e1e8fba..bb72831d92e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/utils.tsx @@ -1,4 +1,37 @@ -import { MCPEnvVar, MCPEnvVarScope } from "@/components/mcp_tools/types"; +import { MCPEnvVar, MCPEnvVarScope, type MCPServer } from "@/components/mcp_tools/types"; + +export const getMCPNetworkAccess = ( + server: Pick, +): { + readonly label: "All Networks" | "Internal Only" | "Unknown"; + readonly dotClassName: string; + readonly description: string; +} => { + const explicitlyPublished = server.mcp_info?.is_public_explicit; + if (server.available_on_public_internet === true || explicitlyPublished === true) { + return { + label: "All Networks", + dotClassName: "bg-success", + description: + server.available_on_public_internet === true + ? "Allows requests from public and internal IPs. Authentication and access permissions still apply" + : "Allows requests from public and internal IPs because this server is published in MCP Hub. Authentication and access permissions still apply", + }; + } + if (server.available_on_public_internet === false && explicitlyPublished === false) { + return { + label: "Internal Only", + dotClassName: "bg-warning", + description: + "Allows requests only from internal/private IP ranges. Authentication and access permissions still apply", + }; + } + return { + label: "Unknown", + dotClassName: "bg-border", + description: "The proxy did not report enough network and publication settings to determine allowed client IPs", + }; +}; export const extractMCPToken = (url: string): { token: string | null; baseUrl: string } => { try { diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx index 1032b03a3ce..fee58e10fda 100644 --- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.test.tsx @@ -1,4 +1,4 @@ -import { render, screen } from "@testing-library/react"; +import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import { DataTable } from "@/components/shared/DataTable"; @@ -63,6 +63,23 @@ describe("getMCPHubTableColumns", () => { expect(screen.getByText("Auth Type")).toBeInTheDocument(); }); + it("shows hub membership separately from the network setting", () => { + renderTable(vi.fn(), [ + { ...mockServer, available_on_public_internet: false, mcp_info: { is_public: true } }, + { + ...mockServer, + server_id: "network-only", + server_name: "Network-only server", + available_on_public_internet: true, + mcp_info: { is_public: false }, + }, + ]); + + expect(screen.getByText("Hub listing")).toBeInTheDocument(); + expect(within(screen.getByRole("row", { name: /exa_test/ })).getByText("Listed")).toBeInTheDocument(); + expect(within(screen.getByRole("row", { name: /Network-only server/ })).getByText("Unlisted")).toBeInTheDocument(); + }); + it("does not expose a URL column", () => { renderTable(); expect(screen.queryByText("URL")).not.toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx index 20a14bcb476..db53e97569b 100644 --- a/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/MCPHubTableColumns.tsx @@ -203,8 +203,8 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps) { id: "is_public", accessorFn: (row) => row.mcp_info?.is_public === true, - meta: { title: "Public", skeleton: "badge", className: "hidden md:table-cell" }, - header: ({ column }) => , + meta: { title: "Hub listing", skeleton: "badge", className: "hidden md:table-cell" }, + header: ({ column }) => , size: 100, enableSorting: true, sortingFn: (rowA, rowB) => { @@ -214,7 +214,7 @@ export const getMCPHubTableColumns = ({ onServerClick }: MCPHubTableColumnsDeps) }, cell: ({ row }) => { const isPublic = row.original.mcp_info?.is_public === true; - return ; + return ; }, }, { diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx index 27fe2330acd..1f052b8932a 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.test.tsx @@ -1,5 +1,7 @@ import * as networking from "@/components/networking"; import userEvent from "@testing-library/user-event"; +import { act } from "@testing-library/react"; +import type { MCPServerData } from "./MCPHubTableColumns"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils"; import ModelHubTable from "./ModelHubTable"; @@ -18,6 +20,7 @@ vi.mock("@/components/networking", () => ({ getProxyBaseUrl: vi.fn(() => "http://localhost:4000"), getAgentsList: vi.fn(), fetchMCPServers: vi.fn(), + makeMCPPublicCall: vi.fn(), getUiSettings: vi.fn(), getClaudeCodePluginsList: vi.fn(() => Promise.resolve({ plugins: [] })), })); @@ -202,13 +205,13 @@ describe("ModelHubTable", () => { }); describe("hub tabs", () => { - const renderHub = async (agents: object[] = []) => { + const renderHub = async (agents: object[] = [], mcpServers: Promise = Promise.resolve([])) => { vi.mocked(networking.modelHubCall).mockResolvedValue({ data: [{ model_group: "claude-opus-4-8", providers: ["anthropic"], mode: "chat" }], }); vi.mocked(networking.getConfigFieldSetting).mockResolvedValue({ field_value: false }); vi.mocked(networking.getAgentsList).mockResolvedValue({ agents }); - vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + vi.mocked(networking.fetchMCPServers).mockReturnValue(mcpServers); vi.mocked(networking.getUiSettings).mockResolvedValue({ values: {} }); mockUseUISettings.mockReturnValue({ data: { values: {} }, isLoading: false }); @@ -219,6 +222,29 @@ describe("ModelHubTable", () => { return { user, search: await screen.findByPlaceholderText("Search model names...") }; }; + it("requires a fresh MCP publication list before and after saving", async () => { + const servers = Promise.withResolvers(); + const { user } = await renderHub([], servers.promise); + await user.click(screen.getByRole("tab", { name: "MCP Hub" })); + + const manageVisibility = screen.getByRole("button", { name: "Manage MCP Hub Visibility" }); + expect(manageVisibility).toBeDisabled(); + await act(async () => servers.resolve([])); + expect(manageVisibility).toBeEnabled(); + + const refresh = Promise.withResolvers(); + vi.mocked(networking.makeMCPPublicCall).mockResolvedValueOnce({}); + vi.mocked(networking.fetchMCPServers).mockReturnValueOnce(refresh.promise); + await user.click(manageVisibility); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(screen.getByRole("button", { name: "Save Publication List" })); + + expect(networking.makeMCPPublicCall).toHaveBeenCalledWith("test-token", []); + expect(manageVisibility).toBeDisabled(); + await act(async () => refresh.reject(new Error("Unable to reload the publication list"))); + expect(manageVisibility).toBeDisabled(); + }); + it("keeps the model filter typed on the Model Hub tab after visiting another hub", async () => { const { user, search } = await renderHub(); diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx index 063850b3e72..c4d776ea2fb 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx @@ -49,6 +49,10 @@ interface ModelHubTableProps { userRole: string | null; } +function isMCPHubVisibilityDisabled(isLoading: boolean, servers: readonly MCPServerData[] | null): boolean { + return isLoading || servers === null; +} + function HubEmptyState({ title, body }: { title: string; body: string }) { return (
@@ -359,10 +363,14 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, if (accessToken) { const fetchMcpData = async () => { try { + setMcpLoading(true); const response = await fetchMCPServers(accessToken); setMcpHubData(response); } catch (error) { + setMcpHubData(null); console.error("Error refreshing MCP server data:", error); + } finally { + setMcpLoading(false); } }; fetchMcpData(); @@ -567,7 +575,12 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, {/* Header with Make Public Button */} {publicPage == false && canModify && (
- +
)} diff --git a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx index 5b96e9ad194..881711c1668 100644 --- a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.test.tsx @@ -1,6 +1,8 @@ import { render, screen, fireEvent, act, waitFor } from "@testing-library/react"; import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; import MakeMCPPublicForm from "./MakeMCPPublicForm"; +import userEvent from "@testing-library/user-event"; +import { toast } from "@/lib/toast"; import { MCPServerData } from "@/components/AIHub/MCPHubTableColumns"; // Mock the networking function @@ -8,6 +10,10 @@ vi.mock("../../networking", () => ({ makeMCPPublicCall: vi.fn(), })); +vi.mock("@/lib/toast", () => ({ + toast: { success: vi.fn(), fromError: vi.fn() }, +})); + // Import the mocked function import { makeMCPPublicCall } from "../../networking"; const mockMakeMCPPublicCall = vi.mocked(makeMCPPublicCall); @@ -28,7 +34,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server1", transport: "http", status: "active", - mcp_info: { is_public: false }, + mcp_info: { is_public: false, is_public_explicit: false }, allowed_tools: ["tool-1", "tool-2"], auth_type: "bearer", credentials: {}, @@ -50,7 +56,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server2", transport: "websocket", status: "inactive", - mcp_info: { is_public: true }, + mcp_info: { is_public: true, is_public_explicit: true }, allowed_tools: [], auth_type: "none", credentials: {}, @@ -80,16 +86,16 @@ describe("MakeMCPPublicForm", () => { it("should render the component", () => { render(); - expect(screen.getByText("Make MCP Servers Public")).toBeInTheDocument(); - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Manage MCP Hub Visibility")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); }); it("should initialize with correct state", () => { render(); // Check that the component renders with the correct title and content - expect(screen.getByText("Make MCP Servers Public")).toBeInTheDocument(); - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Manage MCP Hub Visibility")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); // Check that all server checkboxes are present const checkboxes = screen.getAllByRole("checkbox"); @@ -104,7 +110,7 @@ describe("MakeMCPPublicForm", () => { render(); // Initially on step 1 - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); // Select all servers using the select all checkbox const selectAllCheckbox = screen.getByRole("checkbox", { name: "Select All (2)" }); @@ -123,7 +129,7 @@ describe("MakeMCPPublicForm", () => { // Should move to step 2 await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); }); @@ -145,10 +151,10 @@ describe("MakeMCPPublicForm", () => { // Wait for navigation to complete await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); - const submitButton = screen.getByRole("button", { name: "Make Public" }); + const submitButton = screen.getByRole("button", { name: "Save Publication List" }); await act(async () => { fireEvent.click(submitButton); }); @@ -187,29 +193,105 @@ describe("MakeMCPPublicForm", () => { expect(checkboxes[2]).not.toBeChecked(); }); - it("should show error when no servers selected", async () => { + it("submits an empty publication list after the last server is deselected", async () => { + mockMakeMCPPublicCall.mockResolvedValueOnce({}); render(); - // Deselect all servers first - const checkboxes = screen.getAllByRole("checkbox"); - await act(async () => { - fireEvent.click(checkboxes[0]); // Click select all to select all - }); - await act(async () => { - fireEvent.click(checkboxes[0]); // Click select all again to deselect all - }); + fireEvent.click(screen.getAllByRole("checkbox")[2]); + expect(screen.getByRole("button", { name: "Next" })).toBeEnabled(); + fireEvent.click(screen.getByRole("button", { name: "Next" })); + fireEvent.click(screen.getByRole("button", { name: "Save Publication List" })); - // Try to go to next step - const nextButton = screen.getByRole("button", { name: "Next" }); - await act(async () => { - fireEvent.click(nextButton); - }); - - // Should stay on same step - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + await waitFor(() => expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", [])); + expect(mockProps.onSuccess).toHaveBeenCalled(); }); - it("should display empty state when no servers are available", () => { + it("keeps legacy listings separate from explicitly published selections", () => { + render( + , + ); + + expect(screen.getAllByRole("checkbox")[1]).not.toBeChecked(); + expect(screen.getAllByRole("checkbox")[2]).toBeChecked(); + expect(screen.getByText("Listed by legacy mode")).toBeInTheDocument(); + }); + + it.each([ + { mode: "all missing, stale true", info: { is_public: true }, mixed: false }, + { mode: "all missing, stale false", info: { is_public: false }, mixed: false }, + { mode: "mixed, stale true", info: { is_public: true }, mixed: true }, + { mode: "mixed, stale false", info: { is_public: false }, mixed: true }, + { mode: "null explicit status", info: { is_public: true, is_public_explicit: null }, mixed: true }, + { mode: "nonboolean explicit status", info: { is_public: true, is_public_explicit: "true" }, mixed: true }, + ])("blocks unknown explicit publication metadata: $mode", ({ info, mixed }) => { + const unknownServer = { ...mockProps.mcpHubData[0], mcp_info: info }; + const catalog = mixed ? [unknownServer, mockProps.mcpHubData[1]] : [unknownServer]; + render(); + + expect(screen.getByRole("alert")).toHaveTextContent("explicit publication status"); + expect(screen.queryByRole("checkbox")).not.toBeInTheDocument(); + expect(screen.queryByText("Configure in YAML")).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Copy code" })).not.toBeInTheDocument(); + const nextButton = screen.getByRole("button", { name: "Next" }); + expect(nextButton).toBeDisabled(); + fireEvent.click(nextButton); + expect(screen.queryByText("Confirm MCP Hub Publication")).not.toBeInTheDocument(); + expect(mockMakeMCPPublicCall).not.toHaveBeenCalled(); + }); + + it.each([true, false])("blocks confirmation when explicit metadata disappears with stale listing %s", (listed) => { + const { rerender } = render(); + fireEvent.click(screen.getByRole("button", { name: "Next" })); + expect(screen.getByRole("button", { name: "Save Publication List" })).toBeEnabled(); + + const catalog = [{ ...mockProps.mcpHubData[0], mcp_info: { is_public: listed } }, mockProps.mcpHubData[1]]; + rerender(); + + expect(screen.getByRole("alert")).toHaveTextContent("explicit publication status"); + expect(screen.queryByText("Confirm MCP Hub Publication")).not.toBeInTheDocument(); + const saveButton = screen.getByRole("button", { name: "Save Publication List" }); + expect(saveButton).toBeDisabled(); + fireEvent.click(saveButton); + expect(mockMakeMCPPublicCall).not.toHaveBeenCalled(); + expect(screen.queryByRole("button", { name: "Copy code" })).not.toBeInTheDocument(); + + const refreshedCatalog = [ + { ...mockProps.mcpHubData[0], mcp_info: { is_public: true, is_public_explicit: true } }, + { ...mockProps.mcpHubData[1], mcp_info: { is_public: false, is_public_explicit: false } }, + ]; + rerender(); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Next" })).toBeEnabled(); + expect(screen.getByRole("checkbox", { name: "Publish Test Server 1" })).toBeChecked(); + expect(screen.getByRole("checkbox", { name: "Publish Test Server 2" })).not.toBeChecked(); + }); + + it("copies publication YAML using the selected server IDs", async () => { + const user = userEvent.setup(); + render(); + + await user.click(screen.getByText("Configure in YAML")); + await user.click(screen.getByRole("button", { name: "Copy code" })); + + expect(await navigator.clipboard.readText()).toBe( + 'litellm_settings:\n public_mcp_hub_strict_whitelist: true\n public_mcp_servers:\n - "server-2"', + ); + + await user.click(screen.getByRole("checkbox", { name: "Publish Test Server 2" })); + await user.click(screen.getByRole("button", { name: "Copy code" })); + expect(await navigator.clipboard.readText()).toBe( + "litellm_settings:\n public_mcp_hub_strict_whitelist: true\n public_mcp_servers: []", + ); + }); + + it("allows clearing publication IDs when the loaded server catalog is empty", async () => { + mockMakeMCPPublicCall.mockResolvedValueOnce({}); const emptyProps = { ...mockProps, mcpHubData: [] as MCPServerData[], @@ -223,9 +305,13 @@ describe("MakeMCPPublicForm", () => { const selectAllCheckbox = screen.getByRole("checkbox", { name: "Select All" }); expectDisabledControl(selectAllCheckbox); - // Next button should be disabled const nextButton = screen.getByRole("button", { name: "Next" }); - expect(nextButton).toBeDisabled(); + expect(nextButton).toBeEnabled(); + fireEvent.click(nextButton); + fireEvent.click(screen.getByRole("button", { name: "Save Publication List" })); + + await waitFor(() => expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", [])); + expect(mockProps.onSuccess).toHaveBeenCalled(); }); it("should handle Cancel button functionality", async () => { @@ -252,7 +338,7 @@ describe("MakeMCPPublicForm", () => { // Verify we're on step 1 await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); // Click Previous button @@ -262,7 +348,7 @@ describe("MakeMCPPublicForm", () => { }); // Should go back to step 0 - expect(screen.getByText("Select MCP Servers to Make Public")).toBeInTheDocument(); + expect(screen.getByText("Select MCP Servers for the Hub")).toBeInTheDocument(); }); it("should handle individual server selection", async () => { @@ -322,8 +408,8 @@ describe("MakeMCPPublicForm", () => { }); it("should handle submit error properly", async () => { - const errorMessage = "Network error"; - mockMakeMCPPublicCall.mockRejectedValueOnce(new Error(errorMessage)); + const error = new Error("Update litellm_settings.public_mcp_servers in your YAML configuration"); + mockMakeMCPPublicCall.mockRejectedValueOnce(error); render(); @@ -333,10 +419,10 @@ describe("MakeMCPPublicForm", () => { }); await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); - const submitButton = screen.getByRole("button", { name: "Make Public" }); + const submitButton = screen.getByRole("button", { name: "Save Publication List" }); await act(async () => { fireEvent.click(submitButton); }); @@ -346,6 +432,8 @@ describe("MakeMCPPublicForm", () => { expect(mockMakeMCPPublicCall).toHaveBeenCalledWith("test-token", ["server-2"]); }); + expect(toast.fromError).toHaveBeenCalledWith(error); + // Should not call onSuccess or onClose on error expect(mockProps.onSuccess).not.toHaveBeenCalled(); expect(mockProps.onClose).not.toHaveBeenCalled(); @@ -366,10 +454,10 @@ describe("MakeMCPPublicForm", () => { }); await waitFor(() => { - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); }); - const submitButton = screen.getByRole("button", { name: "Make Public" }); + const submitButton = screen.getByRole("button", { name: "Save Publication List" }); await act(async () => { fireEvent.click(submitButton); }); @@ -381,7 +469,7 @@ describe("MakeMCPPublicForm", () => { expect(mockMakeMCPPublicCall).toHaveBeenCalledTimes(1); expect(mockProps.onSuccess).not.toHaveBeenCalled(); expect(mockProps.onClose).not.toHaveBeenCalled(); - expect(screen.getByText("Confirm Making MCP Servers Public")).toBeInTheDocument(); + expect(screen.getByText("Confirm MCP Hub Publication")).toBeInTheDocument(); resolvePromise({}); await waitFor(() => { @@ -400,7 +488,7 @@ describe("MakeMCPPublicForm", () => { // Modal should not be rendered expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); - expect(screen.queryByText("Make MCP Servers Public")).not.toBeInTheDocument(); + expect(screen.queryByText("Manage MCP Hub Visibility")).not.toBeInTheDocument(); }); it("should preselect already public servers when modal opens", () => { @@ -415,7 +503,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server1", transport: "http", status: "active", - mcp_info: { is_public: false }, // Not public + mcp_info: { is_public: false, is_public_explicit: false }, // Not public allowed_tools: [], auth_type: "bearer", credentials: {}, @@ -437,7 +525,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server2", transport: "websocket", status: "inactive", - mcp_info: { is_public: true }, // Already public + mcp_info: { is_public: true, is_public_explicit: true }, // Already public allowed_tools: [], auth_type: "none", credentials: {}, @@ -459,7 +547,7 @@ describe("MakeMCPPublicForm", () => { url: "http://example.com/server3", transport: "sse", status: "healthy", - mcp_info: { is_public: true }, // Already public + mcp_info: { is_public: true, is_public_explicit: true }, // Already public allowed_tools: [], auth_type: "oauth", credentials: {}, diff --git a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx index 8287cf47f1a..2448732a236 100644 --- a/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/forms/MakeMCPPublicForm.tsx @@ -1,5 +1,6 @@ import React, { useState, useEffect } from "react"; import { Loader2 } from "lucide-react"; +import CodeBlock from "@/components/CodeBlock"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Checkbox } from "@/components/ui/checkbox"; @@ -29,6 +30,11 @@ interface MakeMCPPublicFormProps { onSuccess: () => void; } +interface PublicationSelection { + readonly catalog: MCPServerData[]; + readonly serverIds: Set; +} + const MakeMCPPublicForm: React.FC = ({ visible, onClose, @@ -37,21 +43,28 @@ const MakeMCPPublicForm: React.FC = ({ onSuccess, }) => { const [currentStep, setCurrentStep] = useState(0); - const [selectedServers, setSelectedServers] = useState>(new Set()); + const [selection, setSelection] = useState(null); const [loading, setLoading] = useState(false); + const selectedServers = selection?.serverIds ?? new Set(); + const hasPublicationMetadata = mcpHubData.every((server) => typeof server.mcp_info?.is_public_explicit === "boolean"); + const canManagePublication = hasPublicationMetadata && selection?.catalog === mcpHubData; + const publicationYaml = [ + "litellm_settings:", + " public_mcp_hub_strict_whitelist: true", + selectedServers.size === 0 + ? " public_mcp_servers: []" + : ` public_mcp_servers:\n${Array.from(selectedServers, (id) => ` - ${JSON.stringify(id)}`).join("\n")}`, + ].join("\n"); const handleClose = () => { setCurrentStep(0); - setSelectedServers(new Set()); + setSelection(null); onClose(); }; const handleNext = () => { + if (!canManagePublication) return; if (currentStep === 0) { - if (selectedServers.size === 0) { - toast.fromError("Please select at least one MCP server to make public"); - return; - } setCurrentStep(1); } }; @@ -69,37 +82,32 @@ const MakeMCPPublicForm: React.FC = ({ } else { newSelection.delete(serverId); } - setSelectedServers(newSelection); + setSelection({ catalog: mcpHubData, serverIds: newSelection }); }; const handleSelectAll = (checked: boolean) => { if (checked) { const allServerIds = mcpHubData.map((server) => server.server_id); - setSelectedServers(new Set(allServerIds)); + setSelection({ catalog: mcpHubData, serverIds: new Set(allServerIds) }); } else { - setSelectedServers(new Set()); + setSelection({ catalog: mcpHubData, serverIds: new Set() }); } }; - // Initialize and preselect already public servers when modal opens useEffect(() => { - if (visible && mcpHubData.length > 0) { - // Extract server IDs from servers that are already public - const publicServerIds = mcpHubData - .filter((server) => server.mcp_info?.is_public === true) - .map((server) => server.server_id); - - // Preselect servers that are already public - setSelectedServers(new Set(publicServerIds)); - } - }, [visible]); // Only re-run when modal visibility changes, not when mcpHubData updates - - const handleSubmit = async () => { - if (selectedServers.size === 0) { - toast.fromError("Please select at least one MCP server to make public"); + if (!visible || !hasPublicationMetadata) { + setSelection(null); return; } + const publicServerIds = mcpHubData + .filter((server) => server.mcp_info.is_public_explicit === true) + .map((server) => server.server_id); + setSelection({ catalog: mcpHubData, serverIds: new Set(publicServerIds) }); + setCurrentStep(0); + }, [visible, mcpHubData, hasPublicationMetadata]); + const handleSubmit = async () => { + if (!canManagePublication) return; setLoading(true); try { const serverIdsToMakePublic = Array.from(selectedServers); @@ -107,12 +115,12 @@ const MakeMCPPublicForm: React.FC = ({ // Make batch API call for all servers await makeMCPPublicCall(accessToken, serverIdsToMakePublic); - toast.success(`Successfully made ${serverIdsToMakePublic.length} MCP server(s) public!`); + toast.success("MCP Hub publication list updated"); handleClose(); onSuccess(); } catch (error) { console.error("Error making MCP servers public:", error); - toast.fromError("Failed to make MCP servers public. Please try again."); + toast.fromError(error); } finally { setLoading(false); } @@ -126,7 +134,7 @@ const MakeMCPPublicForm: React.FC = ({ return (
-

Select MCP Servers to Make Public

+

Select MCP Servers for the Hub

- Select the MCP servers you want to be visible on the public model hub. Users will still require a valid - Virtual Key to use these servers. + Select the complete list of MCP servers to publish on the public hub. Uncheck a server to remove it from this + list, or uncheck all to clear it. Authentication and access permissions still apply +

+ +

+ Legacy mode also lists servers with public IP access enabled. Set public_mcp_hub_strict_whitelist to true in + your configuration to use only the publication list

@@ -160,16 +173,22 @@ const MakeMCPPublicForm: React.FC = ({ className="flex items-center space-x-3 p-3 border rounded-lg hover:bg-accent" > handleServerSelection(server.server_id, checked === true)} />

{server.server_name}

- {isPublic && Public} + {isPublic && ( + + {server.mcp_info?.is_public_explicit === false ? "Listed by legacy mode" : "Listed"} + + )} {server.transport} {server.status || "unknown"}
+

{server.server_id}

{server.description || server.url}

@@ -193,6 +212,18 @@ const MakeMCPPublicForm: React.FC = ({
+
+ Configure in YAML +
+

+ Merge these settings into your proxy configuration and reload it. Entries use the server IDs shown above, + not names or aliases. For servers defined in YAML, pin server_id in each existing mcp_servers entry so the + publication list stays stable +

+ +
+
+ {selectedServers.size > 0 && (

@@ -207,19 +238,20 @@ const MakeMCPPublicForm: React.FC = ({ const renderStep2Content = () => { return (

-

Confirm Making MCP Servers Public

+

Confirm MCP Hub Publication

- Warning: Once you make these MCP servers public, anyone who can go to the{" "} - /ui/model_hub_table will be able to know they exist on the proxy. + Anyone who can open /ui/model_hub_table can discover published servers. Explicitly published + server IDs also allow requests from public IPs. Authentication and access permissions still apply

-

MCP Servers to be made public:

+

MCP servers in the publication list:

+ {selectedServers.size === 0 &&

No explicitly published servers

} {Array.from(selectedServers).map((serverId) => { const server = mcpHubData.find((s) => s.server_id === serverId); return ( @@ -248,8 +280,8 @@ const MakeMCPPublicForm: React.FC = ({

- Total: {selectedServers.size} MCP server{selectedServers.size !== 1 ? "s" : ""} will be - made public + Saving replaces the publication list with {selectedServers.size} MCP server + {selectedServers.size !== 1 ? "s" : ""}. Legacy mode may still list servers with public IP access enabled

@@ -257,6 +289,15 @@ const MakeMCPPublicForm: React.FC = ({ }; const renderStepContent = () => { + if (!hasPublicationMetadata) { + return ( +
+ This proxy does not provide explicit publication status for every MCP server. Update the proxy to manage + visibility here, or edit litellm_settings.public_mcp_servers in its existing configuration +
+ ); + } + if (!canManagePublication) return

Loading publication settings

; switch (currentStep) { case 0: return renderStep1Content(); @@ -276,15 +317,15 @@ const MakeMCPPublicForm: React.FC = ({
{currentStep === 0 && ( - )} {currentStep === 1 && ( - )}
@@ -296,7 +337,7 @@ const MakeMCPPublicForm: React.FC = ({ !open && handleClose()} disablePointerDismissal> - Make MCP Servers Public + Manage MCP Hub Visibility
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index be7d39616ca..afeff869b7a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -321,6 +321,8 @@ export interface MCPServerCostInfo { // Define MCP provider info export interface MCPInfo { server_name: string; + is_public?: boolean; + is_public_explicit?: boolean; description?: string; logo_url?: string; mcp_server_cost_info?: MCPServerCostInfo | null; From 4c2458a0b1738472ad66dec20207a9673245d8b4 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:18:23 -0700 Subject: [PATCH 09/51] chore(cost-map): sync openrouter prices for deepseek, minimax, qwen and glm rows (#43384) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 45 ++++++++++--------- model_prices_and_context_window.json | 45 ++++++++++--------- 2 files changed, 46 insertions(+), 44 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 45f5967d372..d807dca329a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -41956,15 +41956,15 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "cache_read_input_token_cost": 8.8e-09, - "input_cost_per_token": 2.64e-07, + "cache_read_input_token_cost": 2.475e-07, + "input_cost_per_token": 2.476e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 7.92e-07, + "output_cost_per_token": 3.5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42318,13 +42318,13 @@ "max_output_tokens": 8000 }, "openrouter/minimax/minimax-m2": { - "input_cost_per_token": 3e-07, + "input_cost_per_token": 2.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.02e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -66867,13 +66867,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 3.794e-07, - "output_cost_per_token": 1.1924e-06, - "cache_read_input_token_cost": 7.046e-08, + "input_cost_per_token": 2.38e-07, + "output_cost_per_token": 7.48e-07, + "cache_read_input_token_cost": 3.91e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67005,7 +67005,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 1.6e-08, - "input_cost_per_token": 2.2e-08, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, @@ -67719,13 +67719,13 @@ "supports_web_search": false }, "openrouter/minimax/minimax-m2.7": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 2.1e-07, + "output_cost_per_token": 8.4e-07, + "cache_read_input_token_cost": 4.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 204800, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 176947, + "max_tokens": 176947, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -68841,12 +68841,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b": { - "input_cost_per_token": 1.2e-07, - "output_cost_per_token": 5e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.3e-07, + "output_cost_per_token": 5.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, - "max_output_tokens": 16384, - "max_tokens": 16384, + "max_output_tokens": 8192, + "max_tokens": 8192, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 45f5967d372..d807dca329a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -41956,15 +41956,15 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "cache_read_input_token_cost": 8.8e-09, - "input_cost_per_token": 2.64e-07, + "cache_read_input_token_cost": 2.475e-07, + "input_cost_per_token": 2.476e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 7.92e-07, + "output_cost_per_token": 3.5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42318,13 +42318,13 @@ "max_output_tokens": 8000 }, "openrouter/minimax/minimax-m2": { - "input_cost_per_token": 3e-07, + "input_cost_per_token": 2.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.02e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -66867,13 +66867,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 3.794e-07, - "output_cost_per_token": 1.1924e-06, - "cache_read_input_token_cost": 7.046e-08, + "input_cost_per_token": 2.38e-07, + "output_cost_per_token": 7.48e-07, + "cache_read_input_token_cost": 3.91e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67005,7 +67005,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "cache_read_input_token_cost": 1.6e-08, - "input_cost_per_token": 2.2e-08, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, @@ -67719,13 +67719,13 @@ "supports_web_search": false }, "openrouter/minimax/minimax-m2.7": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 2.1e-07, + "output_cost_per_token": 8.4e-07, + "cache_read_input_token_cost": 4.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 204800, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 176947, + "max_tokens": 176947, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -68841,12 +68841,13 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b": { - "input_cost_per_token": 1.2e-07, - "output_cost_per_token": 5e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.3e-07, + "output_cost_per_token": 5.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, - "max_output_tokens": 16384, - "max_tokens": 16384, + "max_output_tokens": 8192, + "max_tokens": 8192, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, From e73abe6c72785ad91d4927da26de3a5d1b54300b Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 17:36:34 -0700 Subject: [PATCH 10/51] chore(cost-map): drop stale cache hit field from openrouter deepseek-v4-pro-0813 (#43389) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 1 - model_prices_and_context_window.json | 1 - 2 files changed, 2 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index d807dca329a..09fc442e5a7 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -41958,7 +41958,6 @@ "openrouter/deepseek/deepseek-v4-pro-0813": { "cache_read_input_token_cost": 2.475e-07, "input_cost_per_token": 2.476e-07, - "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index d807dca329a..09fc442e5a7 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -41958,7 +41958,6 @@ "openrouter/deepseek/deepseek-v4-pro-0813": { "cache_read_input_token_cost": 2.475e-07, "input_cost_per_token": 2.476e-07, - "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, From eea1d0f2696d6ab6b67e8b208c85bae9fa624e1e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:05:21 -0700 Subject: [PATCH 11/51] fix(responses): stream guardrail pre-call block as SSE with a typed output item (#42507) * fix(responses): stream guardrail pre-call block as SSE with a typed output item Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): import blocked usage helper from the guardrail utils module Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): drop narrating docstrings and poll without rebinding Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover pre-call guardrail block on /v1/responses stream and json Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): audit cells for responses guardrail block contract Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): observe upstream on the recorded chat route for responses denial cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): tidy responses denial audit cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): wait for worker count to recover after SIGKILL Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): require a replacement worker after SIGKILL Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): type the blocked response test helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yucheng --- .../guardrail_translation/handler.py | 6 +- .../proxy/response_api_endpoints/endpoints.py | 32 +- tests/e2e/coverage_registry/guardrail.yaml | 1 + tests/e2e/guardrails/guardrails_client.py | 29 + ...est_responses_pre_call_block_stream_e2e.py | 154 ++++ .../observability/test_guardrail_effects.py | 695 +++++++++++++++++- .../response_api_endpoints/test_endpoints.py | 148 +++- .../proxy/test_blocked_response_usage.py | 20 +- 8 files changed, 1022 insertions(+), 63 deletions(-) create mode 100644 tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 66cebe0175d..d6d68e0607a 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -1540,7 +1540,7 @@ class OpenAIResponsesHandler(BaseTranslation): from litellm.responses.streaming_iterator import build_synthetic_response_events return build_synthetic_response_events( - transformed=_blocked_response(exc, response_id=f"resp_{uuid.uuid4()}", model=exc.model), + transformed=build_blocked_response(exc), logging_obj=None, chunk_size=max(len(exc.message), 1), ) @@ -1648,6 +1648,10 @@ def _blocked_output_item(exc: "ModifyResponseException") -> GenericResponseOutpu return GenericResponseOutputItem.model_validate(payload) +def build_blocked_response(exc: "ModifyResponseException") -> ResponsesAPIResponse: + return _blocked_response(exc, response_id=f"resp_{uuid.uuid4()}", model=exc.model) + + def _blocked_response( exc: "ModifyResponseException", response_id: str, diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 75eefb2e73b..c5d702ad65a 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1,13 +1,11 @@ import asyncio import contextlib import json -import time -from collections.abc import AsyncIterator, Awaitable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Mapping, Sequence from enum import Enum from functools import partial from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, TypeAlias, cast, get_args -from uuid import uuid4 import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -21,8 +19,9 @@ from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger from litellm.constants import EMPTY_MAPPING from litellm.integrations.custom_guardrail import ModifyResponseException -from litellm.llms.base_llm.guardrail_translation.utils import ( - blocked_responses_api_usage as _blocked_responses_api_usage, +from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + build_blocked_response, ) from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import ( @@ -30,7 +29,7 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth, user_api_key_auth_websocket, ) -from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing, create_response from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_set_request_parsed_body, @@ -440,17 +439,16 @@ async def responses_api( request_data=_data, ) - violation_text: Final = e.message - response_obj: Final = ResponsesAPIResponse( - id=f"resp_{uuid4()}", - object="response", - created_at=int(time.time()), - model=e.model or data.get("model"), - output=cast(Any, [{"content": [{"type": "text", "text": violation_text}]}]), - status="completed", - usage=_blocked_responses_api_usage(e.original_response), - ) - return response_obj + if data.get("stream") is True: + block_chunks: Final = OpenAIResponsesHandler().build_block_sse_chunks(e) + + async def _blocked_stream() -> AsyncGenerator[str, None]: + for chunk in block_chunks: + yield chunk.decode() + yield "data: [DONE]\n\n" + + return await create_response(generator=_blocked_stream(), media_type="text/event-stream", headers={}) + return build_blocked_response(e) except Exception as e: raise await processor._handle_llm_api_exception( e=e, diff --git a/tests/e2e/coverage_registry/guardrail.yaml b/tests/e2e/coverage_registry/guardrail.yaml index f49568c883b..920a288aea6 100644 --- a/tests/e2e/coverage_registry/guardrail.yaml +++ b/tests/e2e/coverage_registry/guardrail.yaml @@ -37,4 +37,5 @@ - {id: guardrail.mcp_security.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/mcp_security", rationale: "MCP protocol security"} - {id: guardrail.llm_as_a_judge.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/llm_as_a_judge", rationale: "LLM-based judgment guardrail"} - {id: guardrail.litellm_content_filter.pre_mcp_call.blocks, module: guardrail, tier: P1, hook_point: pre_mcp_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/litellm_content_filter/content_filter.py:_scan_mcp_tool_call_arguments", rationale: "A general content-filter guardrail configured mode=pre_mcp_call blocks a banned keyword in an MCP tool call's arguments before it reaches the upstream MCP server; a clean argument passes"} +- {id: guardrail.custom_code.pre_call.blocks, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [responses], source: "response_api_endpoints/endpoints.py ModifyResponseException handler", rationale: "A pre_call custom_code block on /v1/responses must answer in the requested shape: SSE response.completed with a completed assistant output_text message item when stream=true, schema-valid JSON when not, both with zero usage"} - {id: guardrail.dispatch.pre_call.rejects_unknown_name, module: guardrail, tier: P1, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "proxy guardrail dispatch (per-request `guardrails` selector)", rationale: "A request naming a guardrail this proxy does not serve must fail closed with a 4xx; today it is silently served unguarded, so a typo'd name drops the protection the caller asked for"} diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index 17223dc36fa..1f4fc43355b 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -111,6 +111,15 @@ class ToolPermissionParamsBody(GuardrailParamsBase): on_disallowed_action: Literal["block", "rewrite"] = "block" +class CustomCodeParamsBody(GuardrailParamsBase): + """Custom-code guardrail params: `custom_code` is the sandboxed source the + proxy compiles, which must define `apply_guardrail(inputs, request_data, + input_type)` returning `allow()` or `block(reason)`.""" + + guardrail: Literal["custom_code"] = "custom_code" + custom_code: str + + GuardrailParamsBody = ( ContentFilterParamsBody | BedrockGuardrailParamsBody @@ -118,6 +127,7 @@ GuardrailParamsBody = ( | BlockCodeExecutionParamsBody | PresidioParamsBody | ToolPermissionParamsBody + | CustomCodeParamsBody ) @@ -174,6 +184,7 @@ class _ResponsesGuardrailBody(BaseModel): model: str input: str guardrails: list[str] | None = None + stream: bool | None = None @dataclass(frozen=True, slots=True) @@ -509,6 +520,24 @@ class GuardrailsClient: json=_ResponsesGuardrailBody(model=model, input=text, guardrails=guardrails), ) + def responses_stream_raw( + self, + key: str, + model: str, + text: str, + *, + guardrails: list[str] | None = None, + ) -> StreamingResponse: + """Drive /v1/responses with stream=true, returning the raw HTTP outcome: + a streamed block is judged on status, content-type, and the SSE event + sequence, not a typed JSON body.""" + return self.proxy.transport.send( + "/v1/responses", + headers=self.proxy.transport.bearer(key), + json=_ResponsesGuardrailBody(model=model, input=text, guardrails=guardrails, stream=True), + stream=True, + ) + def apply_guardrail(self, key: str, *, name: str, text: str) -> Result[ApplyGuardrailResponse]: return self.proxy.transport.post( "/guardrails/apply_guardrail", diff --git a/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py b/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py new file mode 100644 index 00000000000..93512e2a64c --- /dev/null +++ b/tests/e2e/guardrails/test_responses_pre_call_block_stream_e2e.py @@ -0,0 +1,154 @@ +from __future__ import annotations + +import time +from collections.abc import Callable +from typing import Final + +import pytest +from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker +from e2e_http import StreamingResponse +from guardrails_client import CustomCodeParamsBody, GuardrailsClient +from lifecycle import ResourceManager +from pydantic import BaseModel, TypeAdapter + +pytestmark = pytest.mark.e2e + +DENIAL: Final = "This model is not currently available. Please contact support if you think this is a mistake." + +CUSTOM_CODE: Final = f''' +def apply_guardrail(inputs, request_data, input_type): + return block("{DENIAL}") +''' + + +class _ContentPart(BaseModel): + type: str + text: str | None = None + + +class _OutputItem(BaseModel): + type: str | None = None + id: str | None = None + role: str | None = None + status: str | None = None + content: list[_ContentPart] = [] + + +class _Usage(BaseModel): + total_tokens: int = 0 + + +class _ResponseBody(BaseModel): + output: list[_OutputItem] = [] + usage: _Usage | None = None + + +class _EventHead(BaseModel): + type: str + + +class _CompletedEvent(BaseModel): + type: str + response: _ResponseBody + + +_EVENT_HEAD: Final = TypeAdapter(_EventHead) + + +def _denial_delivered(result: StreamingResponse) -> bool: + if not result.ok: + return False + if DENIAL in result.body: + return True + return any(DENIAL in event for event in result.stream_events) + + +def _poll_terminal(result: StreamingResponse) -> bool: + if _denial_delivered(result): + return True + if result.ok: + return False + return "Guardrail not found" not in result.body and result.status_code not in (-1, 401, 429) + + +def _poll_attempt(call: Callable[[], StreamingResponse], deadline: float) -> StreamingResponse: + result: Final = call() + if _poll_terminal(result) or time.monotonic() >= deadline: + return result + time.sleep(POLL_INTERVAL) + return _poll_attempt(call, deadline) + + +def _poll_for_block(call: Callable[[], StreamingResponse]) -> StreamingResponse: + return _poll_attempt(call, time.monotonic() + POLL_TIMEOUT) + + +def _assert_blocked_response(response: _ResponseBody) -> None: + item = next(iter(response.output), None) + assert item is not None, f"blocked response carried no output item: {response.output!r}" + assert item.type == "message", f"output[0] must be a message item, got {item.type!r}: {item!r}" + assert item.role == "assistant", f"output[0] role must be assistant, got {item.role!r}" + assert item.status == "completed", f"output[0] status must be completed, got {item.status!r}" + part = next(iter(item.content), None) + assert part is not None, f"output[0] carried no content part: {item!r}" + assert part.type == "output_text", f"content[0] must be output_text, got {part.type!r}" + assert part.text == DENIAL, f"content[0] text must be the denial, got {part.text!r}" + assert response.usage is not None and response.usage.total_tokens == 0, ( + f"a blocked response never reached a provider, usage must be zero: {response.usage!r}" + ) + + +class TestResponsesPreCallBlock: + def _register_block(self, client: GuardrailsClient, resources: ResourceManager) -> str: + name: Final = f"e2e-custom-code-responses-block-{unique_marker()}" + guardrail_id: Final = client.register( + name, + CustomCodeParamsBody(mode="pre_call", default_on=False, custom_code=CUSTOM_CODE), + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + return name + + @pytest.mark.covers("guardrail.custom_code.pre_call.blocks", exercised_on=["responses"]) + def test_stream_block_is_sse_with_completed_assistant_message( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + name: Final = self._register_block(client, resources) + model: Final = client.create_backend_model( + resources, prefix="e2e-responses-block", backend="openai/gpt-4.1-mini", api_key="os.environ/OPENAI_API_KEY" + ) + + result: Final = _poll_for_block( + lambda: client.responses_stream_raw(scoped_key, model, "say hi", guardrails=[name]) + ) + + assert result.status_code == 200, f"a pre_call block answers 200, got {result.status_code}: {result.body[:400]}" + assert (result.content_type or "").startswith("text/event-stream"), ( + f"stream=true must answer SSE, got content-type {result.content_type!r}: {result.body[:400]}" + ) + events: Final = tuple(_EVENT_HEAD.validate_json(payload).type for payload in result.stream_events) + completed: Final = tuple( + _CompletedEvent.model_validate_json(payload) + for payload, event_type in zip(result.stream_events, events) + if event_type == "response.completed" + ) + assert len(completed) == 1, ( + f"the denial stream must end in exactly one response.completed event, got events {events!r}" + ) + _assert_blocked_response(completed[0].response) + + @pytest.mark.covers("guardrail.custom_code.pre_call.blocks", exercised_on=["responses"]) + def test_non_stream_block_is_schema_valid_json( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + name: Final = self._register_block(client, resources) + model: Final = client.create_backend_model( + resources, prefix="e2e-responses-block", backend="openai/gpt-4.1-mini", api_key="os.environ/OPENAI_API_KEY" + ) + + result: Final = _poll_for_block(lambda: client.responses(scoped_key, model, "say hi", guardrails=[name])) + + assert result.status_code == 200, f"a pre_call block answers 200, got {result.status_code}: {result.body[:400]}" + assert (result.content_type or "").startswith("application/json"), ( + f"a non-streaming block answers JSON, got content-type {result.content_type!r}" + ) + _assert_blocked_response(_ResponseBody.model_validate_json(result.body)) diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index 4fac42a796d..9f5f3da4302 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -1,15 +1,22 @@ import json +import os +import signal +import socket import uuid +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import Final +import httpx +import psutil import pytest import yaml from integration._support.client import Gateway, eventually, object_value from integration._support.database import read_rows from integration._support.mcp import mcp_peer, register_mcp, tool_names -from integration._support.process import owned_proxy +from integration._support.process import group_members, owned_proxy, owned_proxy_process from integration._support.wire import Reply, Request, wire_server +from openai import AsyncOpenAI, OpenAI @pytest.mark.covers("other.observability.guardrails.rewrite_reaches_correct_anthropic_positions") @@ -208,8 +215,6 @@ def test_guardrail_denial_prevents_provider_and_preserves_allowed_control(gatewa with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: model: Final = scenario.model() key: Final = scenario.key(models=[model]) - import httpx - with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: observed.get("/__observations") denied: Final = candidate.request( @@ -500,3 +505,687 @@ def test_request_selected_mcp_guardrail_blocks_direct_and_virtual_calls(gateway: assert len(calls) == 1 assert calls[0]["body"]["params"]["name"] == tool assert calls[0]["body"]["params"]["arguments"] == arguments + + +_RESPONSES_DENIAL: Final = "This model is not currently available." + + +def _deny_guardrail(name: str, denial: str = _RESPONSES_DENIAL) -> dict[str, object]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "custom_code", + "mode": "pre_call", + "default_on": False, + "custom_code": (f"def apply_guardrail(inputs, request_data, input_type):\n return block({denial!r})\n"), + }, + } + + +def _responses_denial_config(tmp_path: Path, identity: str, denial: str = _RESPONSES_DENIAL) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [_deny_guardrail(identity, denial)] + path: Final = tmp_path / "responses-deny.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _assert_blocked_message_item(item: dict[str, object], response: dict[str, object]) -> None: + assert item["type"] == "message", item + assert item["role"] == "assistant", item + assert item["status"] == "completed", item + assert str(item["id"]).startswith("msg_"), item + assert item["content"] == [{"type": "output_text", "text": _RESPONSES_DENIAL, "annotations": []}], item + assert response["status"] == "completed", response + usage: Final = response["usage"] + assert isinstance(usage, dict), response + assert (usage["input_tokens"], usage["output_tokens"], usage["total_tokens"]) == (0, 0, 0), usage + + +def _response_id(index: int, response: httpx.Response) -> str: + assert response.status_code == 200, (index, response.text) + if index % 3 == 0: + assert response.headers["content-type"].startswith("text/event-stream"), response.text + return str(_blocked_stream_events(response.text)[-1]["response"]["id"]) + if index % 3 == 1: + assert response.headers["content-type"].startswith("text/event-stream"), response.text + blocked: Final = _blocked_stream_events(response.text)[-1]["response"] + _assert_blocked_message_item(blocked["output"][0], blocked) + return str(blocked["id"]) + assert response.headers["content-type"].startswith("application/json"), response.text + body: Final = response.json() + _assert_blocked_message_item(body["output"][0], body) + return str(body["id"]) + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_streams_typed_message") +def test_responses_pre_call_denial_streams_sse_with_typed_message_item(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), ( + response.headers["content-type"], + response.text, + ) + lines: Final = tuple(line for line in response.text.split("\n") if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", response.text + events: Final = tuple(json.loads(line.removeprefix("data: ")) for line in lines[:-1]) + kinds: Final = tuple(event["type"] for event in events) + assert tuple(kind for kind in kinds if kind != "response.output_text.delta") == ( + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + ), kinds + assert kinds.index("response.output_text.delta") == kinds.index("response.content_part.added") + 1, kinds + assert "".join(event["delta"] for event in events if event["type"] == "response.output_text.delta") == ( + _RESPONSES_DENIAL + ) + completed: Final = events[-1]["response"] + assert completed["output"] == [events[-2]["item"]], (completed, events[-2]) + _assert_blocked_message_item(completed["output"][0], completed) + assert observed.get("/__observations").json()["requests"] == [] + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_returns_typed_message") +def test_responses_pre_call_denial_returns_json_with_typed_message_item(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json"), response.headers["content-type"] + body: Final = response.json() + assert body["object"] == "response", body + assert len(body["output"]) == 1, body + _assert_blocked_message_item(body["output"][0], body) + assert observed.get("/__observations").json()["requests"] == [] + + +_RESPONSES_OUTPUT_DENIAL: Final = "Output withheld by policy." +_UPSTREAM_INPUT_TOKENS: Final = 20 +_UPSTREAM_OUTPUT_TOKENS: Final = 20 +_UPSTREAM_TOTAL_TOKENS: Final = 40 + + +def _responses_output_denial_config(tmp_path: Path, identity: str, model: str) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "custom_code", + "mode": "post_call", + "default_on": False, + "custom_code": ( + "def apply_guardrail(inputs, request_data, input_type):\n" + f" return block({_RESPONSES_OUTPUT_DENIAL!r})\n" + ), + }, + } + ] + config["policies"] = { + f"{identity}-pipeline": { + "guardrails": {"add": [identity]}, + "pipeline": { + "mode": "post_call", + "steps": [ + { + "guardrail": identity, + "on_pass": "allow", + "on_fail": "modify_response", + "modify_response_message": _RESPONSES_OUTPUT_DENIAL, + } + ], + }, + } + } + config["policy_attachments"] = [{"policy": f"{identity}-pipeline", "models": [model]}] + path: Final = tmp_path / "responses-output-deny.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _blocked_stream_events(text: str) -> tuple[dict[str, object], ...]: + lines: Final = tuple(line for line in text.split("\n") if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", text + return tuple(json.loads(line.removeprefix("data: ")) for line in lines[:-1]) + + +def _dead_api_base() -> str: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + port: Final = reserve.getsockname()[1] + return f"http://127.0.0.1:{port}/v1" + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_openai_sdk_streams_typed_message") +def test_responses_pre_call_denial_openai_sdk_streams_typed_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + client: Final = OpenAI( + base_url=f"{candidate.client.base_url}/v1", api_key=candidate.key, max_retries=0, timeout=15 + ) + events: Final = tuple( + client.responses.create(model=model, input="say hi", stream=True, extra_body={"guardrails": [identity]}) + ) + assert events[-1].type == "response.completed", [event.type for event in events] + completed: Final = events[-1].response + assert completed is not None and len(completed.output) == 1, completed + item: Final = completed.output[0] + assert item.type == "message", item + assert item.role == "assistant" and item.status == "completed", item + assert item.content[0].type == "output_text" and item.content[0].text == _RESPONSES_DENIAL, item.content + assert completed.usage is not None and completed.usage.total_tokens == 0, completed.usage + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_openai_async_sdk_streams_typed_message") +async def test_responses_pre_call_denial_openai_async_sdk_streams_typed_message( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + client: Final = AsyncOpenAI( + base_url=f"{candidate.client.base_url}/v1", api_key=candidate.key, max_retries=0, timeout=15 + ) + stream: Final = await client.responses.create( + model=model, input="say hi", stream=True, extra_body={"guardrails": [identity]} + ) + kinds: Final = [event.type async for event in stream] + assert kinds[-1] == "response.completed", kinds + assert "response.output_text.delta" in kinds, kinds + assert "response.in_progress" in kinds, kinds + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_openai_sdk_returns_typed_message") +def test_responses_pre_call_denial_openai_sdk_returns_typed_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + client: Final = OpenAI( + base_url=f"{candidate.client.base_url}/v1", api_key=candidate.key, max_retries=0, timeout=15 + ) + body: Final = client.responses.create(model=model, input="say hi", extra_body={"guardrails": [identity]}) + assert body.object == "response" and body.status == "completed", body + assert len(body.output) == 1, body.output + item: Final = body.output[0] + assert item.type == "message" and item.role == "assistant", item + assert item.content[0].type == "output_text" and item.content[0].text == _RESPONSES_DENIAL, item.content + assert body.output_text == _RESPONSES_DENIAL, body + assert body.usage is not None and body.usage.total_tokens == 0, body.usage + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_false_returns_json") +def test_responses_pre_call_denial_stream_false_returns_json(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "stream": False, "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json"), response.text + body: Final = response.json() + _assert_blocked_message_item(body["output"][0], body) + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_string_true_returns_json") +def test_responses_pre_call_denial_stream_string_true_returns_json(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "stream": "true", "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json"), ( + response.headers["content-type"], + response.text, + ) + body: Final = response.json() + _assert_blocked_message_item(body["output"][0], body) + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_event_vocabulary") +def test_responses_pre_call_denial_stream_event_vocabulary(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + second: Final = "guardrail-2-" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + loaded: Final = yaml.safe_load(config.read_text()) + loaded["guardrails"].append(_deny_guardrail(second)) + config.write_text(yaml.safe_dump(loaded)) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity, second]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + kinds: Final = {event["type"] for event in events} + assert kinds == { + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + }, kinds + item_done: Final = tuple(event for event in events if event["type"] == "response.output_item.done") + assert len(item_done) == 1, events + assert len(events[-1]["response"]["output"]) == 1, events[-1] + assert observed.get("/__observations").json()["requests"] == [] + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_large_denial_text") +def test_responses_pre_call_denial_stream_large_denial_text(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + denial: Final = ("Denied: " + "mixed ascii and unicode text " * 200 + "fin")[:5000] + config: Final = _responses_denial_config(tmp_path, identity, denial) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + assert "".join(event["delta"] for event in events if event["type"] == "response.output_text.delta") == denial + done: Final = next(event for event in events if event["type"] == "response.output_text.done") + assert done["text"] == denial, done + completed: Final = events[-1]["response"] + assert completed["output"][0]["content"][0]["text"] == denial, completed + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_requests_have_distinct_ids") +def test_responses_pre_call_denial_stream_requests_have_distinct_ids(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + responses: Final = tuple( + candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + for _ in range(2) + ) + completed: Final = tuple(_blocked_stream_events(response.text)[-1]["response"] for response in responses) + for response in responses: + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + assert completed[0]["id"] != completed[1]["id"], completed + assert completed[0]["output"][0]["id"] != completed[1]["output"][0]["id"], completed + assert observed.get("/__observations").json()["requests"] == [] + + +def _register_named_model(candidate: Gateway, name: str, api_base: str | None = None, **parameters: object) -> str: + created: Final = candidate.post( + "/model/new", + { + "model_name": name, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "integration-provider-key", + "api_base": api_base or f"{candidate.upstream_url}/v1", + **parameters, + }, + }, + ) + return str(created["model_info"]["id"]) + + +@pytest.mark.covers("other.observability.guardrails.responses_post_call_pipeline_denial_streams_real_usage") +def test_responses_post_call_pipeline_denial_streams_real_usage(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + model: Final = f"integration-{uuid.uuid4().hex}" + config: Final = _responses_output_denial_config(tmp_path, identity, model) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate: + model_id: Final = _register_named_model(candidate, model, use_chat_completions_api=True) + try: + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": f"say hi {uuid.uuid4().hex}", "stream": True} + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + assert events[-1]["type"] == "response.completed", events + completed: Final = events[-1]["response"] + item: Final = completed["output"][0] + assert item["type"] == "message" and item["role"] == "assistant", item + assert item["content"][0]["type"] == "output_text", item + assert item["content"][0]["text"] == _RESPONSES_OUTPUT_DENIAL, item + usage: Final = completed["usage"] + assert ( + usage["input_tokens"], + usage["output_tokens"], + usage["total_tokens"], + ) == (_UPSTREAM_INPUT_TOKENS, _UPSTREAM_OUTPUT_TOKENS, _UPSTREAM_TOTAL_TOKENS), usage + finally: + candidate.post("/model/delete", {"id": model_id}) + + +@pytest.mark.covers("other.observability.guardrails.responses_post_call_pipeline_denial_returns_real_usage") +def test_responses_post_call_pipeline_denial_returns_real_usage(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + model: Final = f"integration-{uuid.uuid4().hex}" + config: Final = _responses_output_denial_config(tmp_path, identity, model) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate: + model_id: Final = _register_named_model(candidate, model, use_chat_completions_api=True) + try: + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": f"say hi {uuid.uuid4().hex}"} + ) + assert response.status_code == 200, response.text + body: Final = response.json() + item: Final = body["output"][0] + assert item["type"] == "message" and item["role"] == "assistant", item + assert item["content"][0]["type"] == "output_text", item + assert item["content"][0]["text"] == _RESPONSES_OUTPUT_DENIAL, item + usage: Final = body["usage"] + assert ( + usage["input_tokens"], + usage["output_tokens"], + usage["total_tokens"], + ) == (_UPSTREAM_INPUT_TOKENS, _UPSTREAM_OUTPUT_TOKENS, _UPSTREAM_TOTAL_TOKENS), usage + finally: + candidate.post("/model/delete", {"id": model_id}) + + +@pytest.mark.covers("other.observability.guardrails.responses_denial_requires_authentication") +def test_responses_denial_requires_authentication(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "guardrails": [identity]}, key="sk-invalid" + ) + assert response.status_code == 401, (response.status_code, response.text) + assert response.json()["error"]["type"] == "token_not_found_in_db", response.text + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_does_not_reach_upstream") +def test_responses_pre_call_denial_stream_does_not_reach_upstream(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=_dead_api_base()) + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + events: Final = _blocked_stream_events(response.text) + completed: Final = events[-1]["response"] + _assert_blocked_message_item(completed["output"][0], completed) + + +@pytest.mark.covers("other.observability.guardrails.responses_unguarded_stream_reaches_upstream") +def test_responses_unguarded_stream_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + dead: Final = scenario.model(api_base=_dead_api_base()) + denied: Final = candidate.request( + "POST", + "/v1/responses", + {"model": dead, "input": "say hi", "stream": True, "guardrails": [identity]}, + ) + assert denied.status_code == 200, denied.text + model: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": f"say hi {uuid.uuid4().hex}", "stream": True, "guardrails": []}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + assert "response.completed" in response.text, response.text + requests: Final = eventually( + lambda: observed.get("/__observations").json()["requests"], + lambda values: len(values) >= 1, + seconds=30, + ) + assert len(requests) == 1, requests + assert requests[0]["path"] == "/v1/chat/completions", requests + + +@pytest.mark.covers("other.observability.guardrails.chat_pre_call_denial_streams_content_filter") +def test_chat_pre_call_denial_streams_content_filter(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "stream": True, + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + lines: Final = tuple(line for line in response.text.split("\n") if line.startswith("data: ")) + assert lines[-1] == "data: [DONE]", response.text + chunks: Final = tuple(json.loads(line.removeprefix("data: ")) for line in lines[:-1]) + assert chunks[0]["choices"][0]["delta"]["content"] == _RESPONSES_DENIAL, chunks + assert chunks[-1]["choices"][0]["finish_reason"] == "stop", chunks + + +@pytest.mark.covers("other.observability.guardrails.chat_pre_call_denial_returns_content_filter") +def test_chat_pre_call_denial_returns_content_filter(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + choice: Final = body["choices"][0] + assert choice["finish_reason"] == "content_filter", body + assert choice["message"]["content"] == _RESPONSES_DENIAL, body + assert ( + body["usage"]["prompt_tokens"], + body["usage"]["completion_tokens"], + body["usage"]["total_tokens"], + ) == (0, 0, 0), body["usage"] + + +@pytest.mark.covers("other.observability.guardrails.messages_pre_call_denial_returns_message") +def test_messages_pre_call_denial_returns_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "max_tokens": 16, + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["type"] == "message" and body["role"] == "assistant", body + assert body["content"] == [{"type": "text", "text": _RESPONSES_DENIAL}], body + assert body["stop_reason"] == "end_turn", body + assert (body["usage"]["input_tokens"], body["usage"]["output_tokens"]) == (0, 0), body + + +@pytest.mark.covers("other.observability.guardrails.messages_pre_call_denial_streams_message") +def test_messages_pre_call_denial_streams_message(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "messages": [{"role": "user", "content": "say hi"}], + "max_tokens": 16, + "stream": True, + "guardrails": [identity], + }, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text + lines: Final = tuple(line for line in response.text.split("\n") if line.startswith("data: ")) + assert len(lines) == 1, response.text + body: Final = json.loads(lines[0].removeprefix("data: ")) + assert body["type"] == "message" and body["role"] == "assistant", body + assert body["content"] == [{"type": "text", "text": _RESPONSES_DENIAL}], body + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_writes_zero_spend_row") +def test_responses_pre_call_denial_writes_zero_spend_row(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + response: Final = candidate.request( + "POST", "/v1/responses", {"model": model, "input": "say hi", "guardrails": [identity]} + ) + assert response.status_code == 200, response.text + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, total_tokens FROM "LiteLLM_SpendLogs" WHERE model=%s AND call_type=%s', + (model, "aresponses"), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert float(rows[0]["spend"]) == 0, rows + assert rows[0]["total_tokens"] == 0, rows + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_survives_worker_burst") +def test_responses_pre_call_denial_stream_survives_worker_burst(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=_dead_api_base()) + healthy: Final = scenario.model(use_chat_completions_api=True) + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as observed: + observed.get("/__observations") + + def burst(index: int) -> httpx.Response: + if index % 3 == 0: + return candidate.request( + "POST", + "/v1/responses", + {"model": healthy, "input": f"say hi {uuid.uuid4().hex} {index}", "stream": True}, + ) + stream: Final = index % 3 == 1 + return candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": f"say hi {index}", "stream": stream, "guardrails": [identity]}, + ) + + with ThreadPoolExecutor(max_workers=8) as pool: + responses: Final = tuple(pool.map(burst, range(30))) + response_ids: Final = frozenset(_response_id(index, response) for index, response in enumerate(responses)) + assert len(response_ids) == 30, response_ids + assert len(observed.get("/__observations").json()["requests"]) == 10 + + +@pytest.mark.covers("other.observability.guardrails.responses_pre_call_denial_stream_survives_worker_kill") +def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway, tmp_path: Path) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + config: Final = _responses_denial_config(tmp_path, identity) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + with candidate.scenario() as scenario: + model: Final = scenario.model(api_base=_dead_api_base()) + members: Final = tuple( + member for member in group_members(owned.process.pid) if member.pid != owned.process.pid + ) + children: Final = tuple(member.pid for member in members) + workers: Final = tuple( + member.pid for member in members if any("spawn_main" in part for part in member.cmdline()) + ) + assert len(workers) >= 2, workers + os.kill(workers[0], signal.SIGKILL) + expected: Final = len(children) + eventually( + lambda: tuple( + member.pid + for member in group_members(owned.process.pid) + if member.pid != owned.process.pid + and member.is_running() + and member.status() != psutil.STATUS_ZOMBIE + ), + lambda pids: len(pids) >= expected and any(pid not in children for pid in pids), + seconds=30, + ) + + def burst(index: int) -> httpx.Response: + return candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": f"say hi {index}", "stream": True, "guardrails": [identity]}, + ) + + with ThreadPoolExecutor(max_workers=5) as pool: + responses: Final = tuple(pool.map(burst, range(10))) + for response in responses: + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream"), response.text diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index e684aa55b33..656dc33e88c 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -3,6 +3,7 @@ Test for response_api_endpoints/endpoints.py """ import unittest +from collections.abc import Mapping from typing import Any, Final, Literal from unittest.mock import AsyncMock, MagicMock, patch @@ -14,6 +15,7 @@ from httpx import Response import litellm from litellm.proxy.proxy_server import app +from litellm.types.llms.openai import ResponsesAPIResponse @pytest.mark.asyncio @@ -2193,6 +2195,59 @@ class TestCursorGateRecognizesRoutingGroups: assert "reasoning_effort" not in resolved +BLOCK_MESSAGE = "Content flagged by policy, response withheld" + + +def _post_blocked_responses( + original_response: ResponsesAPIResponse | litellm.ModelResponse | None, + payload: Mapping[str, object] | None = None, +) -> httpx.Response: + from litellm.integrations.custom_guardrail import ModifyResponseException + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + exc = ModifyResponseException( + message=BLOCK_MESSAGE, + model="gpt-4o-mini", + request_data={"model": "gpt-4o-mini", "input": "hi"}, + guardrail_name="zero-usage-regression", + original_response=original_response, + ) + mock_proxy_logging = MagicMock() + mock_proxy_logging.post_call_failure_hook = AsyncMock() + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="sk-test", request_route="/v1/responses" + ) + body = {"model": "gpt-4o-mini", "input": "Write a haiku about token accounting"} + if payload: + body.update(payload) + try: + with ( + patch( + "litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=AsyncMock(side_effect=exc), + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), + ): + client = TestClient(app) + return client.post("/v1/responses", json=body, headers={"Authorization": "Bearer sk-1234"}) + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +def _assert_blocked_output_item(item: Mapping[str, object], text: str) -> None: + assert item["type"] == "message" + assert item["id"].startswith("msg_") + assert item["role"] == "assistant" + assert item["status"] == "completed" + assert item["content"][0]["type"] == "output_text" + assert item["content"][0]["text"] == text + + +def _sse_data_frames(text: str) -> list[str]: + return [line.removeprefix("data: ").strip() for line in text.splitlines() if line.startswith("data: ")] + + class TestGuardrailBlockedResponsesUsage: """Regression tests for https://github.com/BerriAI/litellm/issues/36880. @@ -2202,38 +2257,7 @@ class TestGuardrailBlockedResponsesUsage: e.original_response, exactly like /v1/chat/completions already does.""" def _post_blocked_responses(self, original_response): - from litellm.integrations.custom_guardrail import ModifyResponseException - from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy.auth.user_api_key_auth import user_api_key_auth - - exc = ModifyResponseException( - message="Content flagged by policy, response withheld", - model="gpt-4o-mini", - request_data={"model": "gpt-4o-mini", "input": "hi"}, - guardrail_name="zero-usage-regression", - original_response=original_response, - ) - mock_proxy_logging = MagicMock() - mock_proxy_logging.post_call_failure_hook = AsyncMock() - app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( - api_key="sk-test", request_route="/v1/responses" - ) - try: - with ( - patch( - "litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", - new=AsyncMock(side_effect=exc), - ), - patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging), - ): - client = TestClient(app) - return client.post( - "/v1/responses", - json={"model": "gpt-4o-mini", "input": "Write a haiku about token accounting"}, - headers={"Authorization": "Bearer sk-1234"}, - ) - finally: - app.dependency_overrides.pop(user_api_key_auth, None) + return _post_blocked_responses(original_response) def test_post_call_block_reports_real_upstream_usage(self): from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse @@ -2429,6 +2453,66 @@ class TestResponsesInputTokens: assert response.json()["error"]["message"] == "rate limited" +class TestGuardrailBlockedResponsesShape: + """A pre_call block raises ModifyResponseException before any provider call. + + The reply must satisfy the Responses API contract the request selected: + stream=true answers SSE ending in one response.completed whose output[0] is + a completed assistant message item with output_text content, and a plain + POST answers JSON with the same item, both with the usage the blocked call + consumed (zero for pre_call).""" + + def test_non_stream_block_is_a_completed_assistant_message(self): + response = _post_blocked_responses(None) + + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json") + body = response.json() + _assert_blocked_output_item(body["output"][0], BLOCK_MESSAGE) + assert body["usage"]["total_tokens"] == 0 + + def test_stream_block_answers_sse_with_completed_event(self): + response = _post_blocked_responses(None, payload={"stream": True}) + + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("text/event-stream") + frames = _sse_data_frames(response.text) + assert frames[-1] == "[DONE]" + events = [json.loads(frame) for frame in frames[:-1]] + types = [event["type"] for event in events] + assert "response.created" in types + completed = [event for event in events if event["type"] == "response.completed"] + assert len(completed) == 1 + completed_response = completed[0]["response"] + _assert_blocked_output_item(completed_response["output"][0], BLOCK_MESSAGE) + assert completed_response["usage"]["total_tokens"] == 0 + delta_text = "".join(event["delta"] for event in events if event["type"] == "response.output_text.delta") + assert delta_text == BLOCK_MESSAGE + + def test_stream_block_keeps_upstream_usage(self): + from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + original = ResponsesAPIResponse( + id="resp_upstream", + created_at=1, + model="gpt-4o-mini", + object="response", + output=[], + status="completed", + usage=ResponseAPIUsage(input_tokens=14, output_tokens=20, total_tokens=34), + ) + + response = _post_blocked_responses(original, payload={"stream": True}) + + assert response.status_code == 200, response.text + frames = _sse_data_frames(response.text) + completed = [json.loads(frame) for frame in frames[:-1] if json.loads(frame)["type"] == "response.completed"] + usage = completed[0]["response"]["usage"] + assert usage["input_tokens"] == 14 + assert usage["output_tokens"] == 20 + assert usage["total_tokens"] == 34 + + def test_responses_routes_document_response_models_in_openapi_schema(): from typing import cast diff --git a/tests/test_litellm/proxy/test_blocked_response_usage.py b/tests/test_litellm/proxy/test_blocked_response_usage.py index 4f20f35e94b..90d861be8e0 100644 --- a/tests/test_litellm/proxy/test_blocked_response_usage.py +++ b/tests/test_litellm/proxy/test_blocked_response_usage.py @@ -4,7 +4,7 @@ proxy endpoints (/v1/chat/completions, /v1/completions, and /v1/responses). A post-call block replaces the LLM response with the violation message, but the upstream call already consumed tokens. `_blocked_response_usage` (and its -Responses API counterpart `_blocked_responses_api_usage`) reports that real +Responses API counterpart `blocked_responses_api_usage`) reports that real usage (carried on `ModifyResponseException.original_response`) rather than zero; a pre-call block never invoked the LLM, so usage is zero. """ @@ -91,8 +91,8 @@ def test_responses_api_blocked_reply_carries_real_usage(): """ import time - from litellm.proxy.response_api_endpoints.endpoints import ( - _blocked_responses_api_usage, + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_responses_api_usage, ) original_response = ResponsesAPIResponse( @@ -105,7 +105,7 @@ def test_responses_api_blocked_reply_carries_real_usage(): usage=ResponseAPIUsage(input_tokens=14, output_tokens=20, total_tokens=34), ) - usage = _blocked_responses_api_usage(original_response) + usage = blocked_responses_api_usage(original_response) assert usage.input_tokens == 14 assert usage.output_tokens == 20 @@ -114,11 +114,11 @@ def test_responses_api_blocked_reply_carries_real_usage(): def test_responses_api_blocked_reply_zero_usage_when_no_original_response(): """Pre-call block has no original_response, so usage must be zero.""" - from litellm.proxy.response_api_endpoints.endpoints import ( - _blocked_responses_api_usage, + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_responses_api_usage, ) - usage = _blocked_responses_api_usage(None) + usage = blocked_responses_api_usage(None) assert usage.input_tokens == 0 assert usage.output_tokens == 0 @@ -128,14 +128,14 @@ def test_responses_api_blocked_reply_zero_usage_when_no_original_response(): def test_responses_api_blocked_reply_maps_bridged_chat_usage(): """A chat model bridged through /v1/responses blocks with a ModelResponse whose Usage fields must map prompt_tokens -> input_tokens and completion_tokens -> output_tokens.""" - from litellm.proxy.response_api_endpoints.endpoints import ( - _blocked_responses_api_usage, + from litellm.llms.base_llm.guardrail_translation.utils import ( + blocked_responses_api_usage, ) resp = litellm.ModelResponse() resp.usage = litellm.Usage(prompt_tokens=14, completion_tokens=18, total_tokens=32) - usage = _blocked_responses_api_usage(resp) + usage = blocked_responses_api_usage(resp) assert usage.input_tokens == 14 assert usage.output_tokens == 18 From b396b0b72499e7c79b280d8821a44a8b84835a7a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:05:55 -0700 Subject: [PATCH 12/51] feat(e2e): read management routes back from the control plane replicas (#43373) * feat(e2e): read management routes back from the control plane replicas The suite's management read-backs (/key/info, /team/info and friends) polled the same replica list as the data plane. On a componentized stack whose LITELLM_PROXY_REPLICA_URLS names the gateway pods directly, that list answers those routes 404, since a gateway pod trims the management routes at startup. A new LITELLM_CONTROL_PLANE_REPLICA_URLS names the addresses a management read-back polls instead: an exported list wins, and when it is unset the old rule stands, the data-plane replicas while the control plane shares the suite's base URL and the control-plane base alone once it is split. build_proxy_client takes the list as control_replica_urls and read_back_everywhere picks its replicas per path, the way the rest of the client already does. * fix(e2e): derive the control replicas of a client built for another proxy --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- tests/e2e/CONTRIBUTING.md | 2 +- tests/e2e/claude_code/_env.py | 1 + tests/e2e/claude_code/conftest.py | 1 + tests/e2e/e2e_config.py | 61 +++++++++++++++- tests/e2e/mcp/oauth_gateway.py | 1 + tests/e2e/proxy_client.py | 65 +++++++++++------ tests/e2e/test_proxy_client.py | 114 ++++++++++++++++++++++++++++-- 7 files changed, 215 insertions(+), 30 deletions(-) diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index 7e1f516422e..8e221b2da5e 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/`, `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 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. The Buildkite PR stack exports its two gateway pods the same way and, because those pods sit behind one router base that also fronts the backend, names that base in `LITELLM_CONTROL_PLANE_REPLICA_URLS` so management read-backs poll the plane that serves them instead of the gateway pods, which trim management routes at startup and answer them 404. 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 diff --git a/tests/e2e/claude_code/_env.py b/tests/e2e/claude_code/_env.py index 889d8f848dd..431be349a95 100644 --- a/tests/e2e/claude_code/_env.py +++ b/tests/e2e/claude_code/_env.py @@ -100,5 +100,6 @@ def require_proxy_client( master_key=cfg.api_key, control_plane_base_url=cfg.base_url, replica_urls=(cfg.base_url,), + control_replica_urls=(cfg.base_url,), ) return ProxyClientConfig(client=client, api_key=cfg.api_key) diff --git a/tests/e2e/claude_code/conftest.py b/tests/e2e/claude_code/conftest.py index bf226161267..c69f2dae462 100644 --- a/tests/e2e/claude_code/conftest.py +++ b/tests/e2e/claude_code/conftest.py @@ -595,6 +595,7 @@ def _build_control_plane_client(proxy_config: ProxyConfig): master_key=proxy_config.api_key, control_plane_base_url=proxy_config.base_url, replica_urls=(proxy_config.base_url,), + control_replica_urls=(proxy_config.base_url,), ) diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index ca3b74281ae..b2682c04841 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -7,6 +7,7 @@ environment so the same tests run against localhost or a deployed proxy. from __future__ import annotations import os +from dataclasses import dataclass import time import uuid from pathlib import Path @@ -33,12 +34,68 @@ CONTROL_PLANE_BASE_URL = os.environ.get( ).rstrip("/") +def split_replica_urls(raw: str) -> tuple[str, ...]: + return tuple(dict.fromkeys(url.strip().rstrip("/") for url in raw.split(",") if url.strip())) + + def parse_replica_urls(raw: str, fallback: str) -> tuple[str, ...]: - urls: Final = tuple(dict.fromkeys(url.strip().rstrip("/") for url in raw.split(",") if url.strip())) - return urls or (fallback,) + return split_replica_urls(raw) or (fallback,) + + +def parse_control_plane_replica_urls( + raw: str, *, control_plane_base_url: str, base_url: str, replica_urls: tuple[str, ...] +) -> tuple[str, ...]: + """The replicas a management read-back polls. LITELLM_CONTROL_PLANE_REPLICA_URLS + names them outright; unset, they follow the two base URLs: every data-plane + replica when the planes share a base (a monolith serves every route from every + replica) and the control-plane base alone when they differ. A stack sets it when + LITELLM_PROXY_REPLICA_URLS names gateway pods behind a shared router base, since + a gateway trims the management routes at startup and answers them 404.""" + explicit: Final = split_replica_urls(raw) + if explicit: + return explicit + return replica_urls if control_plane_base_url == base_url else (control_plane_base_url,) PROXY_REPLICA_URLS: Final = parse_replica_urls(os.environ.get("LITELLM_PROXY_REPLICA_URLS", ""), PROXY_BASE_URL) +CONTROL_PLANE_REPLICA_URLS: Final = parse_control_plane_replica_urls( + os.environ.get("LITELLM_CONTROL_PLANE_REPLICA_URLS", ""), + control_plane_base_url=CONTROL_PLANE_BASE_URL, + base_url=PROXY_BASE_URL, + replica_urls=PROXY_REPLICA_URLS, +) + + +@dataclass(frozen=True, slots=True) +class StackEndpoints: + base_url: str + control_plane_base_url: str + replica_urls: tuple[str, ...] + control_replica_urls: tuple[str, ...] + + def control_replica_urls_for( + self, *, base_url: str, control_plane_base_url: str, replica_urls: tuple[str, ...] + ) -> tuple[str, ...]: + """The control replicas a client built for these endpoints polls when its caller names none: + this stack's own list for this stack's endpoints, since an exported list describes one stack only, + and the base-URL rule for any other proxy.""" + if (base_url, control_plane_base_url, replica_urls) == ( + self.base_url, + self.control_plane_base_url, + self.replica_urls, + ): + return self.control_replica_urls + return parse_control_plane_replica_urls( + "", control_plane_base_url=control_plane_base_url, base_url=base_url, replica_urls=replica_urls + ) + + +ENV_STACK: Final = StackEndpoints( + base_url=PROXY_BASE_URL, + control_plane_base_url=CONTROL_PLANE_BASE_URL, + replica_urls=PROXY_REPLICA_URLS, + control_replica_urls=CONTROL_PLANE_REPLICA_URLS, +) UI_USERNAME = os.environ.get("E2E_UI_USERNAME", "admin") UI_PASSWORD = os.environ.get("E2E_UI_PASSWORD", MASTER_KEY) diff --git a/tests/e2e/mcp/oauth_gateway.py b/tests/e2e/mcp/oauth_gateway.py index 82bb5f7ba0b..b328c81687b 100644 --- a/tests/e2e/mcp/oauth_gateway.py +++ b/tests/e2e/mcp/oauth_gateway.py @@ -187,6 +187,7 @@ def owned_gateway(idp: Keycloak, directory: Path, cleanup: ExitStack) -> OAuthGa base_url=base_url, control_plane_base_url=base_url, replica_urls=(base_url,), + control_replica_urls=(base_url,), master_key=os.environ["LITELLM_MASTER_KEY"], ), _environment=environment, diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py index 51ea9fbe7bd..4ea83e4b0d3 100644 --- a/tests/e2e/proxy_client.py +++ b/tests/e2e/proxy_client.py @@ -20,6 +20,7 @@ from typing import Final, Literal from e2e_config import ( CONTROL_PLANE_BASE_URL, + ENV_STACK, MASTER_KEY, POLL_INTERVAL, POLL_TIMEOUT, @@ -530,16 +531,16 @@ class ProxyClient: response_type: type[R], converged: Callable[[Result[R]], bool], ) -> Mapping[str, Result[R]]: - """GET `path` under the master key on every replica in PROXY_REPLICA_URLS (the - data-plane URL alone when the stack exports no per-gateway addresses), polling - each to poll_timeout until its read satisfies `converged`. Returns that read per - replica, or fails naming the first replica that never converged and its last - read. Behind a load balancer the single address proves one replica converged, - not all of them; only per-gateway addresses make this a fleet-wide proof.""" + """GET `path` under the master key on every replica that serves it (see + replicas_for), polling each to poll_timeout until its read satisfies + `converged`. Returns that read per replica, or fails naming the first replica + that never converged and its last read. Behind a load balancer the single + address proves one replica converged, not all of them; only per-replica + addresses make this a fleet-wide proof.""" outcomes: Final = await_converged_everywhere( { url: self._body_poller(transport, path, params, response_type) - for url, transport in self.replicas.items() + for url, transport in self.replicas_for(path).items() }, converged=converged, timeout=self.poll_timeout, @@ -765,13 +766,11 @@ class ProxyClient: def replicas_for(self, path: str) -> Mapping[str, Transport]: """The replicas that serve `path`: every data-plane replica for an LLM route, and for a management route the control-plane replicas, since the data-plane - replicas trim management routes and answer them 404. A monolith serves both - from every replica, so a management read-back polls all of them; a split - deployment exposes one control-plane address (there is one backend process - behind it on the stack these suites run against), so it polls that. A - control plane fronting several backends would need its own replica list to - prove each one converged, the way PROXY_REPLICA_URLS does for the gateways. - Never empty: a read-back against no replica would assert nothing and pass.""" + replicas trim management routes and answer them 404. CONTROL_PLANE_REPLICA_URLS + names those (see e2e_config): every data-plane replica for a monolith, the + control plane's own address for a split deployment, and the stack's own list + when its gateway pods sit behind a shared router base. Never empty: a + read-back against no replica would assert nothing and pass.""" replicas: Final = self.control_replicas if is_control_plane_path(path) else self.replicas assert replicas, f"no replica is configured to serve {path}, so a read-back there would prove nothing" return replicas @@ -1132,6 +1131,7 @@ def build_proxy_client( master_key: str = MASTER_KEY, control_plane_base_url: str = CONTROL_PLANE_BASE_URL, replica_urls: tuple[str, ...] = PROXY_REPLICA_URLS, + control_replica_urls: tuple[str, ...] | None = None, ) -> ProxyClient: """The ProxyClient every suite's client is built from: a SplitTransport that routes LLM calls to the data plane (PROXY_BASE_URL) and management/admin calls to the @@ -1139,15 +1139,24 @@ def build_proxy_client( base URLs are the same for a monolithic proxy, so routing is then a no-op. ``replica_urls`` (PROXY_REPLICA_URLS) names every data-plane replica the model barrier polls directly; it is the data-plane URL itself unless the stack - exports each gateway's own address. Management read-backs poll those same - replicas when the two planes share a base URL (a monolith, where every replica - serves every route) and the control plane alone when they differ (a split - deployment, where the data-plane replicas do not serve management routes). + exports each gateway's own address. ``control_replica_urls`` + (CONTROL_PLANE_REPLICA_URLS) names the replicas a management read-back polls: + those same replicas when the two planes share a base URL (a monolith, where + every replica serves every route), the control plane alone when they differ (a + split deployment, where the data-plane replicas do not serve management + routes), or the list the stack exports when its gateway pods sit behind a + shared router base, since a gateway pod trims management routes and its + address cannot stand in for the control plane. The endpoints are injectable for callers that resolve the proxy some other - way than ``e2e_config``'s env names (see ``claude_code/_env.py``); they must - pass all four together, since a caller that overrides only the data plane - would leave management calls and the replica poll pointed at the env defaults. + way than ``e2e_config``'s env names (see ``claude_code/_env.py``); they pass + the three URL parameters together, since a caller that overrides only the + data plane would leave management calls and the replica polls pointed at the + env defaults. An omitted ``control_replica_urls`` is derived from those three + (``ENV_STACK.control_replica_urls_for``): the env stack's own endpoints take + its exported list, any other proxy follows the base-URL rule above, so a + client built for a local test server never reads management state back from + the env proxy. Test-to-proxy traffic always goes over the wire, in every E2E_FIXTURE_MODE: record and replay scope to the proxy's provider-bound calls via the @@ -1170,8 +1179,18 @@ def build_proxy_client( for url in replica_urls } ) - control_replicas: Final = ( - replicas if control_plane_base_url == base_url else MappingProxyType({control_plane_base_url: split.control}) + control_replica_urls_named: Final = ( + control_replica_urls + if control_replica_urls is not None + else ENV_STACK.control_replica_urls_for( + base_url=base_url, control_plane_base_url=control_plane_base_url, replica_urls=replica_urls + ) + ) + control_replicas: Final = MappingProxyType( + { + url: HttpTransport(base_url=url, master_key=master_key, request_timeout=REQUEST_TIMEOUT) + for url in control_replica_urls_named + } ) return ProxyClient( transport=split, diff --git a/tests/e2e/test_proxy_client.py b/tests/e2e/test_proxy_client.py index a8f07ed6dd7..e1615ec5f65 100644 --- a/tests/e2e/test_proxy_client.py +++ b/tests/e2e/test_proxy_client.py @@ -24,7 +24,7 @@ from types import MappingProxyType from typing import Final, cast import pytest -from e2e_config import parse_replica_urls +from e2e_config import StackEndpoints, parse_control_plane_replica_urls, parse_replica_urls from e2e_http import NoBody, Result, Success, without_retries from idp import Keycloak from lifecycle import ResourceManager @@ -110,7 +110,7 @@ def caller_boundary( thread.start() url: Final = f"http://127.0.0.1:{server.server_port}" proxy: Final = build_proxy_client( - base_url=url, control_plane_base_url=url, replica_urls=(url,), master_key="bootstrap" + base_url=url, control_plane_base_url=url, replica_urls=(url,), control_replica_urls=(url,), master_key="bootstrap" ) try: yield ManagementClient(proxy=proxy, master_key="bootstrap"), received @@ -343,6 +343,50 @@ class TestParseReplicaUrls: assert parse_replica_urls(raw, "http://lb") == ("http://127.0.0.1:4010", "http://127.0.0.1:4011") +class TestParseControlPlaneReplicaUrls: + def test_an_exported_list_wins_over_the_base_url_rule(self) -> None: + assert parse_control_plane_replica_urls( + " http://router/, http://router ", + control_plane_base_url="http://router", + base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + ) == ("http://router",) + + def test_unset_with_one_shared_base_follows_the_data_plane_replicas(self) -> None: + assert parse_control_plane_replica_urls( + "", control_plane_base_url="http://lb", base_url="http://lb", replica_urls=("http://pod-1", "http://pod-2") + ) == ("http://pod-1", "http://pod-2") + + def test_unset_with_a_split_control_plane_polls_its_base_alone(self) -> None: + assert parse_control_plane_replica_urls( + "", control_plane_base_url="http://backend", base_url="http://lb", replica_urls=("http://gateway-1",) + ) == ("http://backend",) + + +class TestStackEndpointsControlReplicas: + STACK: Final = StackEndpoints( + base_url="http://router", + control_plane_base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + control_replica_urls=("http://router",), + ) + + def test_the_stacks_own_endpoints_take_its_exported_control_list(self) -> None: + assert self.STACK.control_replica_urls_for( + base_url="http://router", + control_plane_base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + ) == ("http://router",) + + def test_any_other_endpoints_follow_the_base_url_rule(self) -> None: + assert self.STACK.control_replica_urls_for( + base_url="http://router", control_plane_base_url="http://router", replica_urls=("http://10.0.0.1:4000",) + ) == ("http://10.0.0.1:4000",) + assert self.STACK.control_replica_urls_for( + base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1",) + ) == ("http://backend",) + + def _answers(answers: Iterable[str]) -> ReplicaRead[str]: it: Final = iter(answers) return lambda _timeout: next(it) @@ -390,6 +434,7 @@ class TestReplicasFor: base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1", "http://gateway-2"), + control_replica_urls=("http://backend",), ) assert set(client.replicas_for("/key/info")) == {"http://backend"} assert set(client.replicas_for("/project/info")) == {"http://backend"} @@ -400,9 +445,63 @@ class TestReplicasFor: base_url="http://lb", control_plane_base_url="http://lb", replica_urls=("http://pod-1", "http://pod-2"), + control_replica_urls=("http://pod-1", "http://pod-2"), ) assert set(client.replicas_for("/key/info")) == {"http://pod-1", "http://pod-2"} + def test_gateway_pods_behind_one_router_read_management_routes_back_from_the_router(self) -> None: + """The Buildkite PR stack names each gateway pod in PROXY_REPLICA_URLS while + both planes share the router base, so a management read-back polls the + router (CONTROL_PLANE_REPLICA_URLS) rather than the pods, which trim + management routes, while a data-plane read-back still polls every pod.""" + client: Final = build_proxy_client( + base_url="http://router", + control_plane_base_url="http://router", + replica_urls=("http://10.0.0.1:4000", "http://10.0.0.2:4000"), + control_replica_urls=("http://router",), + ) + assert set(client.replicas_for("/key/info")) == {"http://router"} + assert set(client.replicas_for("/v1/models")) == {"http://10.0.0.1:4000", "http://10.0.0.2:4000"} + + def test_a_client_built_for_another_proxy_reads_management_routes_back_from_that_proxy(self) -> None: + """A caller that points the client at its own server (test_provider_cache.py) + names no control list, so the derived one has to follow that server rather + than the env proxy, on a shared base and on split ones alike.""" + local: Final = build_proxy_client( + base_url="http://local", control_plane_base_url="http://local", replica_urls=("http://local",) + ) + assert set(local.replicas_for("/key/info")) == {"http://local"} + assert set(local.replicas_for("/v1/models")) == {"http://local"} + split: Final = build_proxy_client( + base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1",) + ) + assert set(split.replicas_for("/key/info")) == {"http://backend"} + assert set(split.replicas_for("/v1/models")) == {"http://gateway-1"} + + def test_management_read_backs_poll_the_control_replicas_only(self) -> None: + """A gateway pod answers /key/info 404 even after the write landed on the + control plane, so a read-back that polled the data-plane replicas for it + would never converge there.""" + with caller_boundary(status=404) as (pod, pod_headers), caller_boundary() as (router, router_headers): + pod_url: Final = next(iter(pod.proxy.replicas)) + router_url: Final = next(iter(router.proxy.replicas)) + proxy: Final = build_proxy_client( + base_url=router_url, + control_plane_base_url=router_url, + replica_urls=(pod_url,), + control_replica_urls=(router_url,), + master_key="bootstrap", + ) + read: Final = proxy.read_back_everywhere( + "/key/info", + params=NoBody(), + response_type=KeyInfoResponse, + converged=lambda result: isinstance(result, Success), + ) + assert set(read) == {router_url} + assert router_headers.get_nowait() == "Bearer bootstrap" + assert router_headers.empty() and pod_headers.empty() + def test_mcp_admin_routes_read_back_from_every_data_plane_replica(self) -> None: """/v1/mcp/* is a lazily mounted feature, so a data-plane replica serves it too and answers from its own in-memory registry. Routing it to the control @@ -412,6 +511,7 @@ class TestReplicasFor: base_url="http://lb", control_plane_base_url="http://backend", replica_urls=("http://gateway-1", "http://gateway-2"), + control_replica_urls=("http://backend",), ) assert set(client.replicas_for("/v1/mcp/server/abc")) == {"http://gateway-1", "http://gateway-2"} assert set(client.replicas_for("/v1/mcp/toolset/abc")) == {"http://gateway-1", "http://gateway-2"} @@ -531,6 +631,7 @@ class TestSplitCallerPropagation: base_url=data_url, control_plane_base_url=control_url, replica_urls=(data_url,), + control_replica_urls=(control_url,), master_key="bootstrap", ).with_caller(Caller(credential="tenant-token", kind="direct_jwt", role="team_member")) proxy.key_info("owned") @@ -543,8 +644,13 @@ class TestSplitCallerPropagation: response_type=KeyInfoResponse, converged=lambda result: isinstance(result, Success), ) - assert control_headers.get_nowait() == "Bearer tenant-token" - assert control_headers.get_nowait() == "Bearer tenant-token" + proxy.read_back_everywhere( + "/v1/models", + params=NoBody(), + response_type=ModelsListResponse, + converged=lambda result: isinstance(result, Success), + ) + assert tuple(control_headers.get_nowait() for _ in range(3)) == ("Bearer tenant-token",) * 3 assert data_headers.get_nowait() == "Bearer tenant-token" assert control_headers.empty() and data_headers.empty() From 501ef23f4aae19c2bdfee50e7d860107a90e9e7e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:07:09 -0700 Subject: [PATCH 13/51] feat(rust): add the openai_like chat config foundation (#43379) * feat(rust): add the openai_like chat config foundation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): let max_completion_tokens outrank max_tokens and decline refusal responses Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../core/src/chat_completions/common_utils.rs | 2 + litellm-rust/crates/llms/src/lib.rs | 1 + .../crates/llms/src/openai_like/chat/mod.rs | 1 + .../src/openai_like/chat/transformation.rs | 270 ++++++++++++++ .../llms/src/openai_like/common_utils.rs | 58 +++ .../crates/llms/src/openai_like/mod.rs | 2 + .../tests/openai_like_chat_transformation.rs | 343 ++++++++++++++++++ 7 files changed, 677 insertions(+) create mode 100644 litellm-rust/crates/llms/src/openai_like/chat/mod.rs create mode 100644 litellm-rust/crates/llms/src/openai_like/chat/transformation.rs create mode 100644 litellm-rust/crates/llms/src/openai_like/common_utils.rs create mode 100644 litellm-rust/crates/llms/src/openai_like/mod.rs create mode 100644 litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs diff --git a/litellm-rust/crates/core/src/chat_completions/common_utils.rs b/litellm-rust/crates/core/src/chat_completions/common_utils.rs index 4ed39a90366..1c7875c3c33 100644 --- a/litellm-rust/crates/core/src/chat_completions/common_utils.rs +++ b/litellm-rust/crates/core/src/chat_completions/common_utils.rs @@ -3,6 +3,7 @@ use litellm_llms::{ anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG, base_llm::chat::transformation::BaseConfig, bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG, + openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG, }; use serde_json::{Map, Value}; @@ -14,6 +15,7 @@ pub(super) fn chat_completions_provider_config(provider: &str) -> Option<&'stati match provider { "anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG), "bedrock" => Some(&BEDROCK_CHAT_COMPLETIONS_CONFIG), + "openai_like" => Some(&OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG), _ => None, } } diff --git a/litellm-rust/crates/llms/src/lib.rs b/litellm-rust/crates/llms/src/lib.rs index e71a9466c0c..a25b822d0c7 100644 --- a/litellm-rust/crates/llms/src/lib.rs +++ b/litellm-rust/crates/llms/src/lib.rs @@ -7,6 +7,7 @@ pub mod cohere; mod error; pub mod mistral; pub mod openai; +pub mod openai_like; pub mod reducto; pub mod vertex_ai; diff --git a/litellm-rust/crates/llms/src/openai_like/chat/mod.rs b/litellm-rust/crates/llms/src/openai_like/chat/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/chat/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs b/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs new file mode 100644 index 00000000000..4e483ddf8cc --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/chat/transformation.rs @@ -0,0 +1,270 @@ +//! `litellm/llms/openai_like/chat/transformation.py`: the chat config every +//! OpenAI-compatible endpoint shares. The body is already OpenAI-shaped, so +//! parameters pass through verbatim; the port keeps Python's two deviations, +//! the `max_completion_tokens` -> `max_tokens` rename and the usage +//! `*_tokens` null-to-zero sanitize. + +use litellm_auth::{CredentialPlacement, SecretValue}; +use litellm_core_utils::core_helpers::unix_now; +use litellm_types::{ + llms::openai::ChatMessage, + utils::{ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse}, +}; +use serde_json::{Map, Value, json}; + +use crate::{ + Error, + base_llm::{ + auth::AuthScheme, + chat::transformation::{ + BaseConfig, Headers, ProviderChatRequestData, ProviderChatResponseData, + ValidatedEnvironment, + }, + }, + openai_like::common_utils::{complete_openai_like_url, openai_compatible_provider_info}, +}; + +/// OpenAI parameter names the Rust path can place verbatim in the request body. +/// Tool parameters are absent on purpose: the message gate already declines +/// tool-call content, and a `tools` request that did get through would produce +/// a tool-call response this port cannot normalize yet, so it declines before +/// the call instead of after it. +const SUPPORTED_PARAMS: &[(&str, &str)] = &[ + ("frequency_penalty", "frequency_penalty"), + ("logit_bias", "logit_bias"), + ("logprobs", "logprobs"), + ("top_logprobs", "top_logprobs"), + ("max_tokens", "max_tokens"), + ("max_completion_tokens", "max_completion_tokens"), + ("modalities", "modalities"), + ("prediction", "prediction"), + ("n", "n"), + ("presence_penalty", "presence_penalty"), + ("seed", "seed"), + ("stop", "stop"), + ("stream_options", "stream_options"), + ("temperature", "temperature"), + ("top_p", "top_p"), + ("audio", "audio"), + ("web_search_options", "web_search_options"), + ("service_tier", "service_tier"), + ("safety_identifier", "safety_identifier"), + ("prompt_cache_key", "prompt_cache_key"), + ("prompt_cache_retention", "prompt_cache_retention"), + ("store", "store"), + ("response_format", "response_format"), +]; + +/// Call configuration the caller may pass that never enters the request body. +const CONFIG_PARAMS: &[&str] = &["custom_endpoint", "extra_headers", "max_retries"]; + +pub struct OpenAILikeChatConfig; + +pub const OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG: OpenAILikeChatConfig = OpenAILikeChatConfig; + +impl BaseConfig for OpenAILikeChatConfig { + fn supported_openai_param_mappings(&self) -> &'static [(&'static str, &'static str)] { + SUPPORTED_PARAMS + } + + fn get_complete_url( + &self, + api_base: Option<&str>, + _model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + let custom_endpoint = optional_params + .get("custom_endpoint") + .and_then(Value::as_bool) + .unwrap_or(false); + complete_openai_like_url(api_base, custom_endpoint, env_lookup) + } + + fn transform_request( + &self, + model: &str, + messages: Vec, + optional_params: Map, + ) -> Result { + let mut params = Map::from_iter( + optional_params + .into_iter() + .filter(|(key, _)| !CONFIG_PARAMS.contains(&key.as_str())), + ); + // Most OpenAI-compatible endpoints take `max_tokens`, not + // `max_completion_tokens`, so Python's `map_openai_params` renames it + // and lets it overwrite a `max_tokens` the caller also sent. + if let Some(limit) = params.remove("max_completion_tokens") { + params.insert("max_tokens".to_string(), limit); + } + let body = Map::from_iter( + [ + ("model".to_string(), json!(model)), + ("messages".to_string(), json!(messages)), + ] + .into_iter() + .chain(params), + ); + Ok(ProviderChatRequestData { + body: Value::Object(body), + stream_shape: Default::default(), + }) + } + + fn transform_response( + &self, + model: &str, + response: ProviderChatResponseData, + ) -> Result { + let mut body = response.body; + sanitize_usage(&mut body); + let body = body + .as_object() + .ok_or_else(|| Error::InvalidResponse("chat response is not an object".into()))?; + + let choices = body + .get("choices") + .and_then(Value::as_array) + .ok_or(Error::MissingField("choices"))? + .iter() + .enumerate() + .map(|(position, choice)| normalize_choice(position, choice)) + .collect::, _>>()?; + + let usage = body.get("usage").and_then(Value::as_object); + let field = |name: &str| { + usage + .and_then(|usage| usage.get(name)) + .and_then(Value::as_u64) + .unwrap_or(0) + }; + let details = usage.and_then(|usage| usage.get("prompt_tokens_details")); + + Ok(ChatCompletionsResponse { + created: body + .get("created") + .and_then(Value::as_u64) + .unwrap_or_else(unix_now), + model: body + .get("model") + .and_then(Value::as_str) + .unwrap_or(model) + .to_string(), + choices, + usage: litellm_types::utils::ChatCompletionsUsage { + prompt_tokens: field("prompt_tokens"), + completion_tokens: field("completion_tokens"), + total_tokens: field("total_tokens"), + prompt_tokens_details: litellm_types::utils::PromptTokensDetails { + cached_tokens: details + .and_then(|d| d.get("cached_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0), + cache_creation_tokens: details + .and_then(|d| d.get("cache_creation_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0), + text_tokens: details + .and_then(|d| d.get("text_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0), + }, + }, + }) + } + + /// `OpenAILikeBase._validate_environment`: a forwarded `authorization` is + /// the whole credential, and any other call authenticates with the + /// resolved key as a bearer. The key resolves to `""` when neither the + /// deployment nor `OPENAI_LIKE_API_KEY` sets one, because vllm-compatible + /// endpoints take no key; Python still sends `Bearer ` in that case. + fn validate_environment( + &self, + headers: Headers, + api_key: Option<&str>, + _model: &str, + _optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + if headers + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case("authorization")) + { + return Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Forwarded, + }); + } + let (_, key) = openai_compatible_provider_info(None, api_key, env_lookup); + Ok(ValidatedEnvironment { + headers, + auth: AuthScheme::Credential { + placement: CredentialPlacement::Bearer, + secret: SecretValue::new(key.unwrap_or_default()), + }, + }) + } + + fn config_params(&self) -> &'static [&'static str] { + CONFIG_PARAMS + } +} + +/// `OpenAILikeChatConfig._sanitize_usage_obj`: a provider that reports a null +/// `*_tokens` entry breaks OpenAI clients, so nulls become 0. Python scrubs +/// every top-level usage key ending in `_tokens`. +fn sanitize_usage(body: &mut Value) { + if let Some(usage) = body.get_mut("usage").and_then(Value::as_object_mut) { + for (key, value) in usage.iter_mut() { + if key.ends_with("_tokens") && value.is_null() { + *value = json!(0); + } + } + } +} + +fn normalize_choice(position: usize, choice: &Value) -> Result { + let message = choice + .get("message") + .and_then(Value::as_object) + .ok_or(Error::MissingField("message"))?; + if message + .get("tool_calls") + .and_then(Value::as_array) + .is_some_and(|calls| !calls.is_empty()) + { + // Python rewrites the lone tool call into content only under + // `json_mode`, a request flag `transform_response` cannot see, and the + // normalized type cannot carry tool calls at all. Declining is + // terminal at this point, but passing back an empty assistant turn + // would fabricate the reply. + return Err(Error::Unsupported("tool call response")); + } + if message.get("refusal").is_some_and(|value| !value.is_null()) { + return Err(Error::Unsupported("refusal response")); + } + let content = message.get("content"); + if content.is_some_and(|value| !value.is_null() && !value.is_string()) { + return Err(Error::Unsupported("non-text response content")); + } + Ok(ChatCompletionsChoice { + index: choice + .get("index") + .and_then(Value::as_u64) + .unwrap_or(position as u64), + message: ChatCompletionsChoiceMessage { + role: message + .get("role") + .and_then(Value::as_str) + .unwrap_or("assistant") + .to_string(), + content: content.and_then(Value::as_str).map(str::to_string), + }, + finish_reason: choice + .get("finish_reason") + .and_then(Value::as_str) + .unwrap_or("") + .to_string(), + }) +} diff --git a/litellm-rust/crates/llms/src/openai_like/common_utils.rs b/litellm-rust/crates/llms/src/openai_like/common_utils.rs new file mode 100644 index 00000000000..b855e6dc812 --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/common_utils.rs @@ -0,0 +1,58 @@ +//! Shared OpenAI-like credential and endpoint resolution, mirroring +//! `litellm/llms/openai_like/common_utils.py`. + +use crate::Error; + +/// `OpenAILikeChatConfig._get_openai_compatible_provider_info`: the deployment's +/// `api_base` wins over `OPENAI_LIKE_API_BASE`, and the deployment key over +/// `OPENAI_LIKE_API_KEY`, with an empty key allowed because vllm-compatible +/// endpoints do not require one. +pub fn openai_compatible_provider_info( + api_base: Option<&str>, + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> (Option, Option) { + let api_base = api_base + .map(str::to_string) + .or_else(|| env_lookup("OPENAI_LIKE_API_BASE")); + let api_key = api_key + .map(str::to_string) + .or_else(|| env_lookup("OPENAI_LIKE_API_KEY")) + .or(Some(String::new())); + (api_base, api_key) +} + +/// `OpenAILikeBase._validate_environment` requires an api base and, when the +/// caller gave no `custom_endpoint`, appends the route suffix. A caller-supplied +/// `custom_endpoint` base is used as is. +pub fn complete_openai_like_url( + api_base: Option<&str>, + custom_endpoint: bool, + env_lookup: &dyn Fn(&str) -> Option, +) -> Result { + let (api_base, _) = openai_compatible_provider_info(api_base, None, env_lookup); + let api_base = api_base.ok_or_else(|| { + Error::InvalidRequest( + "Missing API Base - A call is being made to LLM Provider but no api base is set either in the environment variables ({LLM_PROVIDER}_API_KEY) or via params" + .to_string(), + ) + })?; + if custom_endpoint { + return Ok(api_base); + } + Ok(format!( + "{}/chat/completions", + api_base.trim_end_matches('/') + )) +} + +/// The api key the call resolves to. `None` means neither the deployment nor the +/// environment supplied one, which is valid for endpoints that take no key. +pub fn resolve_openai_like_api_key( + api_key: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + openai_compatible_provider_info(None, api_key, env_lookup) + .1 + .filter(|key| !key.is_empty()) +} diff --git a/litellm-rust/crates/llms/src/openai_like/mod.rs b/litellm-rust/crates/llms/src/openai_like/mod.rs new file mode 100644 index 00000000000..df0cc73a5b0 --- /dev/null +++ b/litellm-rust/crates/llms/src/openai_like/mod.rs @@ -0,0 +1,2 @@ +pub mod chat; +pub mod common_utils; diff --git a/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs b/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs new file mode 100644 index 00000000000..8ee0654bddb --- /dev/null +++ b/litellm-rust/crates/llms/tests/openai_like_chat_transformation.rs @@ -0,0 +1,343 @@ +use litellm_llms::{ + Error, + base_llm::{ + auth::AuthScheme, + chat::transformation::{BaseConfig, ProviderChatResponseData, Unsupported}, + }, + openai_like::chat::transformation::OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG, +}; +use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse}; +use rstest::rstest; +use serde_json::{Map, Value, json}; + +fn messages(value: Value) -> Vec { + serde_json::from_value(value).expect("valid messages") +} + +fn params(value: Value) -> Map { + match value { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + } +} + +fn no_env(_: &str) -> Option { + None +} + +fn env_with<'a>(name: &'a str, value: &'a str) -> impl Fn(&str) -> Option + 'a { + move |key| (key == name).then(|| value.to_string()) +} + +fn transform(model: &str, msgs: Value, opts: Value) -> Value { + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .transform_request(model, messages(msgs), params(opts)) + .expect("request transforms") + .body +} + +fn transform_response(body: Value) -> Result { + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .transform_response("some-model", ProviderChatResponseData { body }) +} + +fn reason(msgs: Value, opts: Value) -> Option { + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), ¶ms(opts)) +} + +#[rstest] +fn builds_the_openai_shaped_body() { + let body = transform( + "my-model", + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ]), + json!({"temperature": 0.5, "max_tokens": 8}), + ); + assert_eq!(body["model"], json!("my-model")); + assert_eq!( + body["messages"], + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ]) + ); + assert_eq!(body["temperature"], json!(0.5)); + assert_eq!(body["max_tokens"], json!(8)); +} + +#[rstest] +fn renames_max_completion_tokens_to_max_tokens() { + // `OpenAILikeChatConfig.map_openai_params`: most OpenAI-compatible providers + // support `max_tokens`, not `max_completion_tokens`. + let body = transform( + "my-model", + json!([{"role": "user", "content": "hi"}]), + json!({"max_completion_tokens": 12}), + ); + assert_eq!(body["max_tokens"], json!(12)); + assert!(body.get("max_completion_tokens").is_none()); +} + +#[rstest] +fn max_completion_tokens_wins_when_both_limits_are_sent() { + // Python assigns `max_tokens = max_completion_tokens` after copying the + // params, so the renamed value outranks a caller-supplied `max_tokens`. + let body = transform( + "my-model", + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 8, "max_completion_tokens": 12}), + ); + assert_eq!(body["max_tokens"], json!(12)); + assert!(body.get("max_completion_tokens").is_none()); +} + +#[rstest] +fn call_configuration_never_enters_the_body() { + let body = transform( + "my-model", + json!([{"role": "user", "content": "hi"}]), + json!({"custom_endpoint": true, "extra_headers": {"x": "y"}, "max_retries": 2}), + ); + assert_eq!( + body.as_object().unwrap().keys().collect::>(), + vec!["model", "messages"] + ); +} + +#[rstest] +#[case::appends_the_chat_completions_suffix("https://vllm.example.com/v1", json!({}), "https://vllm.example.com/v1/chat/completions")] +#[case::trims_a_trailing_slash("https://vllm.example.com/v1/", json!({}), "https://vllm.example.com/v1/chat/completions")] +#[case::a_custom_endpoint_is_used_as_is("https://vllm.example.com/v1/chat/completions", json!({"custom_endpoint": true}), "https://vllm.example.com/v1/chat/completions")] +fn complete_url(#[case] api_base: &str, #[case] opts: Value, #[case] expected: &str) { + assert_eq!( + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .get_complete_url(Some(api_base), "my-model", ¶ms(opts), &no_env) + .expect("url resolves"), + expected + ); +} + +#[rstest] +fn api_base_falls_back_to_the_environment() { + assert_eq!( + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .get_complete_url( + None, + "my-model", + ¶ms(json!({})), + &env_with("OPENAI_LIKE_API_BASE", "https://env.example.com/v1"), + ) + .expect("url resolves"), + "https://env.example.com/v1/chat/completions" + ); +} + +#[rstest] +fn a_missing_api_base_is_an_error() { + assert!(matches!( + OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .get_complete_url(None, "my-model", ¶ms(json!({})), &no_env), + Err(Error::InvalidRequest(message)) if message.starts_with("Missing API Base") + )); +} + +#[rstest] +fn the_resolved_key_authenticates_as_a_bearer() { + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment( + vec![], + Some("sk-test"), + "my-model", + ¶ms(json!({})), + &no_env, + ) + .expect("validates"); + assert!(matches!( + validated.auth, + AuthScheme::Credential { + placement: litellm_auth::CredentialPlacement::Bearer, + .. + } + )); +} + +#[rstest] +fn the_key_falls_back_to_the_environment() { + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment( + vec![], + None, + "my-model", + ¶ms(json!({})), + &env_with("OPENAI_LIKE_API_KEY", "sk-env"), + ) + .expect("validates"); + let AuthScheme::Credential { secret, .. } = validated.auth else { + panic!("expected a bearer credential"); + }; + assert_eq!(secret.expose(), "sk-env"); +} + +#[rstest] +fn a_forwarded_authorization_is_the_whole_credential() { + // Python adds `Bearer ` only when the caller did not already send + // `Authorization`, so the forwarded header wins over the deployment key. + let headers = vec![("Authorization".to_string(), "Bearer caller".to_string())]; + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment( + headers, + Some("sk-test"), + "my-model", + ¶ms(json!({})), + &no_env, + ) + .expect("validates"); + assert!(matches!(validated.auth, AuthScheme::Forwarded)); +} + +#[rstest] +fn keyless_calls_still_validate_for_endpoints_that_take_no_key() { + // vllm-compatible endpoints require no api key; Python resolves `""` and + // sends `Bearer `, so validation must not fail on the missing key. + let validated = OPENAI_LIKE_CHAT_COMPLETIONS_CONFIG + .validate_environment(vec![], None, "my-model", ¶ms(json!({})), &no_env) + .expect("validates"); + let AuthScheme::Credential { secret, .. } = validated.auth else { + panic!("expected a bearer credential"); + }; + assert_eq!(secret.expose(), ""); +} + +#[rstest] +fn normalizes_an_openai_response() { + let response = transform_response(json!({ + "created": 1_700_000_000, + "model": "served-model-name", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "hello"}, + "finish_reason": "stop", + }], + "usage": {"prompt_tokens": 3, "completion_tokens": 5, "total_tokens": 8}, + })) + .expect("response normalizes"); + assert_eq!(response.created, 1_700_000_000); + assert_eq!(response.model, "served-model-name"); + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("hello") + ); + assert_eq!(response.choices[0].finish_reason, "stop"); + assert_eq!(response.usage.prompt_tokens, 3); + assert_eq!(response.usage.completion_tokens, 5); + assert_eq!(response.usage.total_tokens, 8); +} + +#[rstest] +fn null_token_fields_in_usage_become_zero() { + // `_sanitize_usage_obj`: providers that return null token values break + // OpenAI clients, so the response is scrubbed at the source. + let response = transform_response(json!({ + "model": "m", + "choices": [{"message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 3, "completion_tokens": null, "total_tokens": null}, + })) + .expect("response normalizes"); + assert_eq!(response.usage.completion_tokens, 0); + assert_eq!(response.usage.total_tokens, 0); + assert_eq!(response.usage.prompt_tokens, 3); +} + +#[rstest] +fn a_tool_call_response_declines_instead_of_dropping_the_calls() { + // The `json_mode` rewrite needs a request flag the route does not carry, so + // a tool-call answer falls back to Python rather than losing the calls. + assert_eq!( + transform_response(json!({ + "model": "m", + "choices": [{ + "message": { + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": "call_1", + "type": "function", + "function": {"name": "f", "arguments": "{}"}, + }], + }, + "finish_reason": "tool_calls", + }], + })), + Err(Error::Unsupported("tool call response")) + ); +} + +#[rstest] +fn a_refusal_declines_instead_of_returning_an_empty_reply() { + assert_eq!( + transform_response(json!({ + "model": "m", + "choices": [{ + "message": {"role": "assistant", "content": null, "refusal": "cannot help"}, + "finish_reason": "stop", + }], + })), + Err(Error::Unsupported("refusal response")) + ); +} + +#[rstest] +fn a_non_text_response_content_declines() { + assert_eq!( + transform_response(json!({ + "model": "m", + "choices": [{ + "message": {"role": "assistant", "content": [{"type": "text", "text": "hi"}]}, + "finish_reason": "stop", + }], + })), + Err(Error::Unsupported("non-text response content")) + ); +} + +#[rstest] +#[case::streaming(json!({"stream": true}), "streaming")] +#[case::unrecognized_param(json!({"some_provider_knob": 1}), "unrecognized request parameter")] +fn declines(#[case] opts: Value, #[case] expected: &'static str) { + assert_eq!( + reason(json!([{"role": "user", "content": "hi"}]), opts), + Some(Unsupported(expected)) + ); +} + +#[rstest] +fn accepts_standard_openai_params() { + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({ + "temperature": 0.2, + "top_p": 0.9, + "max_tokens": 16, + "response_format": {"type": "json_object"}, + "custom_endpoint": true, + }), + ), + None + ); +} + +#[rstest] +fn tool_parameters_decline_before_the_call() { + // A `tools` request would come back with tool calls this port cannot + // normalize, so it declines at the gate instead of after the call. + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({"tools": [{"type": "function", "function": {"name": "f"}}]}), + ), + Some(Unsupported("unrecognized request parameter")) + ); +} From de06c937670f8298f9b304e9fbaf4034cd4007fd Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 26 Sep 2026 18:13:13 -0700 Subject: [PATCH 14/51] feat(router): opt in to prompt-cache cost routing (#43232) * feat(router): opt in to prompt-cache cost routing * fix(router): address prompt-cache routing review * ci: include cache-routing regressions in coverage --- .circleci/scripts/unit_selection.sh | 1 + litellm/llms/anthropic/cache_aware_routing.py | 236 +++++++ .../proxy/common_utils/cache_aware_routing.py | 343 ++++++++++ .../common_utils/prompt_cache_prediction.py | 20 + .../common_utils/prompt_cache_pricing.py | 25 +- .../prompt_cache_prediction.py | 113 +--- .../complexity_router/README.md | 41 +- .../complexity_router/complexity_router.py | 54 +- .../complexity_router/config.py | 19 + litellm/types/utils.py | 1 + .../common_utils/test_cache_aware_routing.py | 588 ++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 20 +- 12 files changed, 1337 insertions(+), 124 deletions(-) create mode 100644 litellm/llms/anthropic/cache_aware_routing.py create mode 100644 litellm/proxy/common_utils/cache_aware_routing.py create mode 100644 litellm/proxy/common_utils/prompt_cache_prediction.py create mode 100644 tests/unit/proxy/common_utils/test_cache_aware_routing.py diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index 3f4f5620176..6510b3fd4b5 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -46,6 +46,7 @@ legacy_paths() { echo tests/unit/google_genai echo tests/unit/router_strategy echo tests/unit/router_utils + echo tests/unit/proxy/common_utils/test_cache_aware_routing.py echo tests/unit/enterprise/enterprise_callbacks/send_emails echo tests/unit/enterprise/proxy/test_afile_retrieve_returns_unified_id.py echo tests/unit/enterprise/proxy/test_batch_retrieve_input_file_id.py diff --git a/litellm/llms/anthropic/cache_aware_routing.py b/litellm/llms/anthropic/cache_aware_routing.py new file mode 100644 index 00000000000..e1a50781ace --- /dev/null +++ b/litellm/llms/anthropic/cache_aware_routing.py @@ -0,0 +1,236 @@ +from __future__ import annotations + +import time +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final +from urllib.parse import urlparse + +from pydantic import BaseModel, JsonValue, TypeAdapter + +import litellm +from litellm._internal_context import current_billing_time, pinned_billing_time +from litellm.caching.dual_cache import DualCache +from litellm.llms.anthropic.prompt_cache_prediction import ( + NativePredictionTarget, + PromptPrefix, + TokenCounter, + UnsupportedPredictionTarget, + cache_scope, + count_prompt_tokens, + parse_prompt, + resolve_prediction_target, + supported_prediction_headers, +) +from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens +from litellm.proxy.hooks.prompt_cache_prediction import lookup +from litellm.types.management_endpoints.prompt_cache_prediction import ( + CacheCostScenario, + CacheEvidence, + CachePredictionArm, + CacheTokenBuckets, +) +from litellm.types.router import Deployment +from litellm.utils import get_prompt_cache_min_tokens + +__all__: Final = ("AnthropicCacheRouting", "TokenCounter", "predict_arm") + +_JSON: Final = TypeAdapter(Mapping[str, JsonValue]) +_NATIVE_OPTIONS: Final = frozenset( + ( + "max_tokens", + "system", + "tools", + "tool_choice", + "thinking", + "output_config", + "cache_control", + "speed", + "service_tier", + "temperature", + "top_p", + "top_k", + "stop_sequences", + "stream", + ) +) + + +class _ModelLimits(BaseModel): + max_input_tokens: int | None = None + max_output_tokens: int | None = None + + +@dataclass(frozen=True, slots=True) +class AnthropicCacheRouting: + body: Mapping[str, JsonValue] + prefix: PromptPrefix + requested_output_limit: int + + @staticmethod + def request_body( + url: str, + headers: Mapping[str, str], + body: Mapping[str, JsonValue], + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, + ) -> Mapping[str, JsonValue] | None: + if not urlparse(url).path.endswith("/v1/messages") or not supported_prediction_headers(headers): + return None + return _JSON.validate_python( + MappingProxyType( + { + **body, + **MappingProxyType({key: request_kwargs[key] for key in _NATIVE_OPTIONS if key in request_kwargs}), + "messages": messages, + } + ) + ) + + @classmethod + def from_body(cls, body: Mapping[str, JsonValue]) -> AnthropicCacheRouting | None: + prefix: Final = parse_prompt(body) + limit: Final = body.get("max_tokens") + if prefix is None or not isinstance(limit, int) or isinstance(limit, bool) or limit <= 0: + return None + return cls(body, prefix, limit) + + @staticmethod + def supports(deployment: Deployment) -> bool: + return isinstance(resolve_prediction_target(deployment.litellm_params), NativePredictionTarget) + + async def is_warm(self, deployment: Deployment, caller: str, cache: DualCache, now: float) -> bool: + target: Final = resolve_prediction_target(deployment.litellm_params) + if not isinstance(target, NativePredictionTarget): + return False + scope: Final = cache_scope(caller, deployment.model_info.id or "", target.api_key, target.model) + observation: Final = await lookup(cache, scope, self.prefix, now=now) + return observation is not None and observation.expires_at > now + + @staticmethod + def fits(deployment: Deployment, input_tokens: int, output_tokens: int) -> bool: + target: Final = resolve_prediction_target(deployment.litellm_params) + if not isinstance(target, NativePredictionTarget): + return False + limits: Final = _ModelLimits.model_validate( + MappingProxyType( + { + **litellm.get_model_info(target.model, custom_llm_provider="anthropic"), + **deployment.model_info.model_dump(exclude_none=True), + } + ) + ) + return ( + limits.max_input_tokens is not None + and input_tokens + output_tokens <= limits.max_input_tokens + and limits.max_output_tokens is not None + and output_tokens <= limits.max_output_tokens + ) + + async def predict( + self, + deployment: Deployment, + caller: str, + cache: DualCache, + counter: TokenCounter, + now: float | None, + ) -> CachePredictionArm: + return await predict_arm(deployment, self.body, self.prefix, caller, cache, counter, now=now) + + @staticmethod + def cost(arm: CachePredictionArm, output_tokens: int) -> float | None: + return ( + price_cache_tokens(arm.model or "", arm.deployment_id, arm.estimate.tokens, output_tokens) + if arm.estimate is not None + else None + ) + + @staticmethod + async def count_tokens(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + return await count_prompt_tokens(model, api_key, body) + + +def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets: + return CacheTokenBuckets( + uncached_input_tokens=suffix_tokens, + cache_read_input_tokens=read_tokens, + cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0, + cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0, + ) + + +def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None: + cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens) + return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None + + +async def predict_arm( + deployment: Deployment, + body: Mapping[str, JsonValue], + prefix: PromptPrefix, + caller_key_hash: str, + cache: DualCache, + token_counter: TokenCounter, + now: float | None = None, +) -> CachePredictionArm: + deployment_id: Final = deployment.model_info.id or "" + params: Final = deployment.litellm_params + unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model) + if deployment.model_info.blocked: + return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"})) + target: Final = resolve_prediction_target(params) + if isinstance(target, UnsupportedPredictionTarget): + return unknown.model_copy(update=MappingProxyType({"reason": target.reason})) + model: Final = target.model + api_key: Final = target.api_key + total_count: Final = await token_counter(model, api_key, body) + prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body) + if total_count is None or prefix_count is None or total_count < prefix_count: + return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"})) + scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model) + checked_at: Final = time.time() if now is None else now + observation: Final = await lookup(cache, scope, prefix, now=checked_at) + exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint + cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count + if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable): + return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"})) + suffix: Final = total_count - cacheable + evidence: Final = ( + CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at) + if observation is not None + else None + ) + if cacheable < get_prompt_cache_min_tokens(params.model): + disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count)) + if disabled is None: + return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) + return CachePredictionArm( + deployment_id=deployment_id, + model=model, + cache_state="disabled", + reason="below_cache_minimum", + estimate=disabled, + cold=disabled, + warm=disabled, + token_count_source="anthropic_count_tokens", + ) + fresh: Final = observation is not None and observation.expires_at > checked_at + read: Final = observation.cached_tokens if fresh and observation is not None else 0 + with pinned_billing_time(current_billing_time()): + cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds)) + warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds)) + estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds)) + if cold is None or warm is None or estimate is None: + return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) + return CachePredictionArm( + deployment_id=deployment_id, + model=model, + cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown", + reason=None if fresh else "observation_expired" if observation else "no_compatible_observation", + estimate=estimate, + cold=cold, + warm=warm, + evidence=evidence, + token_count_source="anthropic_count_tokens", + ) diff --git a/litellm/proxy/common_utils/cache_aware_routing.py b/litellm/proxy/common_utils/cache_aware_routing.py new file mode 100644 index 00000000000..4ae2dce2440 --- /dev/null +++ b/litellm/proxy/common_utils/cache_aware_routing.py @@ -0,0 +1,343 @@ +from __future__ import annotations + +import asyncio +import time +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from itertools import chain +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError + +from litellm._logging import verbose_router_logger +from litellm.caching.dual_cache import DualCache +from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs +from litellm.llms.anthropic.cache_aware_routing import AnthropicCacheRouting, TokenCounter +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model +from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig +from litellm.types.router import Deployment, PreRoutingHookResponse +from litellm.types.utils import StandardLoggingRoutingDecision + +if TYPE_CHECKING: + from litellm.router import Router + +_MESSAGES: Final = TypeAdapter(list[Mapping[str, object]]) +_MAPPING: Final = TypeAdapter(Mapping[str, object]) +_DEPLOYMENTS: Final[TypeAdapter[tuple[Deployment, ...] | Deployment]] = TypeAdapter(tuple[Deployment, ...] | Deployment) +_MARKER_OPTIONS: Final = frozenset( + ("model", "complexity_router_config", "rpm", "tpm", "tags", "timeout", "stream_timeout", "num_retries") +) +_CLASSIFIED_CAUSES: Final = frozenset( + { + "heuristic_scorer", + "heuristic_v2", + "reasoning_override", + "llm_classifier", + "llm_v2_classifier", + "jev_classifier", + "capability_classifier", + "heuristic_first_short_circuit", + "hybrid_short_circuit", + "classifier_plugin", + } +) + + +class _ProxyRequest(BaseModel): + model_config = ConfigDict(strict=True) + url: str + body: Mapping[str, JsonValue] + headers: Mapping[str, str] + + +class _CallerSettings(BaseModel): + config: Mapping[str, object] | None = None + + +@dataclass(frozen=True, slots=True) +class CacheAwareChoice: + model: str + tier: str + deployment_id: str + original_cost: float + estimated_cost: float + + +@dataclass(frozen=True, slots=True) +class _Candidate: + model: str + tier: str + deployment: Deployment + + +def eligible_models( + config: ComplexityRouterConfig, decision: StandardLoggingRoutingDecision +) -> tuple[tuple[str, str], ...]: + tier: Final = decision.get("tier") + order: Final = config.tier_names() + tier_entries: Final = chain.from_iterable(config.tier_model_configs.values()) + if ( + tier is None + or tier not in order + or decision.get("cause") not in _CLASSIFIED_CAUSES + or config.has_custom_tiers + or config.plugins + or config.adaptive + or config.session_affinity + or config.classification_mode != "every_request" + or any(entry.litellm_params for entry in tier_entries) + or any(not isinstance(model, str) for model in config.tiers.values()) + ): + return () + floor: Final = order.index(tier) + eligible: Final = tuple( + (name, model) for name, model in config.tiers.items() if isinstance(model, str) and name in order[floor:] + ) + return tuple(entry for index, entry in enumerate(eligible) if entry[1] not in tuple(m for _, m in eligible[:index])) + + +def _candidate(router: Router, tier: str, model: str, request_kwargs: Mapping[str, object]) -> _Candidate | None: + deployments: Final = router.deployments_for_request(model, request_kwargs) + if len(deployments) != 1: + return None + deployment: Final = Deployment.model_validate(deployments[0]) + if deployment.model_info.blocked or not deployment.model_info.id or not AnthropicCacheRouting.supports(deployment): + return None + return _Candidate(model, tier, deployment) + + +async def _available( + candidate: _Candidate, + router: Router, + caller: UserAPIKeyAuth, + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, +) -> bool: + try: + await can_key_call_resolved_model( + model=candidate.model, llm_model_list=router.get_model_list(), valid_token=caller, llm_router=router + ) + healthy: Final = _DEPLOYMENTS.validate_python( + await router.async_get_healthy_deployments( # pyright: ignore[reportUnknownMemberType] # legacy router results are validated at this boundary + model=candidate.model, + messages=_MESSAGES.validate_python(messages) if messages else None, # pyright: ignore[reportArgumentType] # router annotations predate structured native messages + request_kwargs=dict(request_kwargs), # mutable-ok: Router's filtering API accepts a request dictionary + ) + ) + except Exception: # noqa: BLE001 # an unavailable optional candidate must not fail the originally selected route + return False + available: Final = (healthy,) if isinstance(healthy, Deployment) else healthy + return any(entry.model_info.id == candidate.deployment.model_info.id for entry in available) + + +def supported_router_marker(router: Router, alias: str, request_kwargs: Mapping[str, object]) -> bool: + markers: Final = tuple( + Deployment.model_validate(entry) for entry in router.deployments_for_request(alias, request_kwargs) + ) + return bool(markers) and all( + marker.litellm_params.model == "auto_router/complexity_router" + and not frozenset(marker.litellm_params.model_dump(exclude_defaults=True, exclude_none=True)) - _MARKER_OPTIONS + for marker in markers + ) + + +async def select_cached_model( + *, + router: Router, + config: ComplexityRouterConfig, + params_for_model: Callable[[str, str], Mapping[str, object]], + response: PreRoutingHookResponse, + body: Mapping[str, JsonValue], + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, + caller: UserAPIKeyAuth, + cache: DualCache, + counter_for_model: Callable[[str], TokenCounter], + now: float | None = None, +) -> CacheAwareChoice | None: + checked_at: Final = time.time() if now is None else now + decision: Final = response.routing_decision + provider: Final = AnthropicCacheRouting.from_body(body) + if not config.cache_aware_routing or decision is None or provider is None or not caller.api_key: + return None + names: Final = eligible_models(config, decision) + if not names or response.model not in tuple(model for _, model in names): + return None + candidates: Final = tuple( + candidate for tier, model in names if (candidate := _candidate(router, tier, model, request_kwargs)) is not None + ) + original: Final = next((candidate for candidate in candidates if candidate.model == response.model), None) + if original is None: + return None + alternatives: Final = tuple(candidate for candidate in candidates if candidate.model != original.model) + warm_flags: Final = await asyncio.gather( + *(provider.is_warm(candidate.deployment, caller.api_key, cache, checked_at) for candidate in alternatives) + ) + warm: Final = tuple(candidate for candidate, fresh in zip(alternatives, warm_flags) if fresh) + if not warm: + return None + considered: Final = (original, *warm) + availability: Final = await asyncio.gather( + *(_available(candidate, router, caller, request_kwargs, messages) for candidate in considered) + ) + authorized: Final = tuple(candidate for candidate, available in zip(warm, availability[1:]) if available) + if not availability[0] or not authorized: + return None + compared: Final = (original, *authorized) + output_limits: Final = tuple( + params_for_model(candidate.tier, candidate.model).get("max_tokens", provider.requested_output_limit) + for candidate in compared + ) + if any(not isinstance(limit, int) or isinstance(limit, bool) or limit <= 0 for limit in output_limits): + return None + limits: Final = tuple(limit for limit in output_limits if isinstance(limit, int)) + arms: Final = await asyncio.gather( + *( + provider.predict( + candidate.deployment, + caller.api_key, + cache, + counter_for_model(candidate.model), + now=now, + ) + for candidate in compared + ) + ) + costs: Final = tuple( + provider.cost(arm, min(config.cache_aware_routing_output_tokens, limit)) for arm, limit in zip(arms, limits) + ) + original_cost: Final = costs[0] + if original_cost is None: + return None + finished_at: Final = time.time() if now is None else now + qualifying: Final = tuple( + CacheAwareChoice(candidate.model, candidate.tier, arm.deployment_id, original_cost, cost) + for candidate, arm, cost, limit in zip(authorized, arms[1:], costs[1:], limits[1:]) + if cost is not None + and cost < original_cost + and arm.cache_state in ("warm", "partial") + and arm.evidence is not None + and arm.evidence.expires_at > finished_at + and arm.estimate is not None + and provider.fits(candidate.deployment, arm.estimate.tokens.total_tokens, limit) + ) + return min(qualifying, key=lambda choice: choice.estimated_cost, default=None) + + +async def _choose_cached_model( + *, + router: Router, + config: ComplexityRouterConfig, + params_for_model: Callable[[str, str], Mapping[str, object]], + response: PreRoutingHookResponse | None, + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, +) -> CacheAwareChoice | None: + if not config.cache_aware_routing or response is None or response.routing_decision is None: + return None + if not eligible_models(config, response.routing_decision): + return None + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner + ) + from litellm.router_strategy.complexity_router.context_compaction import compaction_pending + + if ( + proxy_server.llm_router is not router + or router.routing_plugins + or has_request_transforms() + or compaction_pending(request_kwargs) + or not supported_router_marker(router, response.routing_decision.get("router_model_name") or "", request_kwargs) + ): + return None + metadata: Final = _MAPPING.validate_python( + request_kwargs.get(get_metadata_variable_name_from_kwargs(request_kwargs)) or MappingProxyType({}) + ) + caller: Final = metadata.get("user_api_key_auth") + if not isinstance(caller, UserAPIKeyAuth): + return None + settings: Final = _CallerSettings.model_validate(caller, from_attributes=True) + if settings.config: + return None + try: + incoming: Final = _ProxyRequest.model_validate(request_kwargs.get("proxy_server_request")) + except ValidationError: + return None + if any( + request_kwargs.get(key) + for key in ( + "guardrails", + "cache_control_injection_points", + "api_key", + "api_base", + "extra_headers", + "prompt_id", + "mock_response", + "model_info", + "custom_llm_provider", + ) + ): + return None + body: Final = AnthropicCacheRouting.request_body( + incoming.url, incoming.headers, incoming.body, request_kwargs, messages + ) + if body is None: + return None + limiter: Final = proxy_server.proxy_logging_obj.get_proxy_hook("parallel_request_limiter") + if not isinstance(limiter, _PROXY_MaxParallelRequestsHandler_v3): + return None + + def counter_for_model(model_name: str) -> TokenCounter: + async def count(model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + try: + async with limiter.request_capacity(caller, model_name, request_data=request_kwargs): + return await AnthropicCacheRouting.count_tokens(model, api_key, body) + except Exception: # noqa: BLE001 # an optional prediction denied capacity is an unavailable estimate + return None + + return count + + return await select_cached_model( + router=router, + config=config, + params_for_model=params_for_model, + response=response, + body=body, + request_kwargs=request_kwargs, + messages=messages, + caller=caller, + cache=proxy_server.proxy_logging_obj.internal_usage_cache.dual_cache, + counter_for_model=counter_for_model, + ) + + +async def choose_cached_model( + *, + router: Router, + config: ComplexityRouterConfig, + params_for_model: Callable[[str, str], Mapping[str, object]], + response: PreRoutingHookResponse | None, + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, +) -> CacheAwareChoice | None: + if not config.cache_aware_routing: + return None + try: + return await asyncio.wait_for( + _choose_cached_model( + router=router, + config=config, + params_for_model=params_for_model, + response=response, + request_kwargs=request_kwargs, + messages=messages, + ), + timeout=config.cache_aware_routing_timeout_ms / 1000, + ) + except Exception: # noqa: BLE001 # cache prediction is optional and must preserve normal routing on failure + verbose_router_logger.debug("Cache-aware routing unavailable; keeping the classified model") + return None diff --git a/litellm/proxy/common_utils/prompt_cache_prediction.py b/litellm/proxy/common_utils/prompt_cache_prediction.py new file mode 100644 index 00000000000..75a9bd40840 --- /dev/null +++ b/litellm/proxy/common_utils/prompt_cache_prediction.py @@ -0,0 +1,20 @@ +from typing import Final + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.anthropic.cache_aware_routing import predict_arm + +__all__: Final = ("has_request_transforms", "predict_arm") + + +def has_request_transforms() -> bool: + from litellm.proxy.hooks import PROXY_HOOKS + + builtins: Final = frozenset(PROXY_HOOKS.values()) + hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook") + callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger) + return any( + type(callback) not in builtins + and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks) + for callback in callbacks + ) diff --git a/litellm/proxy/common_utils/prompt_cache_pricing.py b/litellm/proxy/common_utils/prompt_cache_pricing.py index ff070853b46..1ecc3ac44fe 100644 --- a/litellm/proxy/common_utils/prompt_cache_pricing.py +++ b/litellm/proxy/common_utils/prompt_cache_pricing.py @@ -1,4 +1,5 @@ from collections.abc import Mapping +from datetime import datetime, timezone from math import isfinite from typing import Final @@ -20,12 +21,13 @@ def _valid_price(value: object) -> bool: return isinstance(value, (int, float)) and not isinstance(value, bool) and isfinite(value) and value >= 0 -def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets) -> bool: +def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets, completion_tokens: int = 0) -> bool: required: Final = ( ("input_cost_per_token", True), ("cache_read_input_token_cost", tokens.cache_read_input_tokens > 0), ("cache_creation_input_token_cost", tokens.cache_creation_5m_input_tokens > 0), ("cache_creation_input_token_cost_above_1hr", tokens.cache_creation_1h_input_tokens > 0), + ("output_cost_per_token", completion_tokens > 0), ) if any(needed and not _valid_price(prices.get(key)) for key, needed in required): return False @@ -36,7 +38,9 @@ def _has_required_prices(prices: Mapping[str, object], tokens: CacheTokenBuckets ) -def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> float | None: +def price_cache_tokens( + model: str, deployment_id: str, tokens: CacheTokenBuckets, completion_tokens: int = 0 +) -> float | None: try: selected_model: Final = _select_model_name_for_cost_calc( model=model, @@ -53,12 +57,12 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets if price_entry is None: return None prices: Final = _PRICE_ENTRY.validate_python(price_entry) - if not _has_required_prices(prices, tokens): + if completion_tokens < 0 or not _has_required_prices(prices, tokens, completion_tokens): return None usage: Final = Usage( prompt_tokens=tokens.total_tokens, - completion_tokens=0, - total_tokens=tokens.total_tokens, + completion_tokens=completion_tokens, + total_tokens=tokens.total_tokens + completion_tokens, prompt_tokens_details=PromptTokensDetailsWrapper( cached_tokens=tokens.cache_read_input_tokens, cache_creation_tokens=tokens.cache_creation_5m_input_tokens + tokens.cache_creation_1h_input_tokens, @@ -73,7 +77,7 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets messages=[], # mutable-ok: Logging requires a list stream=False, call_type="completion", - start_time=None, + start_time=datetime.now(timezone.utc), litellm_call_id="prompt-cache-prediction", function_id="prompt-cache-prediction", ) @@ -85,7 +89,12 @@ def price_cache_tokens(model: str, deployment_id: str, tokens: CacheTokenBuckets router_model_id=deployment_id, litellm_logging_obj=logging_obj, ) - cost: Final = logging_obj.cost_breakdown.get("input_cost") if logging_obj.cost_breakdown is not None else None - return cost if cost is not None and _valid_price(cost) else None + breakdown: Final = logging_obj.cost_breakdown + input_cost: Final = breakdown.get("input_cost") if breakdown is not None else None + output_cost: Final = breakdown.get("output_cost") if breakdown is not None else None + if input_cost is None or output_cost is None: + return None + cost: Final = input_cost + output_cost + return cost if _valid_price(cost) else None except Exception: # noqa: BLE001 # the shared pricing owners raise plain Exception for unpriceable models return None diff --git a/litellm/proxy/management_endpoints/prompt_cache_prediction.py b/litellm/proxy/management_endpoints/prompt_cache_prediction.py index 56e844214d6..757880980c9 100644 --- a/litellm/proxy/management_endpoints/prompt_cache_prediction.py +++ b/litellm/proxy/management_endpoints/prompt_cache_prediction.py @@ -1,4 +1,3 @@ -import time from collections.abc import Mapping from types import MappingProxyType from typing import Annotated, Final @@ -6,18 +5,10 @@ from typing import Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Request from pydantic import BaseModel, JsonValue, TypeAdapter -import litellm -from litellm._internal_context import current_billing_time, pinned_billing_time -from litellm.caching.caching import DualCache -from litellm.integrations.custom_logger import CustomLogger from litellm.llms.anthropic.prompt_cache_prediction import ( - PromptPrefix, TokenCounter, - UnsupportedPredictionTarget, - cache_scope, count_prompt_tokens, parse_prompt, - resolve_prediction_target, supported_prediction_headers, ) from litellm.proxy._types import UserAPIKeyAuth @@ -27,22 +18,16 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # canonical parsed-body owner; validate its legacy result at the endpoint boundary ) -from litellm.proxy.common_utils.prompt_cache_pricing import price_cache_tokens +from litellm.proxy.common_utils.prompt_cache_prediction import has_request_transforms, predict_arm from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # use the configured proxy limiter's shared capacity owner ) -from litellm.proxy.hooks.prompt_cache_prediction import lookup from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.types.management_endpoints.prompt_cache_prediction import ( - CacheCostScenario, - CacheEvidence, CachePredictionArm, CachePredictionRequest, CachePredictionResponse, - CacheTokenBuckets, ) -from litellm.types.router import Deployment -from litellm.utils import get_prompt_cache_min_tokens router: Final = APIRouter() _REQUEST_DATA: Final = TypeAdapter(Mapping[str, object]) @@ -52,33 +37,6 @@ class _CallerSettings(BaseModel): config: Mapping[str, object] | None = None -def has_request_transforms() -> bool: - from litellm.proxy.hooks import PROXY_HOOKS - - builtins: Final = frozenset(PROXY_HOOKS.values()) - hooks: Final = ("async_pre_call_hook", "async_pre_request_hook", "async_pre_call_deployment_hook") - callbacks: Final = litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=CustomLogger) - return any( - type(callback) not in builtins - and any(getattr(type(callback), hook) is not getattr(CustomLogger, hook) for hook in hooks) - for callback in callbacks - ) - - -def _buckets(prefix_tokens: int, suffix_tokens: int, read_tokens: int, ttl_seconds: int) -> CacheTokenBuckets: - return CacheTokenBuckets( - uncached_input_tokens=suffix_tokens, - cache_read_input_tokens=read_tokens, - cache_creation_5m_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 300 else 0, - cache_creation_1h_input_tokens=prefix_tokens - read_tokens if ttl_seconds == 3600 else 0, - ) - - -def _scenario(model: str, deployment_id: str, tokens: CacheTokenBuckets) -> CacheCostScenario | None: - cost: Final = price_cache_tokens(model=model, deployment_id=deployment_id, tokens=tokens) - return CacheCostScenario(tokens=tokens, input_cost=cost) if cost is not None else None - - def _capacity_counter( limiter: _PROXY_MaxParallelRequestsHandler_v3, caller: UserAPIKeyAuth, @@ -103,75 +61,6 @@ def _capacity_request_data( return MappingProxyType(data) -async def predict_arm( - deployment: Deployment, - body: Mapping[str, JsonValue], - prefix: PromptPrefix, - caller_key_hash: str, - cache: DualCache, - token_counter: TokenCounter, -) -> CachePredictionArm: - deployment_id: Final = deployment.model_info.id or "" - params: Final = deployment.litellm_params - unknown: Final = CachePredictionArm(deployment_id=deployment_id, model=params.model) - if deployment.model_info.blocked: - return unknown.model_copy(update=MappingProxyType({"reason": "unsupported_deployment_configuration"})) - target: Final = resolve_prediction_target(params) - if isinstance(target, UnsupportedPredictionTarget): - return unknown.model_copy(update=MappingProxyType({"reason": target.reason})) - model: Final = target.model - api_key: Final = target.api_key - total_count: Final = await token_counter(model, api_key, body) - prefix_count: Final = await token_counter(model, api_key, prefix.prefix_body) - if total_count is None or prefix_count is None or total_count < prefix_count: - return unknown.model_copy(update=MappingProxyType({"reason": "token_count_unavailable"})) - scope: Final = cache_scope(caller_key_hash, deployment_id, api_key, model) - observation: Final = await lookup(cache, scope, prefix) - exact: Final = observation is not None and observation.fingerprint == prefix.fingerprint - cacheable: Final = observation.cached_tokens if exact and observation is not None else prefix_count - if cacheable > total_count or (observation is not None and observation.cached_tokens > cacheable): - return unknown.model_copy(update=MappingProxyType({"reason": "inconsistent_prefix_token_count"})) - suffix: Final = total_count - cacheable - evidence: Final = ( - CacheEvidence(observed_at=observation.observed_at, expires_at=observation.expires_at) - if observation is not None - else None - ) - if cacheable < get_prompt_cache_min_tokens(params.model): - disabled: Final = _scenario(model, deployment_id, CacheTokenBuckets(uncached_input_tokens=total_count)) - if disabled is None: - return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) - return CachePredictionArm( - deployment_id=deployment_id, - model=model, - cache_state="disabled", - reason="below_cache_minimum", - estimate=disabled, - cold=disabled, - warm=disabled, - token_count_source="anthropic_count_tokens", - ) - fresh: Final = observation is not None and observation.expires_at > time.time() - read: Final = observation.cached_tokens if fresh and observation is not None else 0 - with pinned_billing_time(current_billing_time()): - cold: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, 0, prefix.ttl_seconds)) - warm: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, cacheable, prefix.ttl_seconds)) - estimate: Final = _scenario(model, deployment_id, _buckets(cacheable, suffix, read, prefix.ttl_seconds)) - if cold is None or warm is None or estimate is None: - return unknown.model_copy(update=MappingProxyType({"reason": "pricing_unavailable"})) - return CachePredictionArm( - deployment_id=deployment_id, - model=model, - cache_state="warm" if fresh and exact else "partial" if fresh else "stale" if observation else "unknown", - reason=None if fresh else "observation_expired" if observation else "no_compatible_observation", - estimate=estimate, - cold=cold, - warm=warm, - evidence=evidence, - token_count_source="anthropic_count_tokens", - ) - - @router.post( "/cost/predict-cache", tags=["Cost Tracking"], # mutable-ok: FastAPI requires a list for OpenAPI tags diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index f023d5001d9..f55362b6c41 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -1,6 +1,6 @@ # Complexity Router -A rule-based routing strategy that classifies requests by complexity and routes them to appropriate models - with zero API calls and sub-millisecond latency. +A routing strategy that classifies requests by complexity and routes them to appropriate models. The default rule-based classifier scores requests locally. Optional classifiers and cache-aware routing can make provider calls ## Overview @@ -68,6 +68,45 @@ still resolve to a deployment in `model_list`; this configuration does not creat - abc ``` +### Opt in to prompt-cache costs + +Set `cache_aware_routing: true` to consider observed prompt-cache savings after classification. This is disabled by default. A warm model in the same or a higher tier can replace the classified model when its estimated input and output cost is strictly lower. Cache savings never lower the required tier + +```yaml +model_list: + - model_name: smart-router + litellm_params: + model: auto_router/complexity_router + complexity_router_config: + cache_aware_routing: true + cache_aware_routing_output_tokens: 1024 + cache_aware_routing_timeout_ms: 2000 + context_compaction: false + tiers: + SIMPLE: haiku + COMPLEX: sonnet + - model_name: haiku + litellm_params: + model: anthropic/claude-haiku-4-5 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: sonnet + litellm_params: + model: anthropic/claude-sonnet-5 + api_key: os.environ/ANTHROPIC_API_KEY +``` + +This first version supports the proxy's native `POST /v1/messages` endpoint with Anthropic, text and client tools, and one explicit message-content `cache_control` breakpoint. Each tier must name one model group with one deployment. The default v3 rate limiter must be enabled. It uses the same observations and token counting as `/cost/predict-cache`; it does not prewarm caches or enable provider caching on the application's behalf + +The proxy must have observed a successful cache read or write for the candidate's matching prefix, under the same caller key, deployment, provider key and model. A fresh observation allows a cache discount; missing or expired evidence does not. Provider eviction can still turn an expected hit into a miss + +The comparison includes uncached input, cache writes at the requested TTL, cache reads and expected output tokens. Set `cache_aware_routing_output_tokens` to your workload's expected response length; it defaults to 1024 and is capped separately by each model's effective output limit. With `max_tokens_from_tier_model: true` (the default), this is the model's known output ceiling; when disabled or unknown, the caller's `max_tokens` applies. The full effective output limit, together with the counted input, must fit the candidate's known limits. Custom deployment prices are respected + +Prediction makes up to two token-count requests per compared model. These use rate and concurrency capacity and add latency. The default total timeout is two seconds; timeout, missing counts or prices, and prediction failures preserve the classified route. No provider count requests run when there is no warm eligible alternative + +Session affinity, user-turn classification, adaptive routing, routing plugins, custom tier ladders, tier pools and per-tier parameter overrides keep their existing behavior without a cache adjustment. The same applies to unsupported providers or prompt shapes, beta headers, custom provider endpoints, request transforms, and pending context compaction. Disable context compaction as in the example so it cannot rewrite the predicted prompt. Alias markers should contain only routing configuration and rate, timeout or tag settings + +When cache costs change the model, the routing decision reports `cause: prompt_cache_cost`. Its signals include the original model, classification cause and both estimated costs + ### Capability forecasting Set `classifier_type: capability` to use diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 9df6306436b..76ee977bf28 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -4390,10 +4390,13 @@ class ComplexityRouter(CustomLogger): resolved_messages=resolved_messages, context_fit=context_fit, ) + cache_adjusted_response: Final = await self._apply_prompt_cache_routing( + routed_response, messages, request_kwargs, context_fit + ) response: Final = ( await self._gate_response_health( await self._gate_response_modality( - routed_response, messages, resolved_messages, request_kwargs, context_fit + cache_adjusted_response, messages, resolved_messages, request_kwargs, context_fit ), messages, input, @@ -4401,7 +4404,7 @@ class ComplexityRouter(CustomLogger): request_kwargs, context_fit, ) - if routed_response is not None + if cache_adjusted_response is not None else None ) # Sentinel presence, not the plan_mode cause, gates the pin write: a plan-mode turn @@ -4425,6 +4428,53 @@ class ComplexityRouter(CustomLogger): ) return self._with_session_deployment_affinity(response) + async def _apply_prompt_cache_routing( + self, + response: PreRoutingHookResponse | None, + messages: Sequence[Mapping[str, object]] | None, + request_kwargs: Mapping[str, object], + context_fit: _RequestContextFit, + ) -> PreRoutingHookResponse | None: + if not self.config.cache_aware_routing or response is None or response.routing_decision is None: + return response + from litellm.proxy.common_utils.cache_aware_routing import choose_cached_model + + choice: Final = await choose_cached_model( + router=self.litellm_router_instance, + config=self.config, + params_for_model=self._litellm_params_for_model, + response=response, + request_kwargs=request_kwargs, + messages=messages, + ) + if choice is None or not context_fit.accepts(choice.model): + return response + params: Final = self._litellm_params_for_model(choice.tier, choice.model) + decision: Final[StandardLoggingRoutingDecision] = { + **response.routing_decision, + "routed_model": choice.model, + "cause": "prompt_cache_cost", + "tier": choice.tier, + "tier_label": (self.config.tier_labels or {}).get(choice.tier, choice.tier), + "tier_litellm_params": params, + "signals": ( + *(response.routing_decision.get("signals") or ()), + f"cache-aware:classified-model={response.model}", + f"cache-aware:classification-cause={response.routing_decision.get('cause')}", + f"cache-aware:estimated-cost={choice.estimated_cost:.8f};original-cost={choice.original_cost:.8f}", + ), + } + verbose_router_logger.info( + "ComplexityRouter: cache-aware choice model=%s original=%s estimated_cost=%s original_cost=%s", + choice.model, + response.model, + choice.estimated_cost, + choice.original_cost, + ) + return response.model_copy( + update=MappingProxyType({"model": choice.model, "litellm_params": params, "routing_decision": decision}) + ) + async def _classify_and_route( self, model: str, diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index dbc70631298..e0427f89fe3 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -1422,6 +1422,25 @@ class ComplexityRouterConfig(BaseModel): ), ) + cache_aware_routing: bool = Field( + default=False, + description=( + "Opt in to comparing prompt-cache costs after classification. On supported native Anthropic proxy requests, " + "an already warm model in the same or a higher tier may replace the classified model when its estimated " + "input and output cost is lower. Unsupported requests and unavailable estimates keep ordinary routing." + ), + ) + cache_aware_routing_output_tokens: int = Field( + default=1024, + ge=0, + description="Expected output tokens used in cache-aware cost comparisons; capped by each model's effective output limit.", + ) + cache_aware_routing_timeout_ms: int = Field( + default=2000, + gt=0, + description="Total time budget for cache-aware predictions; expiry preserves the original routing decision.", + ) + # Session affinity: pin the first turn's routed model for the rest of the session session_affinity: bool = Field( default=False, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index f8b57139b37..4bda1dd53ce 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2970,6 +2970,7 @@ class StandardLoggingRoutingDecisionTierBoundaries(TypedDict): RoutingDecisionCause = Literal[ + "prompt_cache_cost", "heuristic_scorer", "heuristic_v2", # The scorer found 2+ reasoning markers and forced REASONING regardless of score. diff --git a/tests/unit/proxy/common_utils/test_cache_aware_routing.py b/tests/unit/proxy/common_utils/test_cache_aware_routing.py new file mode 100644 index 00000000000..000d6773d04 --- /dev/null +++ b/tests/unit/proxy/common_utils/test_cache_aware_routing.py @@ -0,0 +1,588 @@ +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from typing import Final + +import pytest +from pydantic import JsonValue + +from litellm import Router +from litellm.caching.dual_cache import DualCache +from litellm.llms.anthropic.prompt_cache_prediction import TokenCounter, cache_scope, parse_prompt +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.cache_aware_routing import ( + CacheAwareChoice, + choose_cached_model, + eligible_models, + select_cached_model, +) +from litellm.proxy.hooks.prompt_cache_prediction import CacheObservation, _cache_key +from litellm.router_strategy.complexity_router.complexity_router import ComplexityRouter +from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig +from litellm.types.router import PreRoutingHookResponse + +_CALLER: Final = "test-cache-aware-caller" +_PROVIDER_KEY: Final = "test-cache-aware-provider" +_NOW: Final = 1000.0 + + +@dataclass(frozen=True, slots=True) +class _Counts: + total: int | None = 51000 + prefix: int | None = 50000 + + async def __call__(self, model: str, api_key: str, body: Mapping[str, JsonValue]) -> int | None: + return self.total if "max_tokens" in body else self.prefix + + +def _counter_for_model(model: str) -> TokenCounter: + return _Counts() + + +def _forbidden_counter(model: str) -> TokenCounter: + raise AssertionError("No provider counts should run without a warm eligible alternative") + + +def _body(text: str = "Stable cached context") -> dict[str, JsonValue]: + return { + "max_tokens": 20000, + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": text, "cache_control": {"type": "ephemeral"}}, + {"type": "text", "text": "What is 2 + 2?"}, + ], + } + ], + } + + +def _router( + strong_output_rate: float = 0.000015, *, free: bool = False, cheap_limit: int = 30000, strong_limit: int = 30000 +) -> Router: + return Router( + model_list=[ + { + "model_name": "cheap", + "litellm_params": { + "model": "anthropic/claude-haiku-4-5", + "api_key": _PROVIDER_KEY, + "input_cost_per_token": 0 if free else 0.000001, + "output_cost_per_token": 0 if free else 0.000004, + "cache_read_input_token_cost": 0 if free else 0.0000001, + "cache_creation_input_token_cost": 0 if free else 0.00000125, + }, + "model_info": {"id": "test-cache-cheap", "max_input_tokens": 100000, "max_output_tokens": cheap_limit}, + }, + { + "model_name": "strong", + "litellm_params": { + "model": "anthropic/claude-sonnet-5", + "api_key": _PROVIDER_KEY, + "input_cost_per_token": 0 if free else 0.000003, + "output_cost_per_token": 0 if free else strong_output_rate, + "cache_read_input_token_cost": 0 if free else 0.0000003, + "cache_creation_input_token_cost": 0 if free else 0.00000375, + }, + "model_info": { + "id": "test-cache-strong", + "max_input_tokens": 100000, + "max_output_tokens": strong_limit, + }, + }, + ] + ) + + +def _config(**overrides: object) -> ComplexityRouterConfig: + return ComplexityRouterConfig.model_validate( + { + "tiers": {"SIMPLE": "cheap", "COMPLEX": "strong"}, + "cache_aware_routing": True, + **overrides, + } + ) + + +def _response(tier: str = "SIMPLE", model: str = "cheap") -> PreRoutingHookResponse: + return PreRoutingHookResponse( + model=model, + messages=None, + routing_decision={ + "router_model_name": "smart", + "router_type": "complexity", + "routed_model": model, + "tier": tier, + "cause": "heuristic_scorer", + }, + ) + + +async def _observed(cache: DualCache, *, caller: str = _CALLER, expires_at: float = 1290.0) -> None: + prefix: Final = parse_prompt(_body()) + assert prefix is not None + scope: Final = cache_scope(caller, "test-cache-strong", _PROVIDER_KEY, "claude-sonnet-5") + observation: Final = CacheObservation( + fingerprint=prefix.fingerprint, + cached_tokens=50000, + observed_at=990.0, + expires_at=expires_at, + ) + await cache.async_set_cache(_cache_key(scope, prefix.fingerprint), observation.model_dump_json(), ttl=3600) + + +async def _select( + *, + router: Router, + config: ComplexityRouterConfig, + response: PreRoutingHookResponse, + body: Mapping[str, JsonValue], + request_kwargs: Mapping[str, object], + messages: Sequence[Mapping[str, object]] | None, + caller: UserAPIKeyAuth, + cache: DualCache, + counter_for_model: Callable[[str], TokenCounter], + now: float, +) -> CacheAwareChoice | None: + complexity: Final = ComplexityRouter("smart", router, config.model_dump()) + return await select_cached_model( + router=router, + config=config, + params_for_model=complexity._litellm_params_for_model, + response=response, + body=body, + request_kwargs=request_kwargs, + messages=messages, + caller=caller, + cache=cache, + counter_for_model=counter_for_model, + now=now, + ) + + +@pytest.mark.asyncio +async def test_warm_stronger_model_wins_after_counting_input_and_output_cost() -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + assert choice is not None + assert (choice.model, choice.tier, choice.deployment_id) == ("strong", "COMPLEX", "test-cache-strong") + assert choice.original_cost == pytest.approx(50000 * 0.00000125 + 1000 * 0.000001 + 1024 * 0.000004) + assert choice.estimated_cost == pytest.approx(50000 * 0.0000003 + 1000 * 0.000003 + 1024 * 0.000015) + + +@pytest.mark.asyncio +async def test_output_price_can_outweigh_the_cache_saving() -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(strong_output_rate=0.001), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + assert choice is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("case", ["missing", "expired", "different_caller", "changed_prefix", "unauthorized"]) +async def test_no_cache_discount_without_fresh_authorized_matching_evidence(case: str) -> None: + cache: Final = DualCache() + if case != "missing": + await _observed( + cache, + caller="someone-else" if case == "different_caller" else _CALLER, + expires_at=999.0 if case == "expired" else 1290.0, + ) + choice: Final = await _select( + router=_router(), + config=_config(), + response=_response(), + body=_body("Changed context" if case == "changed_prefix" else "Stable cached context"), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap"] if case == "unauthorized" else ["cheap", "strong"]), + cache=cache, + counter_for_model=_forbidden_counter, + now=_NOW, + ) + assert choice is None + + +@pytest.mark.asyncio +async def test_disabled_setting_does_not_access_prediction_services() -> None: + config: Final = ComplexityRouterConfig(tiers={"SIMPLE": "cheap"}) + assert config.cache_aware_routing is False + assert ( + await choose_cached_model( + router=_router(), + config=config, + params_for_model=ComplexityRouter("smart", _router(), config.model_dump())._litellm_params_for_model, + response=_response(), + request_kwargs={}, + messages=None, + ) + is None + ) + + +def test_cache_prices_cannot_add_a_model_below_the_classified_tier() -> None: + response: Final = _response("COMPLEX", "strong") + assert response.routing_decision is not None + assert eligible_models(_config(), response.routing_decision) == (("COMPLEX", "strong"),) + + +@pytest.mark.parametrize( + "overrides", [{"adaptive": True}, {"session_affinity": True}, {"classification_mode": "user_turn"}] +) +def test_existing_pinned_or_adaptive_policies_are_preserved(overrides: Mapping[str, object]) -> None: + response: Final = _response() + assert response.routing_decision is not None + assert eligible_models(_config(**overrides), response.routing_decision) == () + + +@pytest.mark.parametrize("total,prefix", [(None, 50000), (51000, None), (1000, 50000)]) +@pytest.mark.asyncio +async def test_unavailable_or_inconsistent_counts_keep_the_classified_model( + total: int | None, prefix: int | None +) -> None: + cache: Final = DualCache() + await _observed(cache) + assert ( + await _select( + router=_router(), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=lambda _: _Counts(total, prefix), + now=_NOW, + ) + is None + ) + + +@pytest.mark.asyncio +async def test_output_estimate_is_capped_by_the_requested_limit() -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(), + config=_config(cache_aware_routing_output_tokens=100000, max_tokens_from_tier_model=False), + response=_response(), + body={**_body(), "max_tokens": 1}, + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + assert choice is not None + assert choice.estimated_cost == pytest.approx(50000 * 0.0000003 + 1000 * 0.000003 + 0.000015) + + +@pytest.mark.asyncio +async def test_warm_model_that_cannot_fit_the_request_is_not_selected() -> None: + cache: Final = DualCache() + await _observed(cache) + assert ( + await _select( + router=_router(), + config=_config(max_tokens_from_tier_model=False), + response=_response(), + body={**_body(), "max_tokens": 100000000}, + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + is None + ) + + +def test_repeated_model_in_multiple_tiers_is_only_considered_once() -> None: + decision: Final = _response().routing_decision + assert decision is not None + assert eligible_models(_config(tiers={"SIMPLE": "cheap", "MEDIUM": "strong", "COMPLEX": "strong"}), decision) == ( + ("SIMPLE", "cheap"), + ("MEDIUM", "strong"), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "enabled,behavior,expected", + [ + (False, "success", "cheap"), + (True, "success", "strong"), + (True, "tier_cost", "cheap"), + (True, "tier_context", "cheap"), + (True, "error", "cheap"), + (True, "deadline", "cheap"), + (True, "cancel", None), + (True, "transformed", "cheap"), + (True, "unsupported_shape", "cheap"), + (True, "custom_endpoint", "cheap"), + (True, "compaction", "cheap"), + (True, "guardrail", "cheap"), + ], +) +async def test_router_applies_opt_in_and_preserves_failure_semantics( + monkeypatch: pytest.MonkeyPatch, enabled: bool, behavior: str, expected: str | None +) -> None: + import asyncio + import json + + import httpx + + import litellm + from litellm.caching.llm_caching_handler import LLMClientCache + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.utils import ProxyLogging + from litellm.router_strategy.complexity_router.context_compaction import initialize_compaction_state + + config: Final = _config( + cache_aware_routing=enabled, cache_aware_routing_timeout_ms=1 if behavior == "deadline" else 2000 + ) + models: Final = _router( + strong_output_rate=0.000048 if behavior == "tier_cost" else 0.000015, + cheap_limit=50 if behavior == "tier_cost" else 30000, + strong_limit=60000 if behavior == "tier_context" else 30000, + ).get_model_list() + assert models is not None + router: Final = Router( + model_list=[ + *models, + { + "model_name": "smart", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": config.model_dump(), + **({"temperature": 0.1} if behavior == "transformed" else {}), + }, + }, + ] + ) + logging: Final = ProxyLogging(UserApiKeyCache()) + logging.proxy_hook_mapping["parallel_request_limiter"] = _PROXY_MaxParallelRequestsHandler_v3( + logging.internal_usage_cache + ) + await _observed(logging.internal_usage_cache.dual_cache, expires_at=1e100) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", logging) + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) + + requests: Final = asyncio.Queue[httpx.Request]() + + async def count(request: httpx.Request) -> httpx.Response: + requests.put_nowait(request) + assert enabled and behavior not in ("transformed", "unsupported_shape", "custom_endpoint") + if behavior == "error": + return httpx.Response(503, json={"error": "Provider unavailable"}) + if behavior == "cancel": + raise asyncio.CancelledError() + if behavior == "deadline": + await asyncio.Future() + payload: Final = json.loads(request.content) + assert request.url == "https://api.anthropic.com/v1/messages/count_tokens" + assert request.headers["x-api-key"] == _PROVIDER_KEY + return httpx.Response(200, json={"input_tokens": 51000 if "What is 2 + 2?" in str(payload) else 50000}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(count)) as client: + handler: Final = AsyncHTTPHandler() + await handler.client.aclose() + handler.client = client + litellm.in_memory_llm_clients_cache.set_cache("async_httpx_clientanthropic", handler) + body: Final = { + **_body(), + **({"max_tokens": 1000} if behavior in ("tier_cost", "tier_context") else {}), + **({"thinking": {"type": "enabled", "budget_tokens": 10000}} if behavior == "unsupported_shape" else {}), + } + kwargs: Final = { + "litellm_metadata": { + "user_api_key_auth": UserAPIKeyAuth(api_key=_CALLER, models=["smart", "cheap", "strong"]) + }, + "proxy_server_request": {"url": "http://localhost/v1/messages", "body": body, "headers": {}}, + **({"api_base": "https://custom.example"} if behavior == "custom_endpoint" else {}), + **( + {"_context_compaction_state": initialize_compaction_state({}, "messages")} + if behavior == "compaction" + else {} + ), + **({"guardrails": ["test-guardrail"]} if behavior == "guardrail" else {}), + } + if expected is None: + with pytest.raises(asyncio.CancelledError): + await router.async_pre_routing_hook(model="smart", request_kwargs=kwargs, messages=body["messages"]) + return + response: Final = await router.async_pre_routing_hook( + model="smart", request_kwargs=kwargs, messages=body["messages"] + ) + assert response is not None + assert response.model == expected + if not enabled or behavior in ( + "transformed", + "unsupported_shape", + "custom_endpoint", + "compaction", + "guardrail", + ): + assert requests.qsize() == 0 + if enabled and behavior == "success": + assert requests.qsize() == 4 + assert response.routing_decision is not None + assert response.routing_decision["cause"] == ( + "prompt_cache_cost" if expected == "strong" else "heuristic_scorer" + ) + + +@pytest.mark.asyncio +async def test_equal_costs_keep_the_classified_model() -> None: + cache: Final = DualCache() + await _observed(cache) + assert ( + await _select( + router=_router(free=True), + config=_config(), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + is None + ) + + +@pytest.mark.parametrize( + "cause", ["llm_v2_classifier", "capability_classifier", "heuristic_first_short_circuit", "hybrid_short_circuit"] +) +def test_successful_classifiers_can_consider_cache_costs(cause: str) -> None: + response: Final = PreRoutingHookResponse.model_validate( + { + "model": "cheap", + "messages": None, + "routing_decision": {"tier": "SIMPLE", "cause": cause}, + } + ) + assert response.routing_decision is not None + assert eligible_models(_config(), response.routing_decision) == (("SIMPLE", "cheap"), ("COMPLEX", "strong")) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "cheap_limit,strong_limit,requested,from_tier,output_rate,expected_limits", + [ + (50, 1000, 1000, True, 0.000048, None), + (30000, 30000, 1, True, 0.000049, None), + (30000, 60000, 1, True, 0.000015, None), + (50, 100, 20000, True, 0.000048, (50, 100)), + (30000, 30000, 1, False, 0.000048, (1, 1)), + ], +) +async def test_each_candidate_uses_its_effective_routed_output_limit( + cheap_limit: int, + strong_limit: int, + requested: int, + from_tier: bool, + output_rate: float, + expected_limits: tuple[int, int] | None, +) -> None: + cache: Final = DualCache() + await _observed(cache) + choice: Final = await _select( + router=_router(strong_output_rate=output_rate, cheap_limit=cheap_limit, strong_limit=strong_limit), + config=_config(max_tokens_from_tier_model=from_tier), + response=_response(), + body={**_body(), "max_tokens": requested}, + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong"]), + cache=cache, + counter_for_model=_counter_for_model, + now=_NOW, + ) + if expected_limits is None: + assert choice is None + return + assert choice is not None + assert choice.original_cost == pytest.approx(50000 * 0.00000125 + 1000 * 0.000001 + expected_limits[0] * 0.000004) + assert choice.estimated_cost == pytest.approx( + 50000 * 0.0000003 + 1000 * 0.000003 + expected_limits[1] * output_rate + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("warm,authorized", [(False, True), (True, True), (True, False)]) +async def test_authorization_only_runs_for_original_and_warm_alternatives_before_provider_counts( + monkeypatch: pytest.MonkeyPatch, warm: bool, authorized: bool +) -> None: + from unittest.mock import AsyncMock + + from litellm.proxy.common_utils import cache_aware_routing + + cache: Final = DualCache() + if warm: + await _observed(cache) + models: Final = _router().get_model_list() + assert models is not None + router: Final = Router( + model_list=[ + *models, + {**models[0], "model_name": "cold", "model_info": {"id": "test-cache-cold"}}, + ] + ) + authorization: Final = AsyncMock(wraps=cache_aware_routing.can_key_call_resolved_model) + monkeypatch.setattr(cache_aware_routing, "can_key_call_resolved_model", authorization) + + def counter_for_model(model: str) -> TokenCounter: + assert warm and authorized + assert authorization.await_count == 2 + return _Counts() + + choice: Final = await _select( + router=router, + config=_config(tiers={"SIMPLE": "cheap", "MEDIUM": "cold", "COMPLEX": "strong"}), + response=_response(), + body=_body(), + request_kwargs={}, + messages=None, + caller=UserAPIKeyAuth(api_key=_CALLER, models=["cheap", "strong", "cold"] if authorized else ["cheap", "cold"]), + cache=cache, + counter_for_model=counter_for_model, + now=_NOW, + ) + assert (choice is not None) == (warm and authorized) + assert tuple(call.kwargs["model"] for call in authorization.await_args_list) == ( + ("cheap", "strong") if warm else () + ) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b5d515bd1fb..6d45f6e691c 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -39483,6 +39483,24 @@ export interface components { adaptive_eligible: "all" | "classified_tier"; /** @description Quality vs cost weights for adaptive selection (used when adaptive=True) */ adaptive_weights?: components["schemas"]["AdaptiveRouterWeights"]; + /** + * Cache Aware Routing + * @description Opt in to comparing prompt-cache costs after classification. On supported native Anthropic proxy requests, an already warm model in the same or a higher tier may replace the classified model when its estimated input and output cost is lower. Unsupported requests and unavailable estimates keep ordinary routing. + * @default false + */ + cache_aware_routing: boolean; + /** + * Cache Aware Routing Output Tokens + * @description Expected output tokens used in cache-aware cost comparisons; capped by each model's effective output limit. + * @default 1024 + */ + cache_aware_routing_output_tokens: number; + /** + * Cache Aware Routing Timeout Ms + * @description Total time budget for cache-aware predictions; expiry preserves the original routing decision. + * @default 2000 + */ + cache_aware_routing_timeout_ms: number; /** @description Probability threshold policy required when classifier_type is 'capability'. The classifier forecasts p_solve for efficient_tier, adjusts base_threshold using the capability-card boundary, and otherwise routes to capable_tier */ capability_classifier_config?: components["schemas"]["CapabilityClassifierConfig"] | null; /** @@ -42605,7 +42623,7 @@ export interface components { * Cause * @enum {string} */ - cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "capability_classifier" | "jev_classifier" | "llm_v2_classifier" | "llm_v2_fallback" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "capability_classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit"; + cause?: "prompt_cache_cost" | "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "capability_classifier" | "jev_classifier" | "llm_v2_classifier" | "llm_v2_fallback" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "capability_classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit"; /** Classifier Calibrated Capable P Solve */ classifier_calibrated_capable_p_solve?: number; /** Classifier Calibrated Efficient P Solve */ From 7aba77197dc53737f8e882bfceab493397a424b0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 18:15:45 -0700 Subject: [PATCH 15/51] feat(otel): add SigNoz preset for OpenTelemetry v2 (#43296) * feat(otel): add SigNoz preset for OpenTelemetry v2 Adds the signoz callback (OTLP/HTTP exporter, GenAI vocabulary, key and team level dynamic ingestion endpoint and key) as an OpenTelemetry v2 preset, with the preset factory accepting the allow_missing_credentials kwarg the V2 registry always passes so construction no longer falls back silently to legacy OpenTelemetry. Ships the deterministic tests/integration/observability/test_signoz_delivery.py audit suite Absorbs the work from https://github.com/BerriAI/litellm/pull/38206 Co-authored-by: Nagesh Bansal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(otel): drop explanatory comments from the SigNoz preset Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(types): keep signoz dynamic param lines within ruff format width Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(signoz): assert the missing-endpoint boot path directly instead of in an except block Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(ui): regenerate schema.d.ts for the signoz health service Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(otel): allowlist SigNoz key/team endpoints and route keyless collectors without the operator key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(otel): terminate the SigNoz shutdown cell before the flush and drop test docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(otel): keep the shared tenant routing untouched and require an ingestion key for SigNoz key/team endpoints Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(otel): warn about a keyless SigNoz team endpoint from the header resolver so the shared cache actually reaches it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Nagesh Bansal Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/__init__.py | 1 + litellm/integrations/callback_configs.json | 21 + litellm/integrations/otel/model/config.py | 1 + litellm/integrations/otel/presets/__init__.py | 9 + litellm/integrations/otel/presets/signoz.py | 95 ++ .../custom_logger_registry.py | 1 + .../initialize_dynamic_callback_params.py | 6 + litellm/litellm_core_utils/litellm_logging.py | 32 + .../_experimental/out/assets/logos/signoz.svg | 1 + litellm/proxy/_types.py | 6 + .../health_endpoints/_health_endpoints.py | 2 + litellm/proxy/litellm_pre_call_utils.py | 2 + litellm/types/utils.py | 3 + .../observability/test_signoz_delivery.py | 980 ++++++++++++++++++ .../proxy/test_litellm_pre_call_utils.py | 28 + .../integrations/otel/test_otel_v2_dynamic.py | 66 ++ .../integrations/otel/test_otel_v2_presets.py | 66 ++ .../test_litellm_logging.py | 97 ++ .../public/assets/logos/signoz.svg | 1 + .../src/components/callback_info_helpers.tsx | 12 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 21 files changed, 1431 insertions(+), 1 deletion(-) create mode 100644 litellm/integrations/otel/presets/signoz.py create mode 100644 litellm/proxy/_experimental/out/assets/logos/signoz.svg create mode 100644 tests/integration/observability/test_signoz_delivery.py create mode 100644 ui/litellm-dashboard/public/assets/logos/signoz.svg diff --git a/litellm/__init__.py b/litellm/__init__.py index 5a7d6e8125d..5d10737e876 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -172,6 +172,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "levo", "compression_interception", "newrelic", + "signoz", ] cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json index 4e72075dc5c..190c283d087 100644 --- a/litellm/integrations/callback_configs.json +++ b/litellm/integrations/callback_configs.json @@ -502,6 +502,27 @@ }, "description": "S3 Bucket (AWS) Logging Integration" }, + { + "id": "signoz", + "displayName": "SigNoz", + "logo": "signoz.svg", + "supports_key_team_logging": true, + "dynamic_params": { + "signoz_ingestion_endpoint": { + "type": "text", + "ui_name": "SigNoz Ingestion Endpoint", + "description": "Ingestion endpoint for this team, e.g. https://ingest.us.signoz.cloud:443 for SigNoz Cloud or your own collector. Leave blank to use the proxy's configured endpoint. Regions: https://signoz.io/docs/ingestion/signoz-cloud/overview/", + "required": false + }, + "signoz_ingestion_key": { + "type": "password", + "ui_name": "SigNoz Ingestion Key (optional)", + "description": "Ingestion key for this team, so its traces land in its own SigNoz account. Not needed for self-hosted SigNoz. Keys: https://signoz.io/docs/ingestion/signoz-cloud/keys/", + "required": false + } + }, + "description": "SigNoz Logging Integration. Setup: https://signoz.io/docs/litellm-observability/" + }, { "id": "sqs", "displayName": "SQS", diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index 5447a8ee80a..5a3965862e0 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -41,6 +41,7 @@ class ExporterOwner(str, Enum): LEVO = "levo" AGENTOPS = "agentops" NEWRELIC = "newrelic" + SIGNOZ = "signoz" class _OTelV2Flag(BaseSettings): diff --git a/litellm/integrations/otel/presets/__init__.py b/litellm/integrations/otel/presets/__init__.py index a0cd5b3fd98..7c891c29409 100644 --- a/litellm/integrations/otel/presets/__init__.py +++ b/litellm/integrations/otel/presets/__init__.py @@ -30,6 +30,11 @@ from litellm.integrations.otel.presets.phoenix import ( phoenix_preset, phoenix_project_headers, ) +from litellm.integrations.otel.presets.signoz import ( + signoz_dynamic_endpoint, + signoz_dynamic_headers, + signoz_preset, +) from litellm.integrations.otel.presets.weave import weave_dynamic_headers, weave_preset from litellm.types.utils import StandardCallbackDynamicParams @@ -44,6 +49,7 @@ PRESET_BY_CALLBACK: Final[Mapping[str, Preset]] = MappingProxyType( "langtrace": langtrace_preset, "levo": levo_preset, "newrelic": newrelic_preset, + "signoz": signoz_preset, "weave_otel": weave_preset, } ) @@ -58,6 +64,7 @@ DYNAMIC_HEADERS_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynami "arize": arize_dynamic_headers, "langfuse_otel": langfuse_dynamic_headers, "newrelic": newrelic_dynamic_headers, + "signoz": signoz_dynamic_headers, "weave_otel": weave_dynamic_headers, } ) @@ -71,6 +78,7 @@ DYNAMIC_ENDPOINT_BY_CALLBACK: Final[Mapping[str, Callable[[StandardCallbackDynam MappingProxyType( { "newrelic": newrelic_dynamic_endpoint, + "signoz": signoz_dynamic_endpoint, } ) ) @@ -153,5 +161,6 @@ __all__ = [ "newrelic_preset", "phoenix_preset", "project_routing_headers", + "signoz_preset", "weave_preset", ] diff --git a/litellm/integrations/otel/presets/signoz.py b/litellm/integrations/otel/presets/signoz.py new file mode 100644 index 00000000000..c4d7ed48a38 --- /dev/null +++ b/litellm/integrations/otel/presets/signoz.py @@ -0,0 +1,95 @@ +from functools import lru_cache +from types import MappingProxyType +from typing import Final + +from pydantic import Field +from pydantic_settings import BaseSettings, SettingsConfigDict + +import litellm +from litellm._logging import verbose_logger +from litellm.integrations.otel.model.config import ( + ExporterOwner, + ExporterSpec, + OpenTelemetryV2Config, +) +from litellm.integrations.otel.presets.utils import ensure_mappers +from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host +from litellm.types.utils import StandardCallbackDynamicParams + +SIGNOZ_INGESTION_ENDPOINT_ENV: Final = "SIGNOZ_INGESTION_ENDPOINT" + + +class _SigNozSettings(BaseSettings): + model_config = SettingsConfigDict(case_sensitive=False, extra="ignore") + + endpoint: str | None = Field(default=None, validation_alias=SIGNOZ_INGESTION_ENDPOINT_ENV) + ingestion_key: str | None = Field(default=None, validation_alias="SIGNOZ_INGESTION_KEY") + + +def signoz_preset( + *, + config_overrides: OpenTelemetryV2Config | None = None, + allow_missing_credentials: bool = False, +) -> OpenTelemetryV2Config: + settings: Final = _SigNozSettings() + base: Final = config_overrides or OpenTelemetryV2Config() + key: Final = settings.ingestion_key + spec: Final = ExporterSpec( + kind="otlp_http", + endpoint=settings.endpoint, + headers=(f"signoz-ingestion-key={key}" if key else None), + owner=ExporterOwner.SIGNOZ, + requires_headers=bool(key), + ) + return base.model_copy( + update=MappingProxyType( + { + "exporters": (*base.exporters, spec), + "mapper_names": ensure_mappers(base.mapper_names, "genai"), + } + ) + ) + + +@lru_cache(maxsize=128) +def _warn_host_not_allowlisted(endpoint: str) -> None: + verbose_logger.warning( + "SigNoz: not exporting to key/team endpoint '%s'. Add its host to " + "litellm_settings.provider_url_destination_allowed_hosts to permit it", + endpoint, + ) + + +@lru_cache(maxsize=128) +def _warn_endpoint_without_key(endpoint: str) -> None: + verbose_logger.warning( + "SigNoz: not exporting to key/team endpoint '%s'. Set signoz_ingestion_key alongside it; " + "a keyless collector needs the global callback", + endpoint, + ) + + +def _tenant_endpoint_is_unusable(params: StandardCallbackDynamicParams) -> bool: + return bool(params.get("signoz_ingestion_endpoint")) and signoz_dynamic_endpoint(params) is None + + +def signoz_dynamic_endpoint(params: StandardCallbackDynamicParams) -> str | None: + endpoint: Final = params.get("signoz_ingestion_endpoint") + if not endpoint or not endpoint.startswith(("http://", "https://")): + return None + if not params.get("signoz_ingestion_key"): + _warn_endpoint_without_key(endpoint) + return None + if not is_url_destination_allowed_by_host(endpoint, litellm.provider_url_destination_allowed_hosts): + _warn_host_not_allowlisted(endpoint) + return None + return endpoint + + +def signoz_dynamic_headers( + params: StandardCallbackDynamicParams, +) -> dict[str, str]: # mutable-ok: DYNAMIC_HEADERS_BY_CALLBACK returns a dict + key: Final = params.get("signoz_ingestion_key") + if _tenant_endpoint_is_unusable(params) or not key: + return {} # mutable-ok: same registry contract + return {"signoz-ingestion-key": key} # mutable-ok: same registry contract diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 7049fdd1f39..1d277995211 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -89,6 +89,7 @@ class CustomLoggerRegistry: "langtrace": OpenTelemetry, "weave_otel": OpenTelemetry, "levo": OpenTelemetry, + "signoz": OpenTelemetry, "mlflow": MlflowLogger, "langfuse": LangfusePromptManagement, "otel": OpenTelemetry, diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index 00ab05aba77..3100ca6fba1 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -113,6 +113,8 @@ _supported_callback_params: Final[tuple[str, ...]] = ( "dd_agent_port", "newrelic_api_key", "newrelic_region", + "signoz_ingestion_endpoint", + "signoz_ingestion_key", "turn_off_message_logging", ) @@ -126,6 +128,8 @@ _request_blocked_callback_params: Final = frozenset( "dd_agent_port", "newrelic_api_key", "newrelic_region", + "signoz_ingestion_endpoint", + "signoz_ingestion_key", } ) @@ -138,6 +142,8 @@ _trusted_overlay_callback_params: Final = frozenset( { "newrelic_api_key", "newrelic_region", + "signoz_ingestion_endpoint", + "signoz_ingestion_key", } ) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 9ee7a7b0a7a..152fd54e55d 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4931,6 +4931,38 @@ def _init_custom_logger_compatible_class( _in_memory_loggers.append(_otel_logger) return _otel_logger + elif logging_integration == "signoz": + from litellm.integrations.otel.presets.signoz import ( + SIGNOZ_INGESTION_ENDPOINT_ENV, + ) + + _signoz_endpoint: Final = os.getenv(SIGNOZ_INGESTION_ENDPOINT_ENV) + if not _signoz_endpoint: + raise ValueError(f"{SIGNOZ_INGESTION_ENDPOINT_ENV} not found in environment variables") + + _signoz_v2: Final = _maybe_construct_otel_v2("signoz", _in_memory_loggers) + if _signoz_v2 is not None: + return _signoz_v2 + + from litellm.integrations.opentelemetry import ( + OpenTelemetry, + OpenTelemetryConfig, + ) + + _signoz_base: Final = _signoz_endpoint.rstrip("/") + _signoz_key: Final = os.getenv("SIGNOZ_INGESTION_KEY") + _signoz_config: Final = OpenTelemetryConfig( + exporter="otlp_http", + endpoint=(_signoz_base if _signoz_base.endswith("/v1/traces") else f"{_signoz_base}/v1/traces"), + headers=(f"signoz-ingestion-key={_signoz_key}" if _signoz_key else None), + ) + for callback in _in_memory_loggers: + if isinstance(callback, OpenTelemetry) and callback.callback_name == "signoz": + return callback + _signoz_logger: Final = OpenTelemetry(config=_signoz_config, callback_name="signoz") + _in_memory_loggers.append(_signoz_logger) + return _signoz_logger + elif logging_integration == "mlflow": for callback in _in_memory_loggers: if isinstance(callback, MlflowLogger): diff --git a/litellm/proxy/_experimental/out/assets/logos/signoz.svg b/litellm/proxy/_experimental/out/assets/logos/signoz.svg new file mode 100644 index 00000000000..9064cb86bd6 --- /dev/null +++ b/litellm/proxy/_experimental/out/assets/logos/signoz.svg @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 14aa42afefd..34d7fc1e0f0 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4035,6 +4035,12 @@ class AllCallbacks(LiteLLMPydanticObjectBase): ], ) + signoz: CallbackOnUI = CallbackOnUI( + litellm_callback_name="signoz", + ui_callback_name="SigNoz", + litellm_callback_params=("SIGNOZ_INGESTION_ENDPOINT", "SIGNOZ_INGESTION_KEY"), + ) + zerobus: CallbackOnUI = CallbackOnUI( litellm_callback_name="zerobus", ui_callback_name="Databricks Zerobus", diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index fbd4d57bf77..07be73d7573 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -221,6 +221,7 @@ services = ( "galileo", "newrelic", "pointfive", + "signoz", "sqs", ] | str @@ -309,6 +310,7 @@ async def health_services_endpoint( "galileo", "newrelic", "pointfive", + "signoz", "sqs", ]: raise HTTPException( diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 56f647d5acc..866d84ca8f2 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -924,6 +924,8 @@ def convert_key_logging_metadata_to_callback( # must not export to it. if var.startswith("newrelic_") and data.callback_name != "newrelic": continue + if var.startswith("signoz_") and data.callback_name != "signoz": + continue if team_callback_settings_obj.callback_vars is None: team_callback_settings_obj.callback_vars = {} team_callback_settings_obj.callback_vars[var] = str(value) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 4bda1dd53ce..8862df9dc22 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3679,6 +3679,9 @@ class StandardCallbackDynamicParams(TypedDict, total=False): newrelic_api_key: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns into the dict newrelic_region: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns into the dict + signoz_ingestion_endpoint: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns it + signoz_ingestion_key: str | None # writable-ok: initialize_standard_callback_dynamic_params assigns it + # Logging settings turn_off_message_logging: bool | None # when true will not log messages litellm_disabled_callbacks: list[str] | None diff --git a/tests/integration/observability/test_signoz_delivery.py b/tests/integration/observability/test_signoz_delivery.py new file mode 100644 index 00000000000..f3d715fe5cf --- /dev/null +++ b/tests/integration/observability/test_signoz_delivery.py @@ -0,0 +1,980 @@ +import asyncio +import base64 +import json +import os +import re +import signal +import threading +import uuid +from collections import deque +from collections.abc import Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from functools import partial +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment, object_value, 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, wire_server +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest +from opentelemetry.proto.common.v1.common_pb2 import AnyValue +from pydantic import JsonValue, TypeAdapter + +MARKER: Final = re.compile(rb"signoz-[0-9a-f]{32}") +JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +RESPONSE_ID: Final = "gen_ai.response.id" +INGESTION_HEADER: Final = "signoz-ingestion-key" +OPERATOR_KEY: Final = "operator-ingestion-" + uuid.uuid4().hex +TENANT_KEY: Final = "tenant-ingestion-" + uuid.uuid4().hex + + +def _marker() -> str: + return "signoz-" + uuid.uuid4().hex + + +def _chat_reply(identity: str, stream: bool) -> Reply: + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "signoz ok"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + chunk: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + return Reply( + content_type="text/event-stream", + chunks=( + b"data: " + + json.dumps( + {**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "signoz"}}]} + ).encode() + + b"\n\n", + b"data: " + + json.dumps( + {**chunk, "choices": [{"index": 0, "delta": {"content": " ok"}, "finish_reason": "stop"}]} + ).encode() + + b"\n\n", + b"data: " + + json.dumps( + {**chunk, "choices": [], "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}} + ).encode() + + b"\n\n", + b"data: [DONE]\n\n", + ), + ) + + +def _responses_reply(identity: str, stream: bool) -> Reply: + response: Final[dict[str, JsonValue]] = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": "msg_" + identity, + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "signoz ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}, + } + if not stream: + return Reply(body=json.dumps(response).encode()) + events: Final[tuple[dict[str, JsonValue], ...]] = ( + { + "type": "response.created", + "sequence_number": 0, + "response": {**response, "status": "in_progress", "output": []}, + }, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": "signoz ok", + }, + {"type": "response.completed", "sequence_number": 2, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _upstream(request: Request) -> Reply: + found: Final = MARKER.search(request.body) + if found is None: + return Reply(status=404, body=b'{"error":"no marker"}') + if request.headers.get("authorization") == "Bearer revoked-provider-key": + return Reply( + status=401, body=b'{"error":{"message":"Incorrect API key provided","type":"invalid_request_error"}}' + ) + marker: Final = found.group(0).decode() + stream: Final = object_value(JSON.validate_json(request.body)).get("stream") is True + if request.target.endswith("/responses"): + return _responses_reply(f"resp_{marker}", stream) + return _chat_reply(f"chatcmpl-{marker}", stream) + + +def _decoded_responses_id(identity: str) -> str: + try: + return base64.b64decode(identity.removeprefix("resp_").encode()).decode() + except (ValueError, UnicodeDecodeError): + return identity + + +def _canonical_id(identity: str) -> str: + return _decoded_responses_id(identity).rpartition("response_id:")[2] + + +def _sse_events(text: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(JSON.validate_json(line[6:])) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + + +def _text_at(payload: JsonValue, *path: str) -> str: + if not path: + return string_value(payload) + return _text_at(object_value(payload)[path[0]], *path[1:]) + + +def _body_id(response: httpx.Response) -> str: + return _text_at(JSON.validate_json(response.content), "id") + + +@dataclass(frozen=True, slots=True) +class Span: + target: str + ingestion_key: str | None + attributes: Mapping[str, str] + + +def _attribute_text(value: AnyValue) -> str: + match value.WhichOneof("value"): + case "string_value": + return value.string_value + case "int_value": + return str(value.int_value) + case "double_value": + return str(value.double_value) + case "bool_value": + return str(value.bool_value) + case _: + return "" + + +@dataclass(frozen=True, slots=True) +class Collector: + wire: Wire + outage: threading.Event + rejection: threading.Event + missing: threading.Event + slow: threading.Event + release: threading.Event + accepted: Sequence[Request] + refused: Sequence[Request] + guard: threading.Lock + + def refused_batch_carrying(self, response_id: str) -> Request: + def carrying() -> tuple[Request, ...]: + with self.guard: + return tuple(batch for batch in self.refused if response_id.encode() in batch.body) + + return eventually(carrying, lambda found: len(found) >= 1, seconds=30)[0] + + def refused_batches(self) -> tuple[Request, ...]: + def refused() -> tuple[Request, ...]: + with self.guard: + return tuple(self.refused) + + return eventually(refused, lambda found: len(found) >= 1, seconds=30) + + def spans(self) -> tuple[Span, ...]: + with self.guard: + batches: Final = tuple(self.accepted) + return tuple( + Span( + batch.target, + batch.headers.get(INGESTION_HEADER), + {attribute.key: _attribute_text(attribute.value) for attribute in span.attributes}, + ) + for batch in batches + for resource in ExportTraceServiceRequest.FromString(batch.body).resource_spans + for scope in resource.scope_spans + for span in scope.spans + ) + + def spans_for(self, response_id: str) -> tuple[Span, ...]: + return tuple( + span + for span in self.spans() + if RESPONSE_ID in span.attributes + and _canonical_id(span.attributes[RESPONSE_ID]) == _canonical_id(response_id) + ) + + def single_span(self, response_id: str, *, elsewhere: "Collector | None" = None) -> Span: + found: Final = eventually( + lambda: self.spans_for(response_id), lambda spans: len(spans) == 1, seconds=30, return_last_on_timeout=True + ) + assert len(found) == 1, ( + f"{len(found)} spans for {response_id} at this sink; other sink saw " + f"{elsewhere.landed((response_id,)) if elsewhere else 'n/a'}" + ) + return found[0] + + def landed(self, response_ids: Sequence[str]) -> dict[str, int]: + spans: Final = self.spans() + return { + _canonical_id(identity): sum( + 1 + for span in spans + if RESPONSE_ID in span.attributes + and _canonical_id(span.attributes[RESPONSE_ID]) == _canonical_id(identity) + ) + for identity in response_ids + } + + +def _collector() -> Iterator[Collector]: + outage: Final = threading.Event() + rejection: Final = threading.Event() + missing: Final = threading.Event() + slow: Final = threading.Event() + release: Final = threading.Event() + accepted: Final[deque[Request]] = deque() # mutable-ok: the sink thread records each accepted batch as it arrives + refused: Final[deque[Request]] = deque() # mutable-ok: the sink thread records each refused batch as it arrives + guard: Final = threading.Lock() + + def refuse(request: Request, status: int, body: bytes) -> Reply: + with guard: + refused.append(request) + return Reply(status=status, body=body) + + def sink(request: Request) -> Reply: + if slow.is_set(): + release.wait(timeout=30) + if outage.is_set(): + return refuse(request, 503, b'{"error":"sink down"}') + if rejection.is_set(): + return refuse(request, 403, b'{"error":"forbidden"}') + if missing.is_set(): + return refuse(request, 404, b'{"error":"not found"}') + with guard: + accepted.append(request) + return Reply() + + with wire_server(sink) as wire: + yield Collector(wire, outage, rejection, missing, slow, release, accepted, refused, guard) + + +@pytest.fixture(scope="session") +def operator_sink() -> Iterator[Collector]: + yield from _collector() + + +@pytest.fixture(scope="session") +def tenant_sink() -> Iterator[Collector]: + yield from _collector() + + +@pytest.fixture(scope="session") +def provider() -> Iterator[Wire]: + with wire_server(_upstream) as wire: + yield wire + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + process: OwnedProxy + model: str + upstream: Wire + sink: Collector + tenant_sink: Collector + + def openai_client(self) -> openai.OpenAI: + return openai.OpenAI(base_url=str(self.proxy.client.base_url) + "/v1", api_key=self.proxy.key, max_retries=0) + + def async_openai_client(self) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI( + base_url=str(self.proxy.client.base_url) + "/v1", api_key=self.proxy.key, max_retries=0 + ) + + def anthropic_client(self) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=str(self.proxy.client.base_url), api_key=self.proxy.key, max_retries=0) + + def async_anthropic_client(self) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=str(self.proxy.client.base_url), api_key=self.proxy.key, max_retries=0) + + def chat( + self, marker: str, *, headers: Mapping[str, str] | None = None, key: str | None = None, **extra: JsonValue + ) -> httpx.Response: + return self.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": self.model, + "messages": [{"role": "user", "content": marker}], + "cache": {"no-cache": True}, + **extra, + }, + headers=headers, + key=key, + ) + + def upstream_bodies(self, marker: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(JSON.validate_json(request.body)) + for request in self.upstream.drain() + if marker.encode() in request.body + ) + + def spend_rows(self, response_id: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + eventually( + lambda: read_rows( + 'SELECT request_id, spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (response_id,) + ), + lambda values: len(values) == 1, + seconds=70, + ) + ) + + def tenant_logging(self, endpoint: str | None, key: str | None) -> JsonValue: + variables: Final[dict[str, JsonValue]] = { + **({"signoz_ingestion_endpoint": endpoint} if endpoint is not None else {}), + **({"signoz_ingestion_key": key} if key is not None else {}), + } + return [{"callback_name": "signoz", "callback_type": "success", "callback_vars": variables}] + + +@dataclass(frozen=True, slots=True) +class RigFactory: + provider: Wire + sink: Collector + tenant_sink: Collector + directory: Path + otel_v2: bool + workers: int + endpoint: str | None + ingestion_key: str | None = OPERATOR_KEY + + def config_path(self) -> Path: + loaded: Final = object_value( + JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + ) + config: Final = { + **loaded, + "litellm_settings": { + **object_value(loaded["litellm_settings"]), + "callbacks": ["signoz"], + "provider_url_destination_allowed_hosts": [self.tenant_sink.wire.url], + }, + "general_settings": {**object_value(loaded["general_settings"]), "disable_model_info_refresh": True}, + } + path: Final = self.directory / f"signoz-{uuid.uuid4().hex}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + def overrides(self) -> dict[str, str]: + return { + "LITELLM_OTEL_V2": "1" if self.otel_v2 else "0", + "OTEL_BSP_SCHEDULE_DELAY": "300", + **({"SIGNOZ_INGESTION_ENDPOINT": self.endpoint} if self.endpoint is not None else {}), + **({"SIGNOZ_INGESTION_KEY": self.ingestion_key} if self.ingestion_key is not None else {}), + } + + def start(self) -> Iterator[Rig]: + with ( + gateway_from_environment() as gateway, + owned_proxy_process( + gateway, + self.directory, + self.overrides(), + config=self.config_path(), + remove_environment=("SIGNOZ_INGESTION_ENDPOINT", "SIGNOZ_INGESTION_KEY"), + workers=self.workers, + ) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=self.provider.url + "/v1") + yield Rig(owned.gateway, owned, model, self.provider, self.sink, self.tenant_sink) + + +@pytest.fixture(scope="session") +def rig( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Rig]: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz"), False, 2, operator_sink.wire.url + ) + yield from factory.start() + + +@pytest.fixture(scope="session") +def v2_rig( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Rig]: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz-v2"), True, 2, operator_sink.wire.url + ) + yield from factory.start() + + +def _assert_operator_span(rig: Rig, response_id: str, marker: str) -> Span: + span: Final = rig.sink.single_span(response_id) + assert span.target == "/v1/traces", span + assert span.ingestion_key == OPERATOR_KEY, span + assert rig.tenant_sink.landed((response_id,)) == {_canonical_id(response_id): 0} + bodies: Final = rig.upstream_bodies(marker) + assert len(bodies) == 1, bodies + assert "signoz" not in json.dumps(bodies[0]).replace(marker, ""), bodies[0] + return span + + +def test_signoz_is_registered_as_an_opentelemetry_callback(rig: Rig) -> None: + listed: Final = rig.proxy.request("GET", "/active/callbacks") + assert listed.status_code == 200, listed.text + assert "OpenTelemetry" in json.dumps(listed.json()), listed.text + log: Final = rig.process.log.read_text() + assert "SIGNOZ_INGESTION_ENDPOINT not found" not in log + + +def test_chat_completion_sdk_span_lands_at_the_operator_sink_with_the_ingestion_key(rig: Rig) -> None: + marker: Final = _marker() + completion: Final = rig.openai_client().chat.completions.create( + model=rig.model, messages=[{"role": "user", "content": marker}] + ) + assert completion.id == f"chatcmpl-{marker}" + span: Final = _assert_operator_span(rig, completion.id, marker) + assert span.attributes.get("gen_ai.request.model") or span.attributes.get("llm.request.model"), span + rows: Final = rig.spend_rows(completion.id) + assert rows[0]["request_id"] == completion.id, rows + + +def test_chat_stream_async_sdk_span_lands_once_after_the_stream_is_consumed(rig: Rig) -> None: + marker: Final = _marker() + + async def consume() -> frozenset[str]: + stream: Final = await rig.async_openai_client().chat.completions.create( + model=rig.model, messages=[{"role": "user", "content": marker}], stream=True + ) + return frozenset([chunk.id async for chunk in stream]) + + identities: Final = asyncio.run(consume()) + assert identities == {f"chatcmpl-{marker}"}, identities + _assert_operator_span(rig, f"chatcmpl-{marker}", marker) + + +def test_messages_sdk_span_lands_at_the_operator_sink(rig: Rig) -> None: + marker: Final = _marker() + message: Final = rig.anthropic_client().messages.create( + model=rig.model, max_tokens=16, messages=[{"role": "user", "content": marker}] + ) + _assert_operator_span(rig, message.id, marker) + + +def test_messages_stream_async_sdk_span_lands_once_after_the_stream_is_consumed(rig: Rig) -> None: + marker: Final = _marker() + + async def consume() -> str: + async with rig.async_anthropic_client().messages.stream( + model=rig.model, max_tokens=16, messages=[{"role": "user", "content": marker}] + ) as stream: + async for _ in stream: + pass + return (await stream.get_final_message()).id + + identity: Final = asyncio.run(consume()) + _assert_operator_span(rig, identity, marker) + + +def test_responses_sdk_span_lands_at_the_operator_sink(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.openai_client().responses.create(model=rig.model, input=marker) + _assert_operator_span(rig, response.id, marker) + + +def test_responses_stream_raw_httpx_span_lands_once_after_the_stream_is_consumed(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.proxy.request("POST", "/v1/responses", {"model": rig.model, "input": marker, "stream": True}) + assert response.status_code == 200, response.text + _assert_operator_span(rig, _responses_id(response, marker), marker) + + +def test_v2_flag_on_still_delivers_the_operator_span_with_the_ingestion_key(v2_rig: Rig) -> None: + marker: Final = _marker() + response: Final = v2_rig.chat(marker) + assert response.status_code == 200, response.text + _assert_operator_span(v2_rig, _body_id(response), marker) + + +def test_endpoint_already_ending_in_v1_traces_is_not_doubled( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + factory: Final = RigFactory( + provider, + operator_sink, + tenant_sink, + tmp_path_factory.mktemp("signoz-suffixed"), + False, + 2, + operator_sink.wire.url + "/v1/traces", + ) + suffixed: Final = next(started := factory.start()) + marker: Final = _marker() + response: Final = suffixed.chat(marker) + assert response.status_code == 200, response.text + span: Final = suffixed.sink.single_span(_body_id(response)) + assert span.target == "/v1/traces", span + assert tuple(started) == () + + +def test_three_identical_requests_produce_one_span_each(rig: Rig) -> None: + markers: Final = tuple(_marker() for _ in range(3)) + responses: Final = tuple(rig.chat(marker) for marker in markers) + assert all(response.status_code == 200 for response in responses), [response.text for response in responses] + identities: Final = tuple(_body_id(response) for response in responses) + landed: Final = eventually( + lambda: rig.sink.landed(identities), lambda seen: all(count >= 1 for count in seen.values()), seconds=30 + ) + assert landed == {identity: 1 for identity in identities}, landed + assert rig.sink.landed(identities) == landed + + +def test_unauthenticated_request_is_rejected_without_an_upstream_call_and_any_span_records_the_401(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.chat(marker, key="sk-not-a-real-key") + assert response.status_code == 401, response.text + later: Final = rig.chat(_marker()) + assert later.status_code == 200, later.text + rig.sink.single_span(_body_id(later)) + assert rig.upstream_bodies(marker) == () + marker_spans: Final = tuple(span for span in rig.sink.spans() if marker in json.dumps(span.attributes)) + assert all(span.attributes.get("error.code") == "401" for span in marker_spans), marker_spans + assert not any(span.attributes.get(RESPONSE_ID, "").startswith("chatcmpl-") for span in marker_spans), marker_spans + + +def test_request_supplied_signoz_variables_are_refused_before_the_upstream_is_called(rig: Rig) -> None: + marker: Final = _marker() + response: Final = rig.chat( + marker, + metadata={"signoz_ingestion_endpoint": rig.tenant_sink.wire.url, "signoz_ingestion_key": TENANT_KEY}, + ) + assert response.status_code == 401, response.text + assert "signoz_ingestion_endpoint is not allowed in request body" in response.text + assert rig.upstream_bodies(marker) == () + assert not any(marker in json.dumps(span.attributes) for span in rig.tenant_sink.spans()) + + +def test_upstream_401_reaches_the_caller_and_unrelated_traffic_keeps_landing(rig: Rig) -> None: + marker: Final = _marker() + with rig.proxy.scenario() as scenario: + broken: Final = scenario.model(api_base=rig.upstream.url + "/v1", api_key="revoked-provider-key") + failed: Final = rig.proxy.request( + "POST", "/v1/chat/completions", {"model": broken, "messages": [{"role": "user", "content": marker}]} + ) + assert failed.status_code == 401, failed.text + assert "Incorrect API key provided" in failed.text + healthy_marker: Final = _marker() + healthy: Final = rig.chat(healthy_marker) + assert healthy.status_code == 200, healthy.text + _assert_operator_span(rig, _body_id(healthy), healthy_marker) + + +def test_health_services_accepts_signoz(rig: Rig) -> None: + response: Final = rig.proxy.request("GET", "/health/services", params={"service": "signoz"}) + assert response.status_code == 200, response.text + + +def test_sink_answering_403_drops_those_spans_and_later_spans_still_land(rig: Rig) -> None: + rig.sink.rejection.set() + try: + rejected: Final = rig.chat(_marker()) + assert rejected.status_code == 200, rejected.text + rig.sink.refused_batch_carrying(_body_id(rejected)) + finally: + rig.sink.rejection.clear() + later: Final = rig.chat(_marker()) + assert later.status_code == 200, later.text + rig.sink.single_span(_body_id(later)) + assert rig.sink.landed((_body_id(rejected),)) == {_body_id(rejected): 0} + + +def test_sink_answering_404_drops_those_spans_and_later_spans_still_land(rig: Rig) -> None: + rig.sink.missing.set() + try: + dropped: Final = rig.chat(_marker()) + assert dropped.status_code == 200, dropped.text + rig.sink.refused_batch_carrying(_body_id(dropped)) + finally: + rig.sink.missing.clear() + later: Final = rig.chat(_marker()) + assert later.status_code == 200, later.text + rig.sink.single_span(_body_id(later)) + assert rig.sink.landed((_body_id(dropped),)) == {_body_id(dropped): 0} + + +def test_key_level_signoz_destination_routes_the_span_to_the_tenant_sink(v2_rig: Rig) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + token: Final = scenario.key( + metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, TENANT_KEY)} + ) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + identity: Final = _body_id(response) + span: Final = v2_rig.tenant_sink.single_span(identity, elsewhere=v2_rig.sink) + assert span.ingestion_key == TENANT_KEY, span + assert v2_rig.sink.landed((identity,)) == {identity: 0}, "operator sink also received the tenant span" + + +def test_team_level_signoz_destination_routes_the_span_to_the_tenant_sink(v2_rig: Rig) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team( + metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, TENANT_KEY)} + ) + token: Final = scenario.key(team_id=team) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + identity: Final = _body_id(response) + span: Final = v2_rig.tenant_sink.single_span(identity, elsewhere=v2_rig.sink) + assert span.ingestion_key == TENANT_KEY, span + assert v2_rig.sink.landed((identity,)) == {identity: 0}, "operator sink also received the tenant span" + + +def test_key_level_destination_wins_over_the_team_level_destination(v2_rig: Rig) -> None: + marker: Final = _marker() + team_key: Final = "team-" + TENANT_KEY + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team(metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, team_key)}) + token: Final = scenario.key( + team_id=team, metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, TENANT_KEY)} + ) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + span: Final = v2_rig.tenant_sink.single_span(_body_id(response)) + assert span.ingestion_key == TENANT_KEY, span + + +def test_team_endpoint_without_an_ingestion_key_is_ignored_and_the_span_stays_at_the_operator_sink( + v2_rig: Rig, +) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team(metadata={"logging": v2_rig.tenant_logging(v2_rig.tenant_sink.wire.url, None)}) + token: Final = scenario.key(team_id=team) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + _assert_operator_span(v2_rig, _body_id(response), marker) + eventually( + lambda: v2_rig.process.log.read_text(), + lambda text: "Set signoz_ingestion_key alongside it" in text, + seconds=30, + ) + + +def test_team_endpoint_off_the_allowlist_keeps_the_span_at_the_operator_sink(v2_rig: Rig) -> None: + marker: Final = _marker() + with v2_rig.proxy.scenario() as scenario: + team: Final = scenario.team( + metadata={"logging": v2_rig.tenant_logging("http://tenant.invalid:4318/v1/traces", TENANT_KEY)} + ) + token: Final = scenario.key(team_id=team) + response: Final = v2_rig.chat(marker, key=token) + assert response.status_code == 200, response.text + _assert_operator_span(v2_rig, _body_id(response), marker) + eventually( + lambda: v2_rig.process.log.read_text(), + lambda text: "provider_url_destination_allowed_hosts" in text, + seconds=30, + ) + + +def test_legacy_mode_ignores_key_level_signoz_destination_and_keeps_the_operator_sink(rig: Rig) -> None: + marker: Final = _marker() + with rig.proxy.scenario() as scenario: + token: Final = scenario.key(metadata={"logging": rig.tenant_logging(rig.tenant_sink.wire.url, TENANT_KEY)}) + response: Final = rig.chat(marker, key=token) + assert response.status_code == 200, response.text + _assert_operator_span(rig, _body_id(response), marker) + + +@pytest.mark.parametrize( + "endpoint", + ["", "not-a-url", "ftp://tenant.invalid", "x" * 5000, 12345, ["http://tenant.invalid"]], + ids=["empty", "bare", "ftp", "5kb", "int", "list"], +) +def test_hostile_tenant_endpoint_never_breaks_the_request_or_the_operator_sink( + v2_rig: Rig, endpoint: JsonValue +) -> None: + marker: Final = _marker() + created: Final = v2_rig.proxy.request( + "POST", + "/key/generate", + { + "metadata": { + "logging": [ + { + "callback_name": "signoz", + "callback_type": "success", + "callback_vars": {"signoz_ingestion_endpoint": endpoint, "signoz_ingestion_key": TENANT_KEY}, + } + ] + } + }, + ) + assert created.status_code in (200, 400, 422), created.text + if created.status_code != 200: + return + try: + response: Final = v2_rig.chat(marker, key=_text_at(JSON.validate_json(created.content), "key")) + assert response.status_code == 200, response.text + identity: Final = _body_id(response) + eventually( + lambda: v2_rig.sink.landed((identity,))[identity] + v2_rig.tenant_sink.landed((identity,))[identity], + lambda total: total >= 1, + seconds=30, + ) + assert v2_rig.sink.landed((identity,))[identity] + v2_rig.tenant_sink.landed((identity,))[identity] == 1 + assert v2_rig.tenant_sink.landed((identity,)) == {identity: 0}, "unusable endpoint reached the tenant sink" + finally: + v2_rig.proxy.post("/key/delete", {"keys": [created.json()["key"]]}) + + +def test_missing_ingestion_endpoint_fails_loudly_at_boot( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz-noenv"), False, 2, None + ) + started: Final = factory.start() + broken: Final = next(started) + try: + response: Final = broken.chat(_marker()) + assert response.status_code == 200, response.text + eventually( + lambda: broken.process.log.read_text(), + lambda text: "SIGNOZ_INGESTION_ENDPOINT not found" in text, + seconds=30, + ) + assert operator_sink.landed((_body_id(response),)) == {_body_id(response): 0} + finally: + with pytest.raises(StopIteration): + next(started) + + +def test_empty_ingestion_endpoint_is_treated_as_missing( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + factory: Final = RigFactory( + provider, operator_sink, tenant_sink, tmp_path_factory.mktemp("signoz-empty"), False, 2, "" + ) + started: Final = factory.start() + broken: Final = next(started) + try: + response: Final = broken.chat(_marker()) + assert response.status_code == 200, response.text + eventually( + lambda: broken.process.log.read_text(), + lambda text: "SIGNOZ_INGESTION_ENDPOINT not found" in text, + seconds=30, + ) + finally: + with pytest.raises(StopIteration): + next(started) + + +def _is_event_stream(response: httpx.Response) -> bool: + return "content-type" in response.headers and response.headers["content-type"].startswith("text/event-stream") + + +def _chat_id(response: httpx.Response) -> str: + if not _is_event_stream(response): + return _body_id(response) + identities: Final = frozenset(_text_at(event, "id") for event in _sse_events(response.text)) + assert len(identities) == 1, response.text + return next(iter(identities)) + + +def _responses_id(response: httpx.Response, marker: str) -> str: + if not _is_event_stream(response): + return _body_id(response) + completed: Final = tuple( + _text_at(event, "response", "id") + for event in _sse_events(response.text) + if event.get("type") == "response.completed" + ) + assert len(completed) == 1 and completed[0].startswith("resp_"), response.text + return f"resp_{marker}" + + +def _message_id(response: httpx.Response) -> str: + if not _is_event_stream(response): + return _body_id(response) + starts: Final = tuple( + _text_at(event, "message", "id") for event in _sse_events(response.text) if event.get("type") == "message_start" + ) + assert len(starts) == 1, response.text + return starts[0] + + +def _burst(rig: Rig, count: int) -> tuple[tuple[int, str, str | None], ...]: + markers: Final = tuple(_marker() for _ in range(count)) + + def one(index: int) -> tuple[int, str, str | None]: + marker: Final = markers[index] + headers: Final = {"Authorization": f"Bearer {rig.proxy.key}"} + stream: Final = index % 2 == 0 + path, body, identity_of = ( + ("/v1/chat/completions", {"model": rig.model, "messages": [{"role": "user", "content": marker}]}, _chat_id), + ("/v1/responses", {"model": rig.model, "input": marker}, partial(_responses_id, marker=marker)), + ( + "/v1/messages", + {"model": rig.model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}]}, + _message_id, + ), + )[index % 3] + try: + response: Final = rig.proxy.client.post(path, json={**body, "stream": stream}, headers=headers) + response.read() + except httpx.HTTPError as error: + return index, marker, repr(error) + return (index, marker, response.text) if response.status_code != 200 else (index, identity_of(response), None) + + with ThreadPoolExecutor(max_workers=10) as pool: + return tuple(pool.map(one, range(count))) + + +def _assert_exactly_once(rig: Rig, identities: Sequence[str]) -> None: + landed: Final = eventually( + lambda: rig.sink.landed(identities), lambda seen: all(count >= 1 for count in seen.values()), seconds=80 + ) + assert landed == {_canonical_id(identity): 1 for identity in identities}, landed + charged: Final = frozenset(identity for identity in identities if not identity.startswith("resp_")) + spend: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s::text[])', + ("{" + ",".join(charged) + "}",), + ), + lambda rows: {str(row["request_id"]) for row in rows} >= charged, + seconds=70, + ) + assert {str(row["request_id"]) for row in spend} == charged, spend + + +def test_sink_outage_during_a_mixed_burst_lands_every_response_exactly_once_after_recovery(rig: Rig) -> None: + rig.sink.outage.set() + try: + health_down: Final = rig.proxy.request("GET", "/health/services", params={"service": "signoz"}) + results: Final = _burst(rig, 30) + assert all(error is None for _, _, error in results), [error for _, _, error in results if error] + rig.sink.refused_batches() + finally: + rig.sink.outage.clear() + assert health_down.status_code == 200, health_down.text + _assert_exactly_once(rig, tuple(identity for _, identity, _ in results)) + + +def test_slow_sink_during_a_burst_lands_every_response_exactly_once(rig: Rig) -> None: + rig.sink.release.clear() + rig.sink.slow.set() + try: + results: Final = _burst(rig, 20) + assert all(error is None for _, _, error in results), [error for _, _, error in results if error] + identities: Final = tuple(identity for _, identity, _ in results) + assert all(count == 0 for count in rig.sink.landed(identities).values()), "sink accepted while held" + finally: + rig.sink.slow.clear() + rig.sink.release.set() + _assert_exactly_once(rig, identities) + + +def test_killing_one_of_two_workers_mid_burst_keeps_serving_and_never_duplicates_a_span(rig: Rig) -> None: + root: Final = psutil.Process(rig.process.process.pid) + workers: Final = eventually( + lambda: tuple(child for child in root.children() if "resource_tracker" not in " ".join(child.cmdline())), + lambda found: len(found) == 2, + seconds=30, + ) + markers: Final = tuple(_marker() for _ in range(24)) + + def one(index: int) -> tuple[str, str | None]: + if index == 8: + os.kill(workers[0].pid, signal.SIGKILL) + try: + response: Final = rig.chat(markers[index]) + return f"chatcmpl-{markers[index]}", None if response.status_code == 200 else response.text + except httpx.HTTPError as error: + return f"chatcmpl-{markers[index]}", repr(error) + + with ThreadPoolExecutor(max_workers=6) as pool: + results: Final = tuple(pool.map(one, range(24))) + assert rig.process.process.poll() is None, "Proxy root exited after a worker was killed" + after: Final = rig.chat(_marker()) + assert after.status_code == 200, after.text + rig.sink.single_span(_body_id(after)) + failures: Final = tuple(error for _, error in results if error) + assert all(error.startswith(("ReadError(", "RemoteProtocolError(", "ConnectError(")) for error in failures), ( + failures + ) + assert len(failures) <= 6, failures + served: Final = tuple(identity for identity, error in results if error is None) + assert len(served) >= 18, results + settled: Final = tuple(identity for index, (identity, error) in enumerate(results) if index > 14 and not error) + _assert_exactly_once(rig, settled) + assert all(count <= 1 for count in rig.sink.landed(served).values()), rig.sink.landed(served) + + +def test_terminating_the_proxy_right_after_a_burst_flushes_every_span_before_exit( + provider: Wire, operator_sink: Collector, tenant_sink: Collector, tmp_path_factory: pytest.TempPathFactory +) -> None: + pytest.skip("BUG: spans still queued in the OTel batch processor at SIGTERM never reach the sink (4 of 10 lost)") + factory: Final = RigFactory( + provider, + operator_sink, + tenant_sink, + tmp_path_factory.mktemp("signoz-shutdown"), + False, + 2, + operator_sink.wire.url, + ) + started: Final = factory.start() + rig: Final = next(started) + responses: Final = tuple(rig.chat(_marker()) for _ in range(10)) + assert all(response.status_code == 200 for response in responses), [response.text for response in responses] + identities: Final = tuple(_body_id(response) for response in responses) + pending_at_signal: Final = rig.sink.landed(identities) + rig.process.process.terminate() + assert rig.process.process.wait(timeout=40) in (0, -signal.SIGTERM) + with pytest.raises(httpx.ConnectError): + next(started) + landed: Final = rig.sink.landed(identities) + assert landed == {identity: 1 for identity in identities}, ( + f"spans at the sink after exit: {landed}, at the moment of SIGTERM: {pending_at_signal}" + ) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index ee42042bb8e..1a3ffecb0a5 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8470,3 +8470,31 @@ async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custo for name, value in secrets.items(): assert updated["secret_fields"]["raw_headers"][name.lower()] == value assert request.headers[name] == value + + +def test_signoz_callback_vars_are_scoped_to_the_signoz_callback(): + from litellm.proxy._types import AddTeamCallback + from litellm.proxy.litellm_pre_call_utils import convert_key_logging_metadata_to_callback + + under_signoz = convert_key_logging_metadata_to_callback( + data=AddTeamCallback( + callback_name="signoz", + callback_type="success", + callback_vars={"signoz_ingestion_key": "team-key", "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443"}, + ), + team_callback_settings_obj=None, + ) + assert under_signoz.callback_vars == { + "signoz_ingestion_key": "team-key", + "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443", + } + + under_other = convert_key_logging_metadata_to_callback( + data=AddTeamCallback( + callback_name="langfuse", + callback_type="success", + callback_vars={"signoz_ingestion_key": "team-key", "langfuse_host": "https://cloud.langfuse.com"}, + ), + team_callback_settings_obj=None, + ) + assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"} diff --git a/tests/unit/integrations/otel/test_otel_v2_dynamic.py b/tests/unit/integrations/otel/test_otel_v2_dynamic.py index 29772eb92c7..8163ba06317 100644 --- a/tests/unit/integrations/otel/test_otel_v2_dynamic.py +++ b/tests/unit/integrations/otel/test_otel_v2_dynamic.py @@ -1,6 +1,7 @@ """Per-request multi-tenant credential routing (V1 parity).""" import base64 +import logging import pytest from opentelemetry.trace import NoOpTracer @@ -677,3 +678,68 @@ def test_newrelic_key_only_team_routes_to_us_not_operator_region(monkeypatch): ) owned = next(e for e in new_cfg.exporters if e.owner == "newrelic") assert owned.endpoint == "https://otlp.nr-data.net" + + +def test_signoz_dynamic_headers_stamp_ingestion_key(): + from litellm.integrations.otel.presets import dynamic_otlp_headers + + assert dynamic_otlp_headers("signoz", {"signoz_ingestion_key": "team-key"}) == {"signoz-ingestion-key": "team-key"} + # No key means no per-request routing; the caller keeps its default tracer. + assert dynamic_otlp_headers("signoz", {}) is None + + +def test_signoz_dynamic_endpoint_comes_from_team_config_when_its_host_is_allowlisted(monkeypatch): + import litellm + from litellm.integrations.otel.presets import dynamic_otlp_endpoint + + monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", ["ingest.eu.signoz.cloud"]) + assert ( + dynamic_otlp_endpoint( + "signoz", {"signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443", "signoz_ingestion_key": "k"} + ) + == "https://ingest.eu.signoz.cloud:443" + ) + # A team that saved only a key keeps the operator's configured endpoint. + assert dynamic_otlp_endpoint("signoz", {"signoz_ingestion_key": "k"}) is None + assert dynamic_otlp_endpoint("signoz", {}) is None + + +def test_signoz_team_endpoint_off_the_allowlist_is_dropped_along_with_its_key(monkeypatch): + import litellm + from litellm.integrations.otel.presets import dynamic_otlp_endpoint, dynamic_otlp_headers + + monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", []) + params = {"signoz_ingestion_endpoint": "http://169.254.169.254/v1/traces", "signoz_ingestion_key": "k"} + assert dynamic_otlp_endpoint("signoz", params) is None + # The tenant key must not ride to the operator's collector either: the request keeps the default tracer. + assert dynamic_otlp_headers("signoz", params) is None + + +def test_signoz_keyless_team_endpoint_is_ignored_so_the_operator_key_never_reaches_it(monkeypatch, caplog): + import litellm + from litellm.integrations.otel.presets import dynamic_otlp_endpoint, dynamic_otlp_headers + from litellm.integrations.otel.presets.signoz import _warn_endpoint_without_key + + monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", ["collector.team.internal"]) + params = {"signoz_ingestion_endpoint": "http://collector.team.internal:4318"} + _warn_endpoint_without_key.cache_clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert dynamic_otlp_headers("signoz", params) is None + assert "Set signoz_ingestion_key alongside it" in caplog.text + assert dynamic_otlp_endpoint("signoz", params) is None + cache = _cache( + "signoz", + exporters=[ + ExporterSpec( + kind="otlp_http", + endpoint="https://ingest.us.signoz.cloud:443", + headers="signoz-ingestion-key=OPERATOR", + owner="signoz", + requires_headers=True, + ) + ], + ) + routed = cache._routed_config({}, {}, dynamic_otlp_endpoint("signoz", params), "team-service") + owned = next(e for e in routed.exporters if e.owner == "signoz") + assert owned.endpoint == "https://ingest.us.signoz.cloud:443" + assert owned.headers == "signoz-ingestion-key=OPERATOR" diff --git a/tests/unit/integrations/otel/test_otel_v2_presets.py b/tests/unit/integrations/otel/test_otel_v2_presets.py index 58cfc1ceb3f..a060cdf3648 100644 --- a/tests/unit/integrations/otel/test_otel_v2_presets.py +++ b/tests/unit/integrations/otel/test_otel_v2_presets.py @@ -212,3 +212,69 @@ def test_newrelic_preset_unset_content_knob_keeps_default(monkeypatch): from litellm.integrations.otel.presets.newrelic import newrelic_preset assert newrelic_preset().capture_span_content is False + + +def test_signoz_preset_reads_env_endpoint_and_key(monkeypatch): + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "https://ingest.eu.signoz.cloud:443") + monkeypatch.setenv("SIGNOZ_INGESTION_KEY", "env-ingestion-key") + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.kind == "otlp_http" + assert spec.endpoint == "https://ingest.eu.signoz.cloud:443" + assert spec.headers == "signoz-ingestion-key=env-ingestion-key" + assert spec.requires_headers is True + assert "genai" in cfg.mapper_names + + +def test_signoz_preset_without_key_is_self_hosted(monkeypatch): + # A self-hosted collector accepts unauthenticated OTLP, so requiring headers + # would drop exports that would have succeeded. + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "http://signoz-collector.internal:4318") + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint == "http://signoz-collector.internal:4318" + assert spec.headers is None + assert spec.requires_headers is False + + +def test_signoz_preset_has_no_default_endpoint(monkeypatch): + # No region table and no default host: the preset never invents a destination. + monkeypatch.delenv("SIGNOZ_INGESTION_ENDPOINT", raising=False) + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint is None + + +def test_signoz_preset_endpoint_passed_through_verbatim(monkeypatch): + # The plumbing appends the signal path, so pre-appending would double it. + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "https://ingest.us.signoz.cloud:443/v1/traces") + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.plumbing.providers import _otlp_traces_endpoint + from litellm.integrations.otel.presets.signoz import signoz_preset + + cfg = signoz_preset() + spec = next(e for e in cfg.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint == "https://ingest.us.signoz.cloud:443/v1/traces" + assert _otlp_traces_endpoint(spec.endpoint) == "https://ingest.us.signoz.cloud:443/v1/traces" + + +def test_signoz_preset_accepts_the_factory_call_shape(monkeypatch): + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "http://127.0.0.1:1") + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + from litellm.integrations.otel.model.config import ExporterOwner + from litellm.integrations.otel.presets import PRESET_BY_CALLBACK + + cfg = PRESET_BY_CALLBACK["signoz"](allow_missing_credentials=True) + assert any(e.owner == ExporterOwner.SIGNOZ for e in cfg.exporters) diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index c8b02ebc790..2fc747e1b48 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -8932,3 +8932,100 @@ async def test_prompt_management_with_unchanged_variables_replays_a_byte_identic assert json.dumps(messages_n_plus_one[: len(messages_n)], sort_keys=True) == json.dumps(messages_n, sort_keys=True) assert messages_n[0] == {"role": "system", "content": "You are a pirate. Answer in one sentence."} assert len(messages_n_plus_one) == len(messages_n) + 2 + + +def test_signoz_dispatch_prefers_otel_v2_when_flag_on(monkeypatch): + from litellm.integrations.otel.logger import OpenTelemetryV2 + from litellm.integrations.otel.model.config import ExporterOwner, is_otel_v2_enabled + from litellm.litellm_core_utils import litellm_logging as logging_module + + logging_module._in_memory_loggers.clear() + monkeypatch.setenv("LITELLM_OTEL_V2", "true") + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "https://ingest.eu.signoz.cloud:443") + monkeypatch.setenv("SIGNOZ_INGESTION_KEY", "test-key") + is_otel_v2_enabled.cache_clear() + try: + v2_logger = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert isinstance(v2_logger, OpenTelemetryV2) + assert v2_logger.callback_name == "signoz" + spec = next(e for e in v2_logger.config.exporters if e.owner == ExporterOwner.SIGNOZ) + assert spec.endpoint == "https://ingest.eu.signoz.cloud:443" + assert spec.headers == "signoz-ingestion-key=test-key" + again = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert again is v2_logger + finally: + logging_module._in_memory_loggers.clear() + monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) + is_otel_v2_enabled.cache_clear() + + +def test_signoz_dispatch_keeps_legacy_otel_when_flag_off(monkeypatch): + from litellm.integrations.opentelemetry import OpenTelemetry + from litellm.integrations.otel.model.config import is_otel_v2_enabled + from litellm.litellm_core_utils import litellm_logging as logging_module + + logging_module._in_memory_loggers.clear() + monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) + monkeypatch.setenv("SIGNOZ_INGESTION_ENDPOINT", "http://signoz-collector.internal:4318") + monkeypatch.setenv("SIGNOZ_INGESTION_KEY", "legacy-key") + monkeypatch.delenv("OTEL_EXPORTER_OTLP_TRACES_HEADERS", raising=False) + is_otel_v2_enabled.cache_clear() + try: + legacy = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert isinstance(legacy, OpenTelemetry) + assert legacy.callback_name == "signoz" + assert legacy.config.endpoint == "http://signoz-collector.internal:4318/v1/traces" + assert legacy.config.headers == "signoz-ingestion-key=legacy-key" + assert "OTEL_EXPORTER_OTLP_TRACES_HEADERS" not in os.environ + # Same name resolves to the same instance, not a second exporter. + again = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert again is legacy + finally: + logging_module._in_memory_loggers.clear() + is_otel_v2_enabled.cache_clear() + + +def test_signoz_dispatch_requires_an_endpoint(monkeypatch): + from litellm.integrations.otel.model.config import is_otel_v2_enabled + from litellm.litellm_core_utils import litellm_logging as logging_module + + logging_module._in_memory_loggers.clear() + monkeypatch.setenv("LITELLM_OTEL_V2", "true") + monkeypatch.delenv("SIGNOZ_INGESTION_ENDPOINT", raising=False) + monkeypatch.delenv("SIGNOZ_INGESTION_KEY", raising=False) + is_otel_v2_enabled.cache_clear() + try: + created = logging_module._init_custom_logger_compatible_class( + logging_integration="signoz", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert created is None + assert not [ + cb for cb in logging_module._in_memory_loggers if getattr(cb, "callback_name", None) == "signoz" + ] + finally: + logging_module._in_memory_loggers.clear() + monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) + is_otel_v2_enabled.cache_clear() diff --git a/ui/litellm-dashboard/public/assets/logos/signoz.svg b/ui/litellm-dashboard/public/assets/logos/signoz.svg new file mode 100644 index 00000000000..9064cb86bd6 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/signoz.svg @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx index bc9889da724..19020d92066 100644 --- a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx @@ -10,6 +10,7 @@ import newrelicLogo from "../../public/assets/logos/newrelic.png"; import openmeterLogo from "../../public/assets/logos/openmeter.png"; import otelLogo from "../../public/assets/logos/otel.png"; import pointfiveLogo from "../../public/assets/logos/pointfive.png"; +import signozLogo from "../../public/assets/logos/signoz.svg"; import databricksLogo from "../../public/assets/logos/databricks.svg"; interface CallbackConfig { @@ -209,6 +210,17 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ }, description: "S3 Bucket (AWS) Logging Integration", }, + { + id: "signoz", + displayName: "SigNoz", + logo: signozLogo.src, + supports_key_team_logging: true, + dynamic_params: { + signoz_ingestion_endpoint: "text", + signoz_ingestion_key: "password", + }, + description: "SigNoz Logging Integration. Setup: https://signoz.io/docs/litellm-observability/", + }, { id: "SQS", displayName: "SQS", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 6d45f6e691c..2b8ed9aa58d 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -57081,7 +57081,7 @@ export interface operations { parameters: { query: { /** @description Specify the service being hit. */ - service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "ms_teams" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "galileo" | "newrelic" | "pointfive" | "sqs") | string; + service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "ms_teams" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "galileo" | "newrelic" | "pointfive" | "signoz" | "sqs") | string; }; header?: never; path?: never; From 013d5fa0150cd010d82eaa84c0007ec6a8dda296 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 26 Sep 2026 18:46:36 -0700 Subject: [PATCH 16/51] feat(cli): reuse saved agent setup and add reconfigure (#43392) --- litellm/proxy/client/cli/README.md | 21 +- .../proxy/client/cli/commands/configure.py | 561 +++++++---------- .../client/cli/commands/configure_profiles.py | 158 +++++ .../client/cli/commands/configure_setup.py | 439 ++++++++++++++ litellm/proxy/client/cli/main.py | 8 +- .../client/cli/test_configure_commands.py | 566 +++++++++++++++++- 6 files changed, 1391 insertions(+), 362 deletions(-) create mode 100644 litellm/proxy/client/cli/commands/configure_profiles.py create mode 100644 litellm/proxy/client/cli/commands/configure_setup.py diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index a02d7cce0d8..8c05264c6e6 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -546,7 +546,20 @@ lite configure --api-key sk-... --gateway-url https://your-proxy.example.com Select Claude Code, Codex, or both, then choose a gateway model for each selected agent. The wizard validates the key and reads the models your key can access before changing settings. Start either configured agent normally with `claude` or `codex`; the gateway connection persists across terminals without a wrapper or exported API key -`--gateway-url` also accepts a deployment path prefix and a trailing `/v1`. `--base-url` is an alias. If omitted, setup uses `lite --base-url`, `LITELLM_PROXY_URL`, or the saved CLI URL; the wizard asks for a URL when none was provided +Your gateway, virtual key and model choice are saved separately from each agent's undo record. The command saves validated choices before applying them; if applying fails, `lite configure` retries those saved choices. Disconnecting keeps that setup so you can reconnect without repeating the wizard: + +```bash +lite unconfigure +lite configure +``` + +`lite configure` reuses all saved setups for the current agent homes. Name an agent to reconnect only that one, such as `lite configure claude` or `lite configure codex`. Saved setup works without a terminal when the key and model are still valid + +Run `lite reconfigure` to edit your choices with the saved values prefilled, or `lite reconfigure codex` to edit one agent. The agent picker selects which setups to edit; unchecked agents keep their settings. For Claude Code, choose its own default in the wizard or use `lite configure claude --default-model` to remove LiteLLM's model pin. Omitting `--model` keeps your saved choice + +`lite unconfigure --forget` undoes settings it still owns and deletes the saved setups, including their saved keys. `lite unconfigure claude --forget` forgets only Claude Code. Both work after an earlier disconnect. Saved setup files have owner-only permissions and follow the same resolved config-file scope as the undo records, including `CLAUDE_CONFIG_DIR` and `CODEX_HOME`. A pending undo record remains available if an original credential could not safely be restored. If the undo record is missing, agent settings are left unchanged and the command asks you to remove any remaining gateway connection and key manually; forgetting the saved profile does not erase unowned agent settings + +`--gateway-url` also accepts a deployment path prefix and a trailing `/v1`. `--base-url` is an alias. Current command-line or environment options override saved setup values. Otherwise an existing setup supplies its own gateway; first setup falls back to the saved CLI URL or prompts for one. Changing the gateway requires a key for that gateway, so an old saved key is never reused for a different destination For a scripted setup, name the agent and model: @@ -569,9 +582,9 @@ lite --base-url https://your-proxy.example.com configure claude --api-key sk-... claude ``` -The key comes from `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`) and is written into `env.ANTHROPIC_AUTH_TOKEN`; without one the command refuses, since a `lite login` credential expires within a day and keeping it fresh would mean Claude Code running `lite` through `apiKeyHelper` on every credential refresh. The command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (under `claude-router-` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) and you pick between them as usual. Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key and as `env.ANTHROPIC_MODEL`, both of which have to be on `/v1/models` for the key. The second one matters for `claude -c` and `claude --resume`: a resumed session otherwise re-sends the model its transcript recorded, which behind an auto-router with `return_raw_model_name: true` is the tier model that answered, and a key scoped to the router alias gets a 403 for it; `ANTHROPIC_MODEL` outranks the transcript on resume. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute start` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control +The key comes from `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`), or from this agent's saved setup, and is written into `env.ANTHROPIC_AUTH_TOKEN`; without one the command refuses, since a `lite login` credential expires within a day and keeping it fresh would mean Claude Code running `lite` through `apiKeyHelper` on every credential refresh. The command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (under `claude-router-` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) and you pick between them as usual. On first setup without a model choice, Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key and as `env.ANTHROPIC_MODEL`, both of which have to be on `/v1/models` for the key. The second one matters for `claude -c` and `claude --resume`: a resumed session otherwise re-sends the model its transcript recorded, which behind an auto-router with `return_raw_model_name: true` is the tier model that answered, and a key scoped to the router alias gets a 403 for it; `ANTHROPIC_MODEL` outranks the transcript on resume. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute start` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control -Plain `lite configure`, with no agent named, asks which agents to wire and which gateway model each starts on, picked from `/v1/models` with a type-to-filter prompt. All choices and selected config files are checked before the first settings write. If a later filesystem write fails, the output identifies each agent already configured and its undo command +On first use, plain `lite configure` asks which agents to wire and which gateway model each starts on, picked from `/v1/models` with a type-to-filter prompt. Later runs validate and reapply the saved setups. All choices and selected config files are checked before the first settings write. If a later filesystem write fails, the output identifies each agent already configured and its undo command What the command changed is recorded in `~/.litellm/claude_configure_state.json` (previous values plus fingerprints of what was written, never a second copy of the key). `lite unconfigure claude` restores each of those keys only if it still holds what `configure` wrote, so anything you changed since is left alone and named in the output; a `settings.json` or `env` object that only existed because of `configure` is removed again. Ownership moves only by a write: running `configure` again (a re-login is one) refreshes the record only for the keys its merge changed, keeps the original snapshot of a key that still holds what it wrote, and snapshots afresh a key you changed in between, so `unconfigure` brings back whatever the repeat displaced and never adopts your edit as its own. A credential (`env.ANTHROPIC_API_KEY`, `env.ANTHROPIC_AUTH_TOKEN`, `apiKeyHelper`) is put back only when the restored file points at the `ANTHROPIC_BASE_URL` it was captured next to; otherwise it stays removed, the output says which server it belonged to, and the receipt is kept so pointing the URL back and running `unconfigure` again finishes the job. It also undoes `lite login --config-claude`, which writes through the same path. Both refuse to run while a `lite up` or `lite autoroute start` session holds a backup, and that check comes before any request @@ -587,7 +600,7 @@ Claude Opus 5 ██████████████████████ After the first response, the status line uses the latest routed model recorded by `GET /auto_router/session?session_id=...`, so it can show the tier model even when the transcript contains the router alias. If no session record is available, it falls back to Claude Code's transcript. Session records and costs are cached for five seconds under a per-user `$TMPDIR/litellm-statusline-` directory. The gateway records turns asynchronously, so the display can briefly lag a completed turn. Any virtual key may read its own sessions. The baseline is the priciest model in the router's hardest tier, the same counterfactual the auto-router's savings reports use. `lite unconfigure claude` removes the `statusLine` entry only while it still points at that script -After upgrading the CLI, rerun your original `lite configure claude` command with the same gateway, key and model choice to refresh `~/.litellm/statusline.py`. Keep any explicit `--model` value: omitting it removes the earlier model pin. Package upgrades alone do not refresh this installed copy +After upgrading the CLI, run `lite configure claude` to refresh `~/.litellm/statusline.py` using the saved setup. If your setup predates saved profiles, supply the original gateway, key and model once. Package upgrades alone do not refresh this installed copy `lite codex` registers the same script as a Codex `Stop` hook for the launch, so after each turn Codex prints the same block as a system message. Codex asks once to trust the hook; the answer is remembered for later launches. diff --git a/litellm/proxy/client/cli/commands/configure.py b/litellm/proxy/client/cli/commands/configure.py index eca7ba86496..ea24a644084 100644 --- a/litellm/proxy/client/cli/commands/configure.py +++ b/litellm/proxy/client/cli/commands/configure.py @@ -1,284 +1,43 @@ -"""Persistent Claude Code and Codex gateway configuration.""" +"""Commands for saved Claude Code and Codex gateway setup.""" -import os import sys -from collections.abc import Callable, Sequence -from dataclasses import dataclass from pathlib import Path -from types import MappingProxyType from typing import Final import click from InquirerPy import inquirer -from InquirerPy.base.control import Choice from pydantic import BaseModel -from litellm.proxy.common_utils.model_listing_utils import ( - CLAUDE_CODE_CLIENT, - CLAUDE_CODE_PICKER_PATTERN, - GATEWAY_CLIENT_HEADER, -) - -from .agents import codex_config_path from .auth import CliContextObj from .claude_settings import ( - STARTING_MODEL_ROLE, ClaudeSettingsError, - ModelChoice, - StartOn, - StaticToken, UnconfigureOutcome, - UnpinModel, - claude_settings_path, - configure_claude_settings, - configure_state_path, preflight_claude_settings, settings_file_owners, unconfigure_claude_settings, ) -from .codex_settings import ( - CodexSettingsError, - configure_codex_settings, - preflight_codex_settings, - unconfigure_codex_settings, -) +from .codex_settings import CodexSettingsError, unconfigure_codex_settings from .config import normalize_base_url -from .pi import ListedModel, ListingFailure, PiSyncError, fetch_model_listing - -_LISTED_MODELS_SHOWN: Final = 20 -_CLAUDE_TARGET: Final = "claude" -_CODEX_TARGET: Final = "codex" -_TARGETS: Final = ((_CLAUDE_TARGET, "Claude Code (CLI)"), (_CODEX_TARGET, "Codex (CLI)")) -_KEEP_DEFAULT_MODEL: Final = "Keep Claude Code's own default" -_CLAUDE_CODE_VIEW: Final = MappingProxyType( - {"anthropic-version": "2023-06-01", GATEWAY_CLIENT_HEADER: CLAUDE_CODE_CLIENT} +from .configure_profiles import ( + TARGETS, + Target, + forget_saved_setup, + read_saved_setup, + receipt_path_for, + settings_path_for, + setup_locks, + setup_profile_path, ) -_MODEL_OPTION_HELP: Final = ( - f"Proxy model to set as {STARTING_MODEL_ROLE}. Must be listed on /v1/models for the key; without it, " - "Claude Code keeps its own default and a pin an earlier configure made is let go of. Nothing pins Claude " - "Code's sub-agent or background tiers; `lite autoroute start` is the mode that does." +from .configure_setup import ( + MODEL_OPTION_HELP, + ConnectionSettings, + configure_targets, + interactive_configure, + pick_targets, + resolve_credential, ) -def resolve_credential(ctx: click.Context, api_key: str | None) -> StaticToken: - """The long-lived key written into settings.json: --api-key, `lite --api-key` or LITELLM_PROXY_API_KEY. - - A `lite login` credential is never written: it expires within a day, and keeping it fresh would mean - Claude Code running `lite` through `apiKeyHelper` on every credential refresh. - """ - ctx_obj: Final[CliContextObj] = ctx.obj - explicit: Final = api_key or (None if ctx_obj.get("api_key_from_token_file") else ctx_obj.get("api_key")) - if not explicit: - raise ClaudeSettingsError( - "`lite configure` needs a long-lived virtual key: pass --api-key, `lite --api-key`, or set " - "LITELLM_PROXY_API_KEY. Your `lite login` credential expires within a day, so it is not written " - "into agent settings." - ) - if not explicit.strip() or any(ord(char) <= 32 or ord(char) == 127 for char in explicit): - raise ClaudeSettingsError("The virtual key must not be blank or contain whitespace or control characters.") - return StaticToken(explicit) - - -@dataclass(frozen=True, slots=True) -class _Listing: - models: tuple[ListedModel, ...] - - @property - def ids(self) -> tuple[str, ...]: - return tuple(model.id for model in self.models) - - -def _preflight(target: str) -> None: - try: - if target == _CLAUDE_TARGET: - preflight_claude_settings(claude_settings_path(os.environ)) - else: - preflight_codex_settings(codex_config_path(os.environ)) - except (ClaudeSettingsError, CodexSettingsError) as e: - raise click.ClickException(str(e)) from e - - -def _start( - ctx: click.Context, base_url: str, api_key: str | None, target: str = _CLAUDE_TARGET -) -> tuple[StaticToken, _Listing]: - _preflight(target) - try: - credential: Final = resolve_credential(ctx, api_key) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) - return credential, _listed_models(base_url, credential.token, target) - - -def _listing_error(base_url: str, error: PiSyncError, target: str) -> str: - """The hint that fits how the listing failed: only an unreachable proxy gets the "is it running" question.""" - if error.kind is ListingFailure.REJECTED: - return f"LiteLLM rejected your key (HTTP {error.status}). Pass a valid --api-key." - if error.kind is ListingFailure.UNREACHABLE: - return ( - f"Could not connect. Is the proxy at {base_url} running, and is --base-url (or LITELLM_PROXY_URL) correct?" - ) - if error.kind is ListingFailure.EMPTY: - name: Final = "Claude Code" if target == _CLAUDE_TARGET else "Codex" - return f"{error.message} {name} would have nothing to run; give the key access to at least one model." - return f"The proxy at {base_url} answered, so check that it is a LiteLLM proxy and is healthy." - - -def _listed_models(base_url: str, key: str, target: str = _CLAUDE_TARGET) -> _Listing: - listed: Final = fetch_model_listing( - base_url, key, headers=_CLAUDE_CODE_VIEW if target == _CLAUDE_TARGET else MappingProxyType({}) - ) - if isinstance(listed, PiSyncError): - raise click.ClickException(_listing_error(base_url, listed, target)) - return _Listing(listed) - - -def _starting_model(model: str, listing: _Listing) -> str | None: - source: Final = next((listed.id for listed in listing.models if listed.source_model == model), None) - return source or next((listed.id for listed in listing.models if listed.id == model), None) - - -def _model_choice(model: str | None) -> ModelChoice: - return StartOn(model) if model is not None else UnpinModel() - - -def _validated_model(model: str | None, listing: _Listing, base_url: str) -> str | None: - starting: Final = _starting_model(model, listing) if model is not None else None - if model is not None and starting is None: - shown: Final = ", ".join(listing.ids[:_LISTED_MODELS_SHOWN]) - raise click.ClickException(f"{model!r} is not served by {base_url} for this key. /v1/models lists: {shown}.") - return starting - - -def _apply_claude(base_url: str, credential: StaticToken, listing: _Listing, model: str | None) -> None: - listed: Final = listing.ids - starting: Final = _validated_model(model, listing, base_url) - settings_path: Final = claude_settings_path(os.environ) - try: - configure_claude_settings( - base_url, - credential, - _model_choice(starting), - settings_path, - configure_state_path(settings_path), - settings_file_owners(settings_path), - ) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) - in_picker: Final = sum(1 for listed_model in listed if CLAUDE_CODE_PICKER_PATTERN.search(listed_model)) - click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.") - - click.echo("Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN.") - click.echo( - f"Starting model: {starting} ({STARTING_MODEL_ROLE}); switch any time with /model." - if starting is not None - else "Starting model: not pinned (Claude Code's default, or a model you set yourself); switch with /model, or " - "pass --model to start on a proxy model. Without a pin, a resumed session re-sends the model its transcript " - "recorded, which behind a raw-model auto-router is the tier model." - ) - click.echo( - f"/model will list all {len(listed)} of the proxy's models." - if in_picker == len(listed) - else f"/model will list {in_picker} of the proxy's {len(listed)} models: Claude Code shows only ids containing " - "'claude' or 'anthropic', and this proxy does not list the rest under such names." - ) - click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.") - if settings_path.is_symlink(): - click.echo( - f"Note: {settings_path} is a symlink to {settings_path.resolve()}, so your key now lives in " - "that file; keep it out of version control.", - err=True, - ) - - -def _pick_targets() -> tuple[str, ...]: - picked: Final = inquirer.checkbox( - message="Which agents should route through LiteLLM?", - choices=[Choice(value, name=label, enabled=True) for value, label in _TARGETS], - validate=lambda chosen: len(chosen) > 0, - invalid_message="Pick at least one.", - ).execute() - return tuple(str(value) for value in picked) - - -def _pick_model(listed: Sequence[str]) -> str | None: - picked: Final = inquirer.fuzzy( - message="Model Claude Code starts on (type to filter; /model switches any time):", - choices=[_KEEP_DEFAULT_MODEL, *listed], - default=listed[0] if listed else _KEEP_DEFAULT_MODEL, - ).execute() - return None if picked == _KEEP_DEFAULT_MODEL else str(picked) - - -def _pick_codex_model(listed: Sequence[str]) -> str: - choices: Final = list(listed) # mutable-ok: InquirerPy's choices parameter requires a list - return str(inquirer.fuzzy(message="Model Codex starts on (type to filter):", choices=choices).execute()) - - -def _apply_codex(base_url: str, credential: StaticToken, listing: _Listing, model: str) -> None: - _validated_model(model, listing, base_url) - settings_path: Final = codex_config_path(os.environ) - try: - configure_codex_settings(base_url, credential.token, model, settings_path) - except CodexSettingsError as e: - raise click.ClickException(str(e)) from e - click.echo(f"Configured Codex: {settings_path} now routes through {base_url}.") - click.echo(f"Starting model: {model}. Credential: your virtual key, stored in the private provider settings.") - click.echo("Start `codex` from any terminal. Undo with `lite unconfigure codex`.") - if settings_path.is_symlink(): - click.echo(f"Note: your key now lives in {settings_path.resolve()}; keep it out of version control.", err=True) - - -@dataclass(frozen=True, slots=True) -class _Setup: - target: str - listing: _Listing - model: str | None - - -def _choose_setup( - base_url: str, - target: str, - credential: StaticToken, - pick_model: Callable[[Sequence[str]], str | None], - pick_codex_model: Callable[[Sequence[str]], str], -) -> _Setup: - listing: Final = _listed_models(base_url, credential.token, target) - model: Final = ( - pick_model(tuple(item.source_model or item.id for item in listing.models)) - if target == _CLAUDE_TARGET - else pick_codex_model(listing.ids) - ) - _validated_model(model, listing, base_url) - return _Setup(target, listing, model) - - -def interactive_configure( - ctx: click.Context, - pick_targets: Callable[[], tuple[str, ...]] = _pick_targets, - pick_model: Callable[[Sequence[str]], str | None] = _pick_model, - pick_codex_model: Callable[[Sequence[str]], str] = _pick_codex_model, -) -> None: - """`lite configure` with no agent named: ask which agents to wire and which model to pin.""" - targets: Final = pick_targets() - if not targets: - return - for target in targets: - _preflight(target) - try: - credential: Final = resolve_credential(ctx, None) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) from e - base_url: Final[str] = ctx.obj["base_url"] - setups: Final = tuple( - _choose_setup(base_url, target, credential, pick_model, pick_codex_model) for target in targets - ) - for setup in setups: - if setup.target == _CLAUDE_TARGET: - _apply_claude(base_url, credential, setup.listing, setup.model) - elif setup.model is not None: - _apply_codex(base_url, credential, setup.listing, setup.model) - - class _ConnectionOptions(BaseModel): api_key: str | None = None gateway_url: str | None = None @@ -286,21 +45,20 @@ class _ConnectionOptions(BaseModel): def _connection_settings(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> CliContextObj: """The context object a subcommand runs with: its own --api-key / --gateway-url over the group's, over `lite`'s.""" - ctx_obj: Final[CliContextObj] = ctx.obj + ctx_obj: Final = ConnectionSettings.model_validate(ctx.find_object(object)) group: Final = ( _ConnectionOptions.model_validate(ctx.parent.params) - if ctx.parent is not None and ctx.parent.command.name == "configure" + if ctx.parent is not None and ctx.parent.command.name in ("configure", "reconfigure") else _ConnectionOptions() ) key: Final = api_key if api_key is not None else group.api_key url: Final = gateway_url if gateway_url is not None else group.gateway_url - normalized: Final = normalize_base_url(url if url is not None else ctx_obj["base_url"]) + normalized: Final = normalize_base_url(url if url is not None else ctx_obj.base_url) connection: Final[CliContextObj] = { - **ctx_obj, "base_url": normalized.removesuffix("/v1"), - "base_url_explicit": url is not None or ctx_obj.get("base_url_explicit", False), - "api_key": key if key is not None else ctx_obj.get("api_key"), - "api_key_from_token_file": False if key is not None else ctx_obj.get("api_key_from_token_file", False), + "base_url_explicit": url is not None or ctx_obj.base_url_explicit, + "api_key": key if key is not None else ctx_obj.api_key, + "api_key_from_token_file": False if key is not None else ctx_obj.api_key_from_token_file, } return connection @@ -309,108 +67,216 @@ def _connection_context(ctx: click.Context, settings: CliContextObj) -> click.Co return click.Context(ctx.command, parent=ctx.parent, obj=settings) -@click.group(name="configure", invoke_without_command=True) -@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to store in the selected agents.") -@click.option( - "--gateway-url", "--base-url", default=None, help="Gateway URL; defaults to `lite --base-url` / LITELLM_PROXY_URL." -) -@click.pass_context -def configure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> None: - """Persistently route a coding agent through your LiteLLM proxy. +def _require_terminal(command: str) -> None: + if sys.stdin.isatty(): + return + raise click.ClickException( + f"`lite {command}` asks questions, so it needs a terminal. Non-interactively, run " + f"`lite {command} claude --api-key --model ` or " + f"`lite {command} codex --api-key --model `" + ) - With no agent named, asks which agents to wire and which proxy model to pin. - """ + +def _configure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None, edit: bool) -> None: if ctx.invoked_subcommand is not None: return settings: Final = _connection_settings(ctx, api_key, gateway_url) connection: Final = _connection_context(ctx, settings) - if not sys.stdin.isatty(): - raise click.ClickException( - "`lite configure` asks questions, so it needs a terminal. Non-interactively, run " - "`lite configure claude --api-key --model ` or " - "`lite configure codex --api-key --model `." + with setup_locks(TARGETS): + saved_targets: Final[tuple[Target, ...]] = tuple( + target for target in TARGETS if read_saved_setup(target) is not None ) - if settings.get("base_url_explicit"): - interactive_configure(connection) - return - prompted: Final = _connection_settings(connection, None, click.prompt("Gateway URL", default=settings["base_url"])) - interactive_configure(_connection_context(connection, prompted)) + if saved_targets and not edit: + configure_targets(connection, saved_targets) + return + _require_terminal("reconfigure" if edit else "configure") + selected: Final = pick_targets(saved_targets or TARGETS, edit=edit) + if not selected: + return + if edit: + configure_targets(connection, selected, interactive=True, edit_connection=True) + return + prompted: Final = ( + settings + if settings.get("base_url_explicit") + else _connection_settings(connection, None, click.prompt("Gateway URL", default=settings["base_url"])) + ) + configure_targets(_connection_context(connection, prompted), selected, interactive=True) -@click.group(name="unconfigure") -def unconfigure_group() -> None: - """Undo `lite configure` for a coding agent.""" +@click.group(name="configure", invoke_without_command=True) +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to save for the selected agents.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") +@click.pass_context +def configure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> None: + """Apply saved setup, or choose agents and models on the first run.""" + _configure_group(ctx, api_key, gateway_url, False) + + +@click.group(name="reconfigure", invoke_without_command=True) +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to save for the selected agents.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") +@click.pass_context +def reconfigure_group(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> None: + """Edit saved gateway, key and model choices, using current choices as defaults.""" + _configure_group(ctx, api_key, gateway_url, True) + + +def _configure_target( + ctx: click.Context, + target: Target, + api_key: str | None, + gateway_url: str | None, + model: str | None, + default_model: bool = False, + *, + edit: bool = False, +) -> None: + settings: Final = _connection_settings(ctx, api_key, gateway_url) + interactive: Final = edit and model is None and not default_model + if interactive: + _require_terminal("reconfigure") + with setup_locks((target,)): + configure_targets( + _connection_context(ctx, settings), + (target,), + model=model, + default_model=default_model, + interactive=interactive, + edit_connection=interactive, + ) @configure_group.command(name="claude") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help=MODEL_OPTION_HELP) @click.option( - "--api-key", - "api_key", - default=None, - help="Long-lived LiteLLM virtual key written into Claude Code's settings. Defaults to the `lite --api-key` / " - "LITELLM_PROXY_API_KEY value; required, since a `lite login` credential expires within a day.", + "--default-model", is_flag=True, help="Stop pinning a starting model; let Claude Code choose its default." ) -@click.option("--model", default=None, help=_MODEL_OPTION_HELP) -@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") @click.pass_context -def configure_claude(ctx: click.Context, api_key: str | None, model: str | None, gateway_url: str | None) -> None: - """Route every Claude Code session through your LiteLLM proxy until `lite unconfigure claude`. - - Patches ~/.claude/settings.json in place: the proxy URL, your virtual key as a static token, - and gateway model discovery so /model lists the proxy's models; --model picks the one Claude - Code starts on and resumes with. Every other - setting is kept, and what changed is recorded so `lite unconfigure claude` can put it back. - Assumes the proxy is already running. - """ - settings: Final = _connection_settings(ctx, api_key, gateway_url) - credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key) - _apply_claude(settings["base_url"], credential, listing, model) +def configure_claude( + ctx: click.Context, + api_key: str | None, + gateway_url: str | None, + model: str | None, + default_model: bool, +) -> None: + """Apply Claude Code's saved setup, or save the supplied settings.""" + _configure_target(ctx, "claude", api_key, gateway_url, model, default_model) @configure_group.command(name="codex") -@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key to store in Codex's user config.") -@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, including any deployment path prefix.") -@click.option("--model", required=True, help="Gateway model Codex starts on, as listed by /v1/models for your key.") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help="Gateway model to start on; required only for first-time setup.") @click.pass_context -def configure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str) -> None: - """Route plain `codex` through the gateway until `lite unconfigure codex`.""" - settings: Final = _connection_settings(ctx, api_key, gateway_url) - credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key, _CODEX_TARGET) - _apply_codex(settings["base_url"], credential, listing, model) +def configure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str | None) -> None: + """Apply Codex's saved setup, or save the supplied settings.""" + _configure_target(ctx, "codex", api_key, gateway_url, model) -@unconfigure_group.command(name="codex") -def unconfigure_codex() -> None: - """Restore only Codex settings still holding what configure wrote.""" - settings_path: Final = codex_config_path(os.environ) +@reconfigure_group.command(name="claude") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help=MODEL_OPTION_HELP) +@click.option( + "--default-model", is_flag=True, help="Stop pinning a starting model; let Claude Code choose its default." +) +@click.pass_context +def reconfigure_claude( + ctx: click.Context, + api_key: str | None, + gateway_url: str | None, + model: str | None, + default_model: bool, +) -> None: + """Edit Claude Code setup, or supply --model / --default-model to apply directly.""" + _configure_target(ctx, "claude", api_key, gateway_url, model, default_model, edit=True) + + +@reconfigure_group.command(name="codex") +@click.option("--api-key", default=None, help="Long-lived LiteLLM virtual key, or reuse the saved key.") +@click.option("--gateway-url", "--base-url", default=None, help="Gateway URL, or reuse the saved gateway.") +@click.option("--model", default=None, help="Gateway model to start on; omit to open the setup wizard.") +@click.pass_context +def reconfigure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str | None) -> None: + """Edit Codex setup, or supply --model to apply directly.""" + _configure_target(ctx, "codex", api_key, gateway_url, model, edit=True) + + +def _disconnect(target: Target, forget: bool) -> None: + settings_path: Final = settings_path_for(target) + state_path: Final = receipt_path_for(target, settings_path) + profile: Final = setup_profile_path(target, settings_path) try: - outcome: Final = unconfigure_codex_settings(settings_path) - except CodexSettingsError as e: - raise click.ClickException(str(e)) from e - if outcome.file_removed: - click.echo(f"Removed {settings_path}; it held only settings created by `lite configure codex`.") - elif outcome.restored: - click.echo(f"Restored in {settings_path}: {', '.join(outcome.restored)}.") - else: - click.echo(f"Nothing in {settings_path} was still ours to restore.") - if outcome.kept: - click.echo(f"Left as you changed them since: {', '.join(outcome.kept)}.") + if not state_path.exists(): + click.echo(f"No {target} undo receipt at {state_path}; nothing to undo. Agent settings were not changed.") + if settings_path.exists(): + click.echo( + f"Cannot confirm disconnection. Check {settings_path} and remove any remaining gateway " + "connection and key manually.", + err=True, + ) + elif target == "claude": + preflight_claude_settings(settings_path) + outcome: Final = unconfigure_claude_settings(settings_path, state_path, settings_file_owners(settings_path)) + _report_unconfigure(settings_path, state_path, outcome) + else: + codex_outcome: Final = unconfigure_codex_settings(settings_path) + if codex_outcome.file_removed: + click.echo(f"Removed {settings_path}; it held only settings created by `lite configure codex`.") + elif codex_outcome.restored: + click.echo(f"Restored in {settings_path}: {', '.join(codex_outcome.restored)}.") + else: + click.echo(f"Nothing in {settings_path} was still ours to restore.") + if codex_outcome.kept: + click.echo(f"Left as you changed them since: {', '.join(codex_outcome.kept)}.") + except (ClaudeSettingsError, CodexSettingsError) as error: + raise click.ClickException(str(error)) from error + if forget: + forget_saved_setup(target) + click.echo(f"Forgot saved {target} setup, including its saved key.") + elif profile.exists(): + click.echo(f"Saved setup retained. Run `lite configure {target}` to apply it again.") + + +@click.group(name="unconfigure", invoke_without_command=True) +@click.option("--forget", is_flag=True, help="Also delete saved setups and their keys.") +@click.pass_context +def unconfigure_group(ctx: click.Context, forget: bool) -> None: + """Disconnect agents while retaining saved setup for `lite configure`.""" + if ctx.invoked_subcommand is not None: + return + with setup_locks(TARGETS): + for target in TARGETS: + _disconnect(target, forget) + + +class _UnconfigureOptions(BaseModel): + forget: bool = False + + +def _unconfigure_target(ctx: click.Context, target: Target, forget: bool) -> None: + parent: Final = _UnconfigureOptions.model_validate(ctx.parent.params) if ctx.parent else _UnconfigureOptions() + with setup_locks((target,)): + _disconnect(target, forget or parent.forget) @unconfigure_group.command(name="claude") -def unconfigure_claude() -> None: - """Return Claude Code's settings to what they were before `lite configure claude`. +@click.option("--forget", is_flag=True, help="Also delete the saved Claude Code setup and key.") +@click.pass_context +def unconfigure_claude(ctx: click.Context, forget: bool) -> None: + """Restore Claude Code settings, including those applied by `lite login --config-claude`.""" + _unconfigure_target(ctx, "claude", forget) - Also undoes `lite login --config-claude`. Only keys still holding what configure wrote are - put back; anything you changed since is left as it is and named in the output. - """ - settings_path: Final = claude_settings_path(os.environ) - state_path: Final = configure_state_path(settings_path) - try: - outcome: Final = unconfigure_claude_settings(settings_path, state_path, settings_file_owners(settings_path)) - except ClaudeSettingsError as e: - raise click.ClickException(str(e)) - _report_unconfigure(settings_path, state_path, outcome) + +@unconfigure_group.command(name="codex") +@click.option("--forget", is_flag=True, help="Also delete the saved Codex setup and key.") +@click.pass_context +def unconfigure_codex(ctx: click.Context, forget: bool) -> None: + """Restore Codex settings still holding what configure wrote.""" + _unconfigure_target(ctx, "codex", forget) def _report_unconfigure(settings_path: Path, state_path: Path, outcome: UnconfigureOutcome) -> None: @@ -434,4 +300,11 @@ def _report_unconfigure(settings_path: Path, state_path: Path, outcome: Unconfig ) -__all__ = ("configure_group", "interactive_configure", "resolve_credential", "unconfigure_group") +__all__ = ( + "configure_group", + "inquirer", + "interactive_configure", + "reconfigure_group", + "resolve_credential", + "unconfigure_group", +) diff --git a/litellm/proxy/client/cli/commands/configure_profiles.py b/litellm/proxy/client/cli/commands/configure_profiles.py new file mode 100644 index 00000000000..87c4a05dd22 --- /dev/null +++ b/litellm/proxy/client/cli/commands/configure_profiles.py @@ -0,0 +1,158 @@ +"""Reusable agent setup, separate from the settings writers' undo receipts.""" + +import hashlib +import os +from collections.abc import Generator, Sequence +from contextlib import ExitStack, contextmanager +from pathlib import Path +from typing import Final, Literal, TypeAlias + +import click +from filelock import FileLock, Timeout +from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator + +from litellm.litellm_core_utils.private_json import ( + commit_staged_json, + discard_staged_json, + ensure_private_dir, + stage_private_json, +) + +from .agents import codex_config_path +from .claude_settings import claude_settings_path, configure_state_path +from .codex_settings import codex_configure_state_path +from .config import normalize_base_url + +Target: TypeAlias = Literal["claude", "codex"] +TARGETS: Final[tuple[Target, ...]] = ("claude", "codex") + + +class SavedSetup(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", strict=True) + + version: Literal[1] = 1 + target: Target + settings_path: str + base_url: str + api_key: str = Field(repr=False) + model: str | None + + @field_validator("base_url") + @classmethod + def normalized_gateway(cls, value: str) -> str: + try: + if normalize_base_url(value).removesuffix("/v1") != value: + raise ValueError("Gateway must be normalized") + except click.UsageError as error: + raise ValueError("Invalid gateway URL") from error + return value + + @field_validator("api_key") + @classmethod + def valid_key(cls, value: str) -> str: + if not value or any(ord(char) <= 32 or ord(char) == 127 for char in value): + raise ValueError("Invalid virtual key") + return value + + @field_validator("model") + @classmethod + def valid_model(cls, value: str | None) -> str | None: + if value is not None and (not value or any(ord(char) < 32 or ord(char) == 127 for char in value)): + raise ValueError("Invalid model choice") + return value + + +def settings_path_for(target: Target) -> Path: + return claude_settings_path(os.environ) if target == "claude" else codex_config_path(os.environ) + + +def receipt_path_for(target: Target, settings_path: Path) -> Path: + return configure_state_path(settings_path) if target == "claude" else codex_configure_state_path(settings_path) + + +def setup_profile_path(target: Target, settings_path: Path) -> Path: + receipt: Final = receipt_path_for(target, settings_path) + return receipt.with_name(f"{receipt.stem}_profile.json") + + +def read_saved_setup(target: Target) -> SavedSetup | None: + settings_path: Final = settings_path_for(target) + path: Final = setup_profile_path(target, settings_path) + try: + payload: Final = path.read_bytes() + except FileNotFoundError: + return None + except OSError as error: + raise click.ClickException( + f"Could not read saved {target} setup at {path}; no settings were changed" + ) from error + try: + saved: Final = SavedSetup.model_validate_json(payload) + if ( + saved.target != target + or saved.settings_path != str(settings_path.resolve()) + or (target == "codex" and saved.model is None) + ): + raise ValueError("Invalid saved setup") + return saved + except (ValidationError, ValueError, click.UsageError) as error: + raise click.ClickException( + f"Saved {target} setup at {path} is invalid or unsupported. " + f"Run `lite unconfigure {target} --forget` to discard it; no settings were changed" + ) from error + + +def _lock_path(target: Target) -> Path: + digest: Final = hashlib.sha256(f"{target}:{settings_path_for(target).resolve()}".encode()).hexdigest() + return Path.home() / ".litellm" / "setup-locks" / f"{digest}.lock" + + +@contextmanager +def setup_locks(targets: Sequence[Target]) -> Generator[None, None, None]: + with ExitStack() as stack: + try: + for path in tuple(_lock_path(target) for target in sorted(frozenset(targets))): + ensure_private_dir(path.parent) + stack.enter_context(FileLock(str(path), timeout=10, mode=0o600)) + except (OSError, Timeout) as error: + raise click.ClickException( + "Could not lock agent setup; retry when other configure commands finish" + ) from error + yield + + +def save_setup(saved: SavedSetup) -> None: + path: Final = setup_profile_path(saved.target, settings_path_for(saved.target)) + try: + ensure_private_dir(path.parent) + staged: Final = stage_private_json( + str(path), + { # mutable-ok: private_json serializes with json.dump, which requires a dict + "version": saved.version, + "target": saved.target, + "settings_path": saved.settings_path, + "base_url": saved.base_url, + "api_key": saved.api_key, + "model": saved.model, + }, + ) + except OSError as error: + raise click.ClickException( + f"Could not save {saved.target} setup; no {saved.target} settings were changed" + ) from error + try: + commit_staged_json(staged, str(path)) + except OSError as error: + raise click.ClickException( + f"Could not save {saved.target} setup; no {saved.target} settings were changed" + ) from error + finally: + discard_staged_json(staged) + + +def forget_saved_setup(target: Target) -> None: + path: Final = setup_profile_path(target, settings_path_for(target)) + try: + path.unlink(missing_ok=True) + except OSError as error: + raise click.ClickException(f"Could not remove saved {target} setup at {path}") from error diff --git a/litellm/proxy/client/cli/commands/configure_setup.py b/litellm/proxy/client/cli/commands/configure_setup.py new file mode 100644 index 00000000000..bd07c19dff4 --- /dev/null +++ b/litellm/proxy/client/cli/commands/configure_setup.py @@ -0,0 +1,439 @@ +"""Persistent Claude Code and Codex gateway configuration.""" + +import os +import sys +from collections.abc import Callable, Sequence +from dataclasses import dataclass +from functools import partial +from types import MappingProxyType +from typing import Final + +import click +import requests +from InquirerPy import inquirer +from InquirerPy.base.control import Choice +from pydantic import BaseModel, TypeAdapter, ValidationError + +from litellm.proxy.common_utils.model_listing_utils import ( + CLAUDE_CODE_CLIENT, + CLAUDE_CODE_PICKER_PATTERN, + GATEWAY_CLIENT_HEADER, +) + +from .agents import codex_config_path +from .claude_settings import ( + STARTING_MODEL_ROLE, + ClaudeSettingsError, + ModelChoice, + StartOn, + StaticToken, + UnpinModel, + claude_settings_path, + configure_claude_settings, + configure_state_path, + preflight_claude_settings, + settings_file_owners, +) +from .codex_settings import ( + CodexSettingsError, + configure_codex_settings, + preflight_codex_settings, +) +from .config import normalize_base_url +from .configure_profiles import ( + TARGETS, + SavedSetup, + Target, + read_saved_setup, + save_setup, + settings_path_for, + setup_locks, +) +from .pi import ListedModel, ListingFailure, PiSyncError, fetch_model_listing + +_LISTED_MODELS_SHOWN: Final = 20 +_CLAUDE_TARGET: Final = "claude" +_CODEX_TARGET: Final = "codex" +_TARGETS: Final = ((_CLAUDE_TARGET, "Claude Code (CLI)"), (_CODEX_TARGET, "Codex (CLI)")) +_KEEP_DEFAULT_MODEL: Final = "Keep Claude Code's own default" +_CLAUDE_CODE_VIEW: Final = MappingProxyType( + {"anthropic-version": "2023-06-01", GATEWAY_CLIENT_HEADER: CLAUDE_CODE_CLIENT} +) +MODEL_OPTION_HELP: Final = ( + f"Proxy model to set as {STARTING_MODEL_ROLE}. Must be listed on /v1/models for the key; omission keeps " + "the saved choice. Use --default-model to stop pinning a model. Nothing pins Claude " + "Code's sub-agent or background tiers; `lite autoroute start` is the mode that does." +) +_TARGET_SELECTION: Final = TypeAdapter(tuple[Target, ...]) +_MODEL_SELECTION: Final = TypeAdapter(str) + + +class ConnectionSettings(BaseModel): + base_url: str + base_url_explicit: bool = False + api_key: str | None = None + api_key_from_token_file: bool = False + + +def resolve_credential(ctx: click.Context, api_key: str | None) -> StaticToken: + """The long-lived key written into settings.json: --api-key, `lite --api-key` or LITELLM_PROXY_API_KEY. + + A `lite login` credential is never written: it expires within a day, and keeping it fresh would mean + Claude Code running `lite` through `apiKeyHelper` on every credential refresh. + """ + ctx_obj: Final = ConnectionSettings.model_validate(ctx.find_object(object)) + explicit: Final = api_key if api_key is not None else (None if ctx_obj.api_key_from_token_file else ctx_obj.api_key) + if explicit is None: + raise ClaudeSettingsError( + "`lite configure` needs a long-lived virtual key: pass --api-key, `lite --api-key`, or set " + "LITELLM_PROXY_API_KEY. Your `lite login` credential expires within a day, so it is not written " + "into agent settings." + ) + if not explicit.strip() or any(ord(char) <= 32 or ord(char) == 127 for char in explicit): + raise ClaudeSettingsError("The virtual key must not be blank or contain whitespace or control characters.") + return StaticToken(explicit) + + +@dataclass(frozen=True, slots=True) +class _Listing: + models: tuple[ListedModel, ...] + + @property + def ids(self) -> tuple[str, ...]: + return tuple(model.id for model in self.models) + + +def _preflight(target: Target) -> None: + try: + if target == _CLAUDE_TARGET: + preflight_claude_settings(claude_settings_path(os.environ)) + else: + preflight_codex_settings(codex_config_path(os.environ)) + except (ClaudeSettingsError, CodexSettingsError) as e: + raise click.ClickException(str(e)) from e + + +def _listing_error(base_url: str, error: PiSyncError, target: str) -> str: + """The hint that fits how the listing failed: only an unreachable proxy gets the "is it running" question.""" + if error.kind is ListingFailure.REJECTED: + return f"LiteLLM rejected your key (HTTP {error.status}). Pass a valid --api-key." + if error.kind is ListingFailure.UNREACHABLE: + return ( + f"Could not connect. Is the proxy at {base_url} running, and is --base-url (or LITELLM_PROXY_URL) correct?" + ) + if error.kind is ListingFailure.EMPTY: + name: Final = "Claude Code" if target == _CLAUDE_TARGET else "Codex" + return f"{error.message} {name} would have nothing to run; give the key access to at least one model." + return f"The proxy at {base_url} answered, so check that it is a LiteLLM proxy and is healthy." + + +def _fetch_models(base_url: str, key: str, target: Target) -> tuple[ListedModel, ...] | PiSyncError: + return fetch_model_listing( + base_url, + key, + get=partial(requests.get, allow_redirects=False), + headers=_CLAUDE_CODE_VIEW if target == _CLAUDE_TARGET else MappingProxyType({}), + ) + + +def _connection_listing( + ctx: click.Context, + base_url: str, + credential: StaticToken, + target: Target, + repair: bool, +) -> tuple[StaticToken, _Listing]: + listed: Final = _fetch_models(base_url, credential.token, target) + if not isinstance(listed, PiSyncError): + return credential, _Listing(listed) + if not repair or listed.kind is not ListingFailure.REJECTED: + raise click.ClickException(_listing_error(base_url, listed, target)) + replacement: Final = click.prompt("Replacement virtual key", hide_input=True, show_default=False) + try: + refreshed: Final = resolve_credential(ctx, replacement) + except ClaudeSettingsError as error: + raise click.ClickException(str(error)) from error + retried: Final = _fetch_models(base_url, refreshed.token, target) + if isinstance(retried, PiSyncError): + raise click.ClickException(_listing_error(base_url, retried, target)) + return refreshed, _Listing(retried) + + +def _starting_model(model: str, listing: _Listing) -> str | None: + source: Final = next((listed.id for listed in listing.models if listed.source_model == model), None) + return source or next((listed.id for listed in listing.models if listed.id == model), None) + + +def _model_choice(model: str | None) -> ModelChoice: + return StartOn(model) if model is not None else UnpinModel() + + +def _validated_model(model: str | None, listing: _Listing, base_url: str) -> str | None: + starting: Final = _starting_model(model, listing) if model is not None else None + if model is not None and starting is None: + shown: Final = ", ".join(listing.ids[:_LISTED_MODELS_SHOWN]) + raise click.ClickException(f"{model!r} is not served by {base_url} for this key. /v1/models lists: {shown}.") + return starting + + +def _apply_claude(base_url: str, credential: StaticToken, listing: _Listing, model: str | None) -> None: + listed: Final = listing.ids + starting: Final = _validated_model(model, listing, base_url) + settings_path: Final = claude_settings_path(os.environ) + try: + configure_claude_settings( + base_url, + credential, + _model_choice(starting), + settings_path, + configure_state_path(settings_path), + settings_file_owners(settings_path), + ) + except ClaudeSettingsError as e: + raise click.ClickException(str(e)) + in_picker: Final = sum(1 for listed_model in listed if CLAUDE_CODE_PICKER_PATTERN.search(listed_model)) + click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.") + + click.echo("Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN.") + click.echo( + f"Starting model: {starting} ({STARTING_MODEL_ROLE}); switch any time with /model." + if starting is not None + else "Starting model: not pinned (Claude Code's default, or a model you set yourself); switch with /model, or " + "pass --model to start on a proxy model. Without a pin, a resumed session re-sends the model its transcript " + "recorded, which behind a raw-model auto-router is the tier model." + ) + click.echo( + f"/model will list all {len(listed)} of the proxy's models." + if in_picker == len(listed) + else f"/model will list {in_picker} of the proxy's {len(listed)} models: Claude Code shows only ids containing " + "'claude' or 'anthropic', and this proxy does not list the rest under such names." + ) + click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.") + if settings_path.is_symlink(): + click.echo( + f"Note: {settings_path} is a symlink to {settings_path.resolve()}, so your key now lives in " + "that file; keep it out of version control.", + err=True, + ) + + +def _has_targets(chosen: Sequence[object]) -> bool: + return bool(chosen) + + +def pick_targets(defaults: tuple[Target, ...] = ("claude", "codex"), *, edit: bool = False) -> tuple[Target, ...]: + choices: Final = [ # mutable-ok: InquirerPy requires a list + Choice(value, name=label, enabled=value in defaults) for value, label in _TARGETS + ] + picked: Final = _TARGET_SELECTION.validate_python( + inquirer.checkbox( + message="Which agents should be edited? Unselected agents keep their current setup" + if edit + else "Which agents should route through LiteLLM?", + choices=choices, + validate=_has_targets, + invalid_message="Pick at least one.", + ).execute() + ) + return tuple(target for target in TARGETS if target in picked) + + +def _pick_model(listed: Sequence[str], default: str | None = None) -> str | None: + choices: Final = [_KEEP_DEFAULT_MODEL, *listed] # mutable-ok: InquirerPy requires a list + picked: Final = _MODEL_SELECTION.validate_python( + inquirer.fuzzy( + message="Model Claude Code starts on (type to filter; /model switches any time):", + choices=choices, + default=default if default in listed else _KEEP_DEFAULT_MODEL, + ).execute() + ) + return None if picked == _KEEP_DEFAULT_MODEL else picked + + +def _pick_codex_model(listed: Sequence[str], default: str | None = None) -> str: + choices: Final = list(listed) # mutable-ok: InquirerPy's choices parameter requires a list + return _MODEL_SELECTION.validate_python( + inquirer.fuzzy( + message="Model Codex starts on (type to filter):", + choices=choices, + default=default if default in listed else listed[0], + ).execute() + ) + + +def _apply_codex(base_url: str, credential: StaticToken, listing: _Listing, model: str) -> None: + _validated_model(model, listing, base_url) + settings_path: Final = codex_config_path(os.environ) + try: + configure_codex_settings(base_url, credential.token, model, settings_path) + except CodexSettingsError as e: + raise click.ClickException(str(e)) from e + click.echo(f"Configured Codex: {settings_path} now routes through {base_url}.") + click.echo(f"Starting model: {model}. Credential: your virtual key, stored in the private provider settings.") + click.echo("Start `codex` from any terminal. Undo with `lite unconfigure codex`.") + if settings_path.is_symlink(): + click.echo(f"Note: your key now lives in {settings_path.resolve()}; keep it out of version control.", err=True) + + +@dataclass(frozen=True, slots=True) +class PreparedSetup: + saved: SavedSetup + listing: _Listing + + +def _connection(ctx: click.Context, saved: SavedSetup | None) -> tuple[str, StaticToken]: + settings: Final = ConnectionSettings.model_validate(ctx.find_object(object)) + base_url: Final = settings.base_url if saved is None or settings.base_url_explicit else saved.base_url + supplied_key: Final = None if settings.api_key_from_token_file else settings.api_key + reusable_key: Final = saved.api_key if saved is not None and saved.base_url == base_url else None + try: + credential: Final = resolve_credential(ctx, supplied_key if supplied_key is not None else reusable_key) + except ClaudeSettingsError as error: + if saved is not None and saved.base_url != base_url and supplied_key is None: + raise click.ClickException( + "The gateway changed. Pass --api-key for the new gateway; the saved key was not used" + ) from error + raise click.ClickException(str(error)) from error + return base_url, credential + + +def _prompt_connection(ctx: click.Context, target: Target, saved: SavedSetup | None) -> tuple[str, StaticToken]: + settings: Final = ConnectionSettings.model_validate(ctx.find_object(object)) + default_url: Final = settings.base_url if saved is None or settings.base_url_explicit else saved.base_url + base_url: Final = normalize_base_url( + click.prompt(f"{target.capitalize()} gateway URL", default=default_url) + ).removesuffix("/v1") + supplied_key: Final = None if settings.api_key_from_token_file else settings.api_key + kept_key: Final = ( + supplied_key + if supplied_key is not None + else (saved.api_key if saved is not None and saved.base_url == base_url else None) + ) + entered: Final = click.prompt( + "Virtual key (press Enter to keep the current key)" if kept_key is not None else "Virtual key", + default="" if kept_key is not None else None, + show_default=False, + hide_input=True, + ) + key: Final[str | None] = entered or kept_key + try: + return base_url, resolve_credential(ctx, key) + except ClaudeSettingsError as error: + raise click.ClickException(str(error)) from error + + +def _prepare( + ctx: click.Context, + target: Target, + saved: SavedSetup | None, + model: str | None, + default_model: bool, + *, + interactive: bool = False, + edit_connection: bool = False, + pick_model: Callable[[Sequence[str]], str | None] | None = None, + pick_codex_model: Callable[[Sequence[str]], str] | None = None, +) -> PreparedSetup: + default: Final = None if default_model else (model if model is not None else (saved.model if saved else None)) + if target == "codex" and default is None and not interactive: + raise click.UsageError("Missing option '--model'. First-time Codex setup needs a starting model") + base_url, credential = _prompt_connection(ctx, target, saved) if edit_connection else _connection(ctx, saved) + repair: Final = saved is not None and not interactive and sys.stdin.isatty() + active_credential, listing = _connection_listing(ctx, base_url, credential, target, repair) + repair_model: Final = repair and default is not None and _starting_model(default, listing) is None + source_names: Final = tuple(item.source_model or item.id for item in listing.models) + chosen: Final = ( + (pick_model(source_names) if pick_model is not None else _pick_model(source_names, default)) + if (interactive or repair_model) and target == "claude" + else ( + pick_codex_model(listing.ids) if pick_codex_model is not None else _pick_codex_model(listing.ids, default) + ) + if interactive or repair_model + else default + ) + if target == "codex" and chosen is None: + raise click.ClickException("First-time Codex setup needs --model. Run `lite configure` for the model picker") + validated: Final = _validated_model(chosen, listing, base_url) + saved_model: Final = ( + next(item.source_model or item.id for item in listing.models if item.id == validated) + if target == "claude" and validated is not None + else chosen + ) + try: + profile: Final = SavedSetup( + target=target, + settings_path=str(settings_path_for(target).resolve()), + base_url=base_url, + api_key=active_credential.token, + model=saved_model, + ) + except ValidationError as error: + raise click.ClickException("Invalid gateway setup; no settings were changed") from error + return PreparedSetup(profile, listing) + + +def _apply(setup: PreparedSetup) -> None: + saved: Final = setup.saved + save_setup(saved) + try: + if saved.target == "claude": + _apply_claude(saved.base_url, StaticToken(saved.api_key), setup.listing, saved.model) + elif saved.model is not None: + _apply_codex(saved.base_url, StaticToken(saved.api_key), setup.listing, saved.model) + except click.ClickException as error: + raise click.ClickException( + f"{error.format_message()} {saved.target.capitalize()} setup was saved. " + f"Run `lite configure {saved.target}` to retry applying it" + ) from error + click.echo( + f"Setup saved. Edit with `lite reconfigure {saved.target}`; " + f"remove saved settings and key with `lite unconfigure {saved.target} --forget`." + ) + + +def configure_targets( + ctx: click.Context, + targets: tuple[Target, ...], + *, + model: str | None = None, + default_model: bool = False, + interactive: bool = False, + edit_connection: bool = False, + pick_model: Callable[[Sequence[str]], str | None] | None = None, + pick_codex_model: Callable[[Sequence[str]], str] | None = None, +) -> None: + if model is not None and default_model: + raise click.UsageError("--model and --default-model cannot be used together") + for target in targets: + _preflight(target) + setups: Final = tuple( + _prepare( + ctx, + target, + read_saved_setup(target), + model, + default_model, + interactive=interactive, + edit_connection=edit_connection, + pick_model=pick_model, + pick_codex_model=pick_codex_model, + ) + for target in targets + ) + for setup in setups: + _apply(setup) + + +def interactive_configure( + ctx: click.Context, + pick_targets: Callable[[], tuple[str, ...]] = pick_targets, + pick_model: Callable[[Sequence[str]], str | None] | None = None, + pick_codex_model: Callable[[Sequence[str]], str] | None = None, +) -> None: + """Configure selected agents, retaining injectable pickers for embedders.""" + selected: Final = pick_targets() + targets: Final[tuple[Target, ...]] = tuple(target for target in TARGETS if target in selected) + if not targets: + return + with setup_locks(targets): + configure_targets(ctx, targets, interactive=True, pick_model=pick_model, pick_codex_model=pick_codex_model) diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index 6d63acc7479..682dd61ac42 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -21,7 +21,7 @@ from .commands.auth import ( from .commands.autoroute.commands import autoroute_group from .commands.chat import chat from .commands.config import config_commands, get_config_value, hidden_command_names -from .commands.configure import configure_group, unconfigure_group +from .commands.configure import configure_group, reconfigure_group, unconfigure_group from .commands.credentials import credentials from .commands.debug import debug from .commands.encryption import encryption @@ -103,7 +103,8 @@ def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: s # If no API key provided via flag or environment variable, try to load from saved token. # Pass base_url so we only use the stored key when it was issued for this server. - api_key_from_token_file: Final = api_key is None and ctx.invoked_subcommand not in ("configure", "unconfigure") + setup_command: Final = ctx.invoked_subcommand in ("configure", "reconfigure", "unconfigure") + api_key_from_token_file: Final = api_key is None and not setup_command resolved_api_key: Final = ( get_stored_api_key(expected_base_url=base_url, vault=context_secret_vault(ctx)) if api_key_from_token_file @@ -119,7 +120,7 @@ def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: s # "user said localhost:4000 on purpose" so they can fall back to # whatever server the stored token was actually issued for. A base_url # saved via `lite config set` counts as the user saying it. - ctx.obj["base_url_explicit"] = base_url_provided or bool(stored_base_url) + ctx.obj["base_url_explicit"] = base_url_provided or (bool(stored_base_url) and not setup_command) if show_version: print_version(base_url, resolved_api_key) @@ -174,6 +175,7 @@ cli.add_command(autoroute_group, name="autoroute") cli.add_command(config_commands) # Add configure/unconfigure (persistently wire a coding agent to the proxy with a virtual key) cli.add_command(configure_group) +cli.add_command(reconfigure_group) cli.add_command(unconfigure_group) diff --git a/tests/test_litellm/proxy/client/cli/test_configure_commands.py b/tests/test_litellm/proxy/client/cli/test_configure_commands.py index 8f68bb1320b..d8be80ef560 100644 --- a/tests/test_litellm/proxy/client/cli/test_configure_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_configure_commands.py @@ -5,6 +5,7 @@ import stat import time from pathlib import Path from types import SimpleNamespace +from typing import Final, Literal import click import pytest @@ -12,6 +13,8 @@ import requests import responses import tomlkit from click.testing import CliRunner +from InquirerPy.base.control import Choice +from pydantic import JsonValue, TypeAdapter from litellm.proxy.client.cli import cli from litellm.proxy.client.cli.commands import claude_settings as claude_settings_module @@ -393,7 +396,7 @@ class TestConfigureAgents: ) assert (settings_path.read_bytes(), codex_path.read_bytes()) == before assert not state_path.exists() - assert not (codex_path.parent / ".litellm").exists() + assert not tuple((codex_path.parent / ".litellm").glob("*.json")) @responses.activate def test_both_configs_are_preflighted_before_fetching_models_or_writing( @@ -433,7 +436,7 @@ class TestConfigureAgents: assert VALID_KEY not in str(caught.value) assert len(responses.calls) == 0 assert not paths[0].exists() and not paths[1].exists() - assert not codex_path.exists() and not (codex_path.parent / ".litellm").exists() + assert not codex_path.exists() and not tuple((codex_path.parent / ".litellm").glob("*.json")) @responses.activate def test_claude_only_configuration_does_not_require_codex( @@ -509,7 +512,7 @@ class TestConfigureAgents: assert not paths[0].exists() and not codex_path.exists() @responses.activate - def test_configure_and_unconfigure_do_not_read_a_stored_login( + def test_configure_reconfigure_and_unconfigure_do_not_read_a_stored_login( self, runner, paths, codex_path, tmp_path, secret_vault_factory, fake_codex_version ): _mock_agent_models() @@ -528,12 +531,17 @@ class TestConfigureAgents: obj={"secret_vault": vault}, ) assert configured.exit_code == 0, configured.output + reconfigured: Final = runner.invoke( + cli, ["reconfigure", "codex", "--model", "auto"], obj={"secret_vault": vault} + ) + assert reconfigured.exit_code == 0, reconfigured.output fake_codex_version(None, 0) undone = runner.invoke(cli, ["unconfigure", "codex"], obj={"secret_vault": vault}) assert undone.exit_code == 0, undone.output assert vault.reads == 0 and vault.writes == [] and vault.erases == 0 assert not codex_path.exists() and not paths[0].exists() - assert "Removed" in undone.output and "sk-login" not in missing.output + configured.output + undone.output + assert "Removed" in undone.output + assert "sk-login" not in missing.output + configured.output + reconfigured.output + undone.output class TestUnconfigureClaude: @@ -599,9 +607,14 @@ class TestUnconfigureClaude: assert str(state_path) in result.output and state_path.exists() assert "sk-ant" not in result.output - def test_refuses_while_lite_up_holds_a_backup(self, runner, paths, lite_up_backup): - result = runner.invoke(cli, ["unconfigure", "claude"]) - assert result.exit_code != 0 and "lite down" in result.output + def test_disconnected_unconfigure_does_not_touch_lite_up_backup( + self, runner: CliRunner, paths: tuple[Path, Path], lite_up_backup: Path + ) -> None: + result: Final = runner.invoke(cli, ["unconfigure", "claude"]) + assert result.exit_code == 0, result.output + assert "nothing to undo" in result.output + assert "Agent settings were not changed" in result.output + assert lite_up_backup.read_text() == "{}" @responses.activate def test_a_config_dir_is_configured_and_undone_apart_from_the_default_file( @@ -625,12 +638,12 @@ class TestUnconfigureClaude: assert undone.exit_code == 0, undone.output assert json.loads((work_dir / "settings.json").read_text()) == original assert not default_settings.exists() and not default_state.exists() - assert runner.invoke(cli, ["unconfigure", "claude"]).exit_code != 0, "the receipt is gone with the undo" + assert runner.invoke(cli, ["unconfigure", "claude"]).exit_code == 0 - def test_without_a_receipt_it_fails_loudly(self, runner, paths): + def test_without_a_receipt_it_reports_nothing_to_undo(self, runner, paths): result = runner.invoke(cli, ["unconfigure", "claude"]) - assert result.exit_code != 0 - assert "nothing to undo" in result.output + assert result.exit_code == 0, result.output + assert "nothing to undo" in result.output.lower() class TestClaudeCodeView: @@ -702,3 +715,534 @@ class TestClaudeCodeView: result = _configure(runner, "--api-key", VALID_KEY) assert result.exit_code == 0, result.output assert "/model will list 1 of the proxy's 2 models: Claude Code shows only ids containing" in result.output + + +def _saved_profile_path(target: Literal["claude", "codex"], settings_path: Path) -> Path: + from litellm.proxy.client.cli.commands.configure_profiles import setup_profile_path + + return setup_profile_path(target, settings_path) + + +def _configure_saved_agent(runner: CliRunner, target: Literal["claude", "codex"]) -> None: + result: Final = runner.invoke( + cli, + ["configure", "--gateway-url", PROXY, "--api-key", VALID_KEY, target, "--model", "auto"], + ) + assert result.exit_code == 0, result.output + + +def _agent_document(settings_path: Path) -> dict[str, JsonValue]: + adapter: Final = TypeAdapter(dict[str, JsonValue]) + if settings_path.suffix == ".json": + return adapter.validate_json(settings_path.read_text()) + return adapter.validate_python(tomlkit.parse(settings_path.read_text()).unwrap()) + + +def _prompt_answer(answer: str | tuple[str, ...]) -> SimpleNamespace: + def execute() -> str | tuple[str, ...]: + return answer + + return SimpleNamespace(execute=execute) + + +class TestSavedAgentSetup: + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("resume", [("configure",), None], ids=["all", "target"]) + def test_disconnect_then_configure_reuses_connection_and_model_without_prompts( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + target: Literal["claude", "codex"], + resume: tuple[str, ...] | None, + ) -> None: + _mock_agent_models() + settings_path: Final = paths[0] if target == "claude" else codex_path + _configure_saved_agent(runner, target) + configured: Final = _agent_document(settings_path) + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + assert not settings_path.exists() + + resumed: Final = runner.invoke(cli, list(resume or ("configure", target))) + assert resumed.exit_code == 0, resumed.output + assert _agent_document(settings_path) == configured + assert "lite configure" in undone.output and "saved" in undone.output.lower() + repeated: Final = runner.invoke(cli, ["configure", target]) + assert repeated.exit_code == 0, repeated.output + assert _agent_document(settings_path) == configured + restored: Final = runner.invoke(cli, ["unconfigure", target]) + assert restored.exit_code == 0, restored.output + assert not settings_path.exists() + assert VALID_KEY not in resumed.output + repeated.output + restored.output + + @responses.activate + def test_resume_both_agents_captures_the_settings_changed_while_disconnected( + self, runner: CliRunner, paths: tuple[Path, Path], codex_path: Path + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + codex_url: Final = "https://codex-gateway.test/prefix" + responses.get( + f"{codex_url}/v1/models", + json={"data": [{"id": "auto"}]}, + match=[responses.matchers.header_matcher({"Authorization": "Bearer sk-codex"})], + ) + codex_setup: Final = runner.invoke( + cli, + ["configure", "codex", "--gateway-url", codex_url, "--api-key", "sk-codex", "--model", "auto"], + ) + assert codex_setup.exit_code == 0, codex_setup.output + undone: Final = runner.invoke(cli, ["unconfigure"]) + assert undone.exit_code == 0, undone.output + paths[0].write_text('{"theme": "light", "model": "personal-claude"}') + codex_path.write_text('model = "personal-codex"\napproval_policy = "on-request"\n') + + resumed: Final = runner.invoke(cli, ["configure"]) + assert resumed.exit_code == 0, resumed.output + assert json.loads(paths[0].read_text())["model"] == "claude-router-6175746f" + assert tomlkit.parse(codex_path.read_text())["model"] == "auto" + assert responses.calls[-1].request.url == f"{codex_url}/v1/models" + restored: Final = runner.invoke(cli, ["unconfigure"]) + assert restored.exit_code == 0, restored.output + assert json.loads(paths[0].read_text()) == {"theme": "light", "model": "personal-claude"} + assert tomlkit.parse(codex_path.read_text()) == { + "model": "personal-codex", "approval_policy": "on-request" + } + + @responses.activate + @pytest.mark.parametrize("disconnected", [False, True], ids=["active", "disconnected"]) + @pytest.mark.parametrize( + "forget, forgotten", + [ + (("unconfigure", "--forget", "claude"), ("claude",)), + (("unconfigure", "codex", "--forget"), ("codex",)), + (("unconfigure", "--forget"), ("claude", "codex")), + ], + ids=["group-option-target", "leaf-option", "all"], + ) + def test_forget_removes_only_selected_saved_setups_even_after_disconnect( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + disconnected: bool, + forget: tuple[str, ...], + forgotten: tuple[str, ...], + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + _configure_saved_agent(runner, "codex") + if disconnected: + undone: Final = runner.invoke(cli, ["unconfigure"]) + assert undone.exit_code == 0, undone.output + result: Final = runner.invoke(cli, list(forget)) + assert result.exit_code == 0, result.output + for target, settings_path in (("claude", paths[0]), ("codex", codex_path)): + assert _saved_profile_path(target, settings_path).exists() == (target not in forgotten) + resumed: Final = runner.invoke(cli, ["configure", target]) + assert (resumed.exit_code == 0) == (target not in forgotten), resumed.output + assert settings_path.exists() == (target not in forgotten) + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("source", ["leaf", "global", "environment"]) + def test_saved_key_never_follows_a_gateway_override_without_a_replacement( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + source: str, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + replacement_url: Final = "https://replacement.test/gateway" + if source == "environment": + monkeypatch.setenv("LITELLM_PROXY_URL", replacement_url) + args: Final = ( + ["--base-url", replacement_url, "configure", target] + if source == "global" + else ["configure", target, "--gateway-url", replacement_url] + if source == "leaf" + else ["configure", target] + ) + refused: Final = runner.invoke(cli, args) + assert refused.exit_code != 0, refused.output + assert "--api-key" in refused.output and VALID_KEY not in refused.output + assert len(responses.calls) == 1 + assert not paths[0].exists() and not codex_path.exists() + + responses.get( + f"{replacement_url}/v1/models", + json={"data": [{"id": "auto"}]}, + match=[responses.matchers.header_matcher({"Authorization": "Bearer sk-replacement"})], + ) + replaced: Final = runner.invoke(cli, [*args, "--api-key", "sk-replacement"]) + assert replaced.exit_code == 0, replaced.output + assert len(responses.calls) == 2 + assert responses.calls[-1].request.url == f"{replacement_url}/v1/models" + assert VALID_KEY not in replaced.output and "sk-replacement" not in replaced.output + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + def test_saved_setup_is_private_and_scoped_to_the_resolved_agent_home( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + assert stat.S_IMODE(profile_path.stat().st_mode) == 0o600 + assert stat.S_IMODE(profile_path.parent.stat().st_mode) & 0o077 == 0 + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + alternate_home: Final = tmp_path / f"other-{target}" + alternate_settings: Final = alternate_home / settings_path.name + environment: Final = "CLAUDE_CONFIG_DIR" if target == "claude" else "CODEX_HOME" + monkeypatch.setenv(environment, str(alternate_home)) + missing: Final = runner.invoke(cli, ["configure", target]) + assert missing.exit_code != 0, missing.output + assert not alternate_settings.exists() and len(responses.calls) == 1 + assert profile_path.exists() + stored_url: Final = runner.invoke(cli, ["config", "set", "base_url", "https://other-default.test"]) + assert stored_url.exit_code == 0, stored_url.output + monkeypatch.setenv(environment, str(settings_path.parent)) + resumed: Final = runner.invoke(cli, ["configure", target]) + assert resumed.exit_code == 0, resumed.output + assert settings_path.exists() and not alternate_settings.exists() + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("fault", ["json", "version", "target", "path"]) + def test_invalid_saved_setup_fails_without_network_or_secret_output_and_can_be_forgotten( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + target: Literal["claude", "codex"], + fault: str, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + profile: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + corrupted: Final = ( + "{ " + VALID_KEY + if fault == "json" + else json.dumps({**profile, "version": 999}) + if fault == "version" + else json.dumps({**profile, "target": "codex" if target == "claude" else "claude"}) + if fault == "target" + else json.dumps({**profile, "settings_path": str(settings_path.parent / "another-file")}) + ) + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + profile_path.write_text(corrupted) + failed: Final = runner.invoke(cli, ["configure", target]) + assert failed.exit_code != 0, failed.output + assert "saved" in failed.output.lower() and "--forget" in failed.output + assert VALID_KEY not in failed.output + assert not settings_path.exists() and len(responses.calls) == 1 + forgotten: Final = runner.invoke(cli, ["unconfigure", "--forget", target]) + assert forgotten.exit_code == 0, forgotten.output + assert not profile_path.exists() and not settings_path.exists() + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + def test_reconfigure_prefills_saved_choices_and_changes_only_the_selected_agent( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + _configure_saved_agent(runner, "codex") + untouched: Final = codex_path if target == "claude" else paths[0] + before: Final = untouched.read_bytes() + responses.replace( + responses.GET, f"{PROXY}/v1/models", json={"data": [{"id": "auto"}, {"id": "replacement"}]} + ) + + def checkbox(**kwargs: object) -> SimpleNamespace: + choices: Final = kwargs["choices"] + assert isinstance(choices, list) and len(choices) == 2 + for choice in choices: + assert isinstance(choice, Choice) and choice.enabled + return _prompt_answer((target,)) + + def fuzzy(**kwargs: object) -> SimpleNamespace: + assert kwargs["default"] == "auto" + return _prompt_answer("replacement") + + monkeypatch.setattr(configure_module.inquirer, "checkbox", checkbox) + monkeypatch.setattr(configure_module.inquirer, "fuzzy", fuzzy) + changed: Final = runner.invoke(cli, ["reconfigure"], input=_TerminalInput(b"\n\n")) + assert changed.exit_code == 0, changed.output + assert PROXY in changed.output and VALID_KEY not in changed.output + assert untouched.read_bytes() == before + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + resumed: Final = runner.invoke(cli, ["configure", target]) + assert resumed.exit_code == 0, resumed.output + if target == "claude": + assert json.loads(paths[0].read_text())["model"] == "replacement" + else: + assert tomlkit.parse(codex_path.read_text())["model"] == "replacement" + assert untouched.read_bytes() == before + + @responses.activate + def test_reconfigure_cancel_preserves_every_agents_settings_and_saved_choices( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + _configure_saved_agent(runner, "codex") + files: Final = ( + paths[0], codex_path, _saved_profile_path("claude", paths[0]), _saved_profile_path("codex", codex_path) + ) + before: Final = tuple(path.read_bytes() for path in files) + + def checkbox(**kwargs: object) -> SimpleNamespace: + return _prompt_answer(("claude", "codex")) + + def fuzzy(**kwargs: object) -> SimpleNamespace: + assert tuple(path.read_bytes() for path in files) == before + if "Codex" in str(kwargs["message"]): + raise KeyboardInterrupt() + return _prompt_answer("Keep Claude Code's own default") + + monkeypatch.setattr(configure_module.inquirer, "checkbox", checkbox) + monkeypatch.setattr(configure_module.inquirer, "fuzzy", fuzzy) + cancelled: Final = runner.invoke(cli, ["reconfigure"], input=_TerminalInput(b"\n\n\n\n")) + assert cancelled.exit_code != 0, cancelled.output + assert "Aborted" in cancelled.output + assert tuple(path.read_bytes() for path in files) == before + + @responses.activate + def test_explicit_default_model_unpins_claude_and_remains_the_saved_choice( + self, runner: CliRunner, paths: tuple[Path, Path] + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, "claude") + changed: Final = runner.invoke(cli, ["reconfigure", "claude", "--default-model"]) + assert changed.exit_code == 0, changed.output + assert "model" not in json.loads(paths[0].read_text()) + undone: Final = runner.invoke(cli, ["unconfigure", "claude"]) + assert undone.exit_code == 0, undone.output + resumed: Final = runner.invoke(cli, ["configure", "claude"]) + assert resumed.exit_code == 0, resumed.output + assert "model" not in json.loads(paths[0].read_text()) + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("fault", ["key", "model"]) + def test_terminal_resume_repairs_only_the_rejected_saved_choice( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + fault: str, + ) -> None: + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + before: Final = profile_path.read_bytes() + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + responses.reset() + if fault == "key": + responses.get( + f"{PROXY}/v1/models", status=401, + match=[responses.matchers.header_matcher({"Authorization": f"Bearer {VALID_KEY}"})], + ) + responses.get( + f"{PROXY}/v1/models", + json={"data": [{"id": "auto" if fault == "key" else "replacement"}]}, + match=[responses.matchers.header_matcher({ + "Authorization": "Bearer sk-repaired" if fault == "key" else f"Bearer {VALID_KEY}" + })], + ) + failed: Final = runner.invoke(cli, ["configure", target]) + assert failed.exit_code != 0, failed.output + assert not settings_path.exists() and profile_path.read_bytes() == before + assert len(responses.calls) == 1 + + def checkbox(**kwargs: object) -> SimpleNamespace: + raise AssertionError("Saved resume must not ask which agents to configure") + + def fuzzy(**kwargs: object) -> SimpleNamespace: + assert fault == "model", "A rejected key must not discard the saved model" + return _prompt_answer("replacement") + + monkeypatch.setattr(configure_module.inquirer, "checkbox", checkbox) + monkeypatch.setattr(configure_module.inquirer, "fuzzy", fuzzy) + resumed: Final = runner.invoke( + cli, ["configure"], input=_TerminalInput(b"sk-repaired\n" if fault == "key" else b"") + ) + assert resumed.exit_code == 0, resumed.output + assert "gateway URL" not in resumed.output + assert VALID_KEY not in resumed.output and "sk-repaired" not in resumed.output + assert settings_path.exists() + saved: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + assert saved["api_key"] == ("sk-repaired" if fault == "key" else VALID_KEY) + assert saved["model"] == ("auto" if fault == "key" else "replacement") + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("lost_receipt", [False, True], ids=["malformed-settings", "lost-receipt"]) + def test_forget_without_receipt_preserves_agent_settings_and_reports_unknown_connection( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + target: Literal["claude", "codex"], + lost_receipt: bool, + ) -> None: + from litellm.proxy.client.cli.commands.configure_profiles import receipt_path_for + + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + if lost_receipt: + receipt_path_for(target, settings_path).unlink() + else: + undone: Final = runner.invoke(cli, ["unconfigure", target]) + assert undone.exit_code == 0, undone.output + settings_path.write_text("[invalid") + before: Final = settings_path.read_bytes() + forgotten: Final = runner.invoke(cli, ["unconfigure", target, "--forget"]) + assert forgotten.exit_code == 0, forgotten.output + assert not profile_path.exists() + assert settings_path.read_bytes() == before + assert "Cannot confirm disconnection" in forgotten.output + assert "gateway connection and key manually" in forgotten.output + assert str(settings_path) in forgotten.output + assert "already disconnected" not in forgotten.output and VALID_KEY not in forgotten.output + assert len(responses.calls) == 1 + + @responses.activate + @pytest.mark.parametrize("target", ["claude", "codex"]) + @pytest.mark.parametrize("failure", ["stage_private_json", "commit_staged_json", "apply"]) + def test_failed_setup_write_preserves_saved_intent_and_plain_configure_retries_it( + self, + runner: CliRunner, + paths: tuple[Path, Path], + codex_path: Path, + monkeypatch: pytest.MonkeyPatch, + target: Literal["claude", "codex"], + failure: str, + ) -> None: + from litellm.proxy.client.cli.commands import configure_profiles, configure_setup + + _mock_agent_models() + _configure_saved_agent(runner, target) + settings_path: Final = paths[0] if target == "claude" else codex_path + profile_path: Final = _saved_profile_path(target, settings_path) + receipt_path: Final = configure_profiles.receipt_path_for(target, settings_path) + before: Final = (settings_path.read_bytes(), profile_path.read_bytes(), receipt_path.read_bytes()) + original_settings: Final = _agent_document(settings_path) + original_profile: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + replacement_url: Final = "https://replacement.test/gateway" + replacement_key: Final = "sk-replacement" + responses.get( + f"{replacement_url}/v1/models", + json={"data": [{"id": "replacement"}]}, + match=[responses.matchers.header_matcher({"Authorization": f"Bearer {replacement_key}"})], + ) + + def fail_write(*args: object, **kwargs: object) -> str: + raise OSError(f"simulated disk error {VALID_KEY}") + + def fail_apply(*args: object, **kwargs: object) -> None: + error: Final = ( + configure_setup.ClaudeSettingsError if target == "claude" else configure_setup.CodexSettingsError + ) + raise error("simulated agent settings write failure") + + with monkeypatch.context() as patch: + if failure == "apply": + patch.setattr(configure_setup, f"configure_{target}_settings", fail_apply) + else: + patch.setattr(configure_profiles, failure, fail_write) + failed: Final = runner.invoke( + cli, + [ + "reconfigure", target, "--gateway-url", replacement_url, + "--api-key", replacement_key, "--model", "replacement", + ], + ) + assert failed.exit_code != 0, failed.output + assert VALID_KEY not in failed.output and replacement_key not in failed.output + assert (settings_path.read_bytes(), receipt_path.read_bytes()) == (before[0], before[2]) + saved: Final = TypeAdapter(dict[str, JsonValue]).validate_json(profile_path.read_text()) + if failure == "apply": + assert saved == { + **original_profile, "base_url": replacement_url, "api_key": replacement_key, "model": "replacement" + } + assert "simulated agent settings write failure" in failed.output + assert "setup was saved" in failed.output and f"lite configure {target}" in failed.output + else: + assert "could not save" in failed.output.lower() + assert profile_path.read_bytes() == before[1] + retried: Final = runner.invoke(cli, ["configure", target]) + assert retried.exit_code == 0, retried.output + written: Final = _agent_document(settings_path) + if failure != "apply": + assert written == original_settings + elif target == "claude": + environment: Final = written["env"] + assert isinstance(environment, dict) + assert (environment["ANTHROPIC_BASE_URL"], environment["ANTHROPIC_AUTH_TOKEN"], written["model"]) == ( + replacement_url, replacement_key, "replacement" + ) + else: + providers: Final = written["model_providers"] + assert isinstance(providers, dict) + provider: Final = providers["litellm"] + assert isinstance(provider, dict) + headers: Final = provider["http_headers"] + assert isinstance(headers, dict) + assert (provider["base_url"], headers["Authorization"], written["model"]) == ( + f"{replacement_url}/v1", f"Bearer {replacement_key}", "replacement" + ) + + @responses.activate + def test_contended_setup_lock_blocks_requests_and_agent_writes( + self, runner: CliRunner, paths: tuple[Path, Path] + ) -> None: + from litellm.proxy.client.cli.commands.configure_profiles import setup_locks + + _mock_agent_models() + with setup_locks(("claude",)): + blocked: Final = runner.invoke( + cli, + ["configure", "claude", "--gateway-url", PROXY, "--api-key", VALID_KEY, "--model", "auto"], + ) + assert blocked.exit_code != 0, blocked.output + assert "Could not lock agent setup" in blocked.output + assert len(responses.calls) == 0 + assert not paths[0].exists() and not paths[1].exists() + assert not _saved_profile_path("claude", paths[0]).exists() From be35b22dfc37a9f96852dec3262d78908799d164 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 26 Sep 2026 18:53:12 -0700 Subject: [PATCH 17/51] fix(streaming): keep litellm Usage on text-completion usage chunks (#43047) * fix(streaming): keep litellm Usage on text-completion usage chunks * fix(streaming): convert provider usage to litellm Usage instead of dropping it --- .../litellm_core_utils/streaming_handler.py | 7 +++- .../test_streaming_handler.py | 41 +++++++++++++++++++ 2 files changed, 46 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index fa687b585f5..fa4650aec4f 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1504,8 +1504,11 @@ class CustomStreamWrapper: self.tool_call = True - if hasattr(chunk, "usage") and chunk.usage is not None: - model_response.usage = chunk.usage + chunk_usage: Final = getattr(chunk, "usage", None) + if isinstance(chunk_usage, Usage): + model_response.usage = chunk_usage + elif isinstance(chunk_usage, BaseModel): + model_response.usage = Usage(**chunk_usage.model_dump()) ## RETURN ARG result: Final = self.return_processed_chunk_logic( diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index 6557811b530..62d8b0e203f 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -2859,6 +2859,47 @@ def test_dispatch_text_completion_openai_with_usage( assert model_response.usage.total_tokens == 8 +@pytest.mark.parametrize("custom_llm_provider", ["text-completion-openai", "azure_text"]) +def test_text_completion_usage_chunk_keeps_provider_usage_as_litellm_usage( + initialized_custom_stream_wrapper: CustomStreamWrapper, + custom_llm_provider: str, +): + from openai.types.completion import Completion + from openai.types.completion_usage import CompletionUsage + + initialized_custom_stream_wrapper.custom_llm_provider = custom_llm_provider + initialized_custom_stream_wrapper.model = "gpt-3.5-turbo-instruct" + initialized_custom_stream_wrapper.send_stream_usage = True + initialized_custom_stream_wrapper.received_finish_reason = "length" + provider_usage: Final = CompletionUsage.model_validate( + { + "prompt_tokens": 7, + "completion_tokens": 4, + "total_tokens": 11, + "completion_tokens_details": {"reasoning_tokens": 3}, + "prompt_tokens_details": {"cached_tokens": 2}, + "cost": 0.0123, + } + ) + chunk: Final = Completion.model_construct( + id="cmpl-usage", + choices=[], + created=1, + model="gpt-3.5-turbo-instruct", + object="text_completion", + usage=provider_usage, + ) + + returned: Final = initialized_custom_stream_wrapper.chunk_creator(chunk=chunk) + + assert isinstance(returned.usage, Usage) + dumped: Final = returned.model_dump()["usage"] + assert (dumped["prompt_tokens"], dumped["completion_tokens"], dumped["total_tokens"]) == (7, 4, 11) + assert dumped["cost"] == provider_usage.model_dump()["cost"] + assert dumped["completion_tokens_details"]["reasoning_tokens"] == 3 + assert dumped["prompt_tokens_details"]["cached_tokens"] == 2 + + @pytest.mark.asyncio async def test_custom_stream_wrapper_anext_does_not_block_event_loop_for_sync_iterators( logging_obj: Logging, From 9ba552d527883d7dad9778e63b73422dccbfbaf4 Mon Sep 17 00:00:00 2001 From: agustin18 Date: Sat, 26 Sep 2026 23:51:45 -0300 Subject: [PATCH 18/51] fix(vertex_ai): consider tools when validating context caching min tokens (#43319) * fix(vertex_ai): consider tools when validating context caching min tokens Pass tools to is_prompt_caching_valid_prompt in both sync and async check_and_create_cache before popping them into the cachedContents request body. This allows agent-shaped requests with heavy tool schemas and small message histories to reach the minimum token threshold and benefit from prompt caching. Fixes #42804 * test(vertex_ai): avoid doubles on internal code and assert tools in cache payload --- .../vertex_ai_context_caching.py | 2 + .../test_vertex_ai_context_caching.py | 123 ++++++++++++++++++ 2 files changed, 125 insertions(+) diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 75d4ffbed86..2d35dd9b480 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -322,6 +322,7 @@ class ContextCachingEndpoints(VertexBase): if not is_prompt_caching_valid_prompt( model=model, messages=cached_messages, + tools=optional_params.get("tools"), custom_llm_provider=custom_llm_provider, ): verbose_logger.debug( @@ -481,6 +482,7 @@ class ContextCachingEndpoints(VertexBase): if not is_prompt_caching_valid_prompt( model=model, messages=cached_messages, + tools=optional_params.get("tools"), custom_llm_provider=custom_llm_provider, ): verbose_logger.debug( diff --git a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 7913700c8a7..283ed3710d0 100644 --- a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1503,6 +1503,129 @@ class TestContextCachingEndpoints: # Restart the patcher so teardown_method can stop it cleanly self._token_check_patcher.start() + @pytest.mark.parametrize("is_async", [False, True]) + @pytest.mark.parametrize( + "custom_llm_provider", ["gemini", "vertex_ai"] + ) + @pytest.mark.asyncio + async def test_check_and_create_cache_considers_tools_for_min_tokens( + self, custom_llm_provider, is_async + ): + """Test that context caching accounts for tools when validating minimum token count. + + Fixes #42804: When messages alone are below the threshold, but tools push the total + over the minimum token count, context caching must proceed and include tools. + """ + self._token_check_patcher.stop() + + short_cached_messages = [ + { + "role": "system", + "content": "Short system instruction.", + "cache_control": {"type": "ephemeral"}, + } + ] + non_cached_messages = [ + {"role": "user", "content": "Hello world"}, + ] + all_messages = short_cached_messages + non_cached_messages + + large_tools = [ + { + "type": "function", + "function": { + "name": f"synthetic_tool_{i}", + "description": "A very descriptive explanation of a synthetic tool designed to add tokens to the prompt cache prefix " * 8, + "parameters": { + "type": "object", + "properties": { + f"arg_{j}": {"type": "string", "description": "Argument description for caching verification " * 4} + for j in range(10) + }, + "required": [f"arg_{j}" for j in range(5)], + }, + }, + } + for i in range(12) + ] + + optional_params = { + **self.sample_optional_params, + "tools": large_tools, + } + + mock_response = MagicMock() + mock_response.json.return_value = { + "name": "cachedContents/test_cache_id", + "model": "gemini-1.5-pro", + } + mock_response.status_code = 200 + self.mock_client.post.return_value = mock_response + self.mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch.object( + self.context_caching, + "_get_token_and_url_context_caching", + return_value=("fake_token", "https://fake.url/cachedContents"), + ), patch.object( + self.context_caching, + "check_cache", + return_value=None, + ), patch.object( + self.context_caching, + "async_check_cache", + new_callable=AsyncMock, + return_value=None, + ): + if is_async: + result = await self.context_caching.async_check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-1.5-pro", + client=self.mock_async_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider=custom_llm_provider, + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + else: + result = self.context_caching.check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-1.5-pro", + client=self.mock_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider=custom_llm_provider, + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == non_cached_messages + assert returned_cache == "cachedContents/test_cache_id" + assert "tools" not in returned_params + + post_mock = self.mock_async_client.post if is_async else self.mock_client.post + post_mock.assert_called_once() + call_kwargs = post_mock.call_args.kwargs + assert call_kwargs["json"]["tools"] == large_tools + assert call_kwargs["json"]["contents"] == [ + {"role": "user", "parts": [{"text": "Short system instruction."}]} + ] + + self._token_check_patcher.start() + + def _model_turn_final_messages(self, final_cached_role): tool_call = { "id": "call_abc123", From ba2d1c2785255f4143c26f35b18c37693e44d4eb Mon Sep 17 00:00:00 2001 From: Jeremy Schoemaker Date: Sat, 26 Sep 2026 22:17:44 -0500 Subject: [PATCH 19/51] =?UTF-8?q?fix(anthropic):=20drop=20thinking=20block?= =?UTF-8?q?s=20with=20empty=20thinking=20text,=20not=20just=20missing=20si?= =?UTF-8?q?gnature=20=F0=9F=A7=A0=F0=9F=9A=AB=20(#38049)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit _is_unsignable_thinking_block() only checked block["signature"], so a thinking block with a valid-looking signature but empty (or whitespace-only) thinking text sailed through _drop_unsignable_thinking_blocks and into anthropic_messages_pt(). Anthropic rejects that with: 400 messages.N.content.M.thinking: each thinking block must contain thinking This is reachable whenever a thinking_blocks history item gets replayed through this Anthropic-shaped request path (e.g. a non-Anthropic reasoning turn with no summary text), the same class of bug PR #36033 fixed on the Responses adapter's own separate code path. Now the signature check runs first (unsigned blocks are still dropped, same as before), then an additional check drops the block if `thinking` is missing, not a string, or strips to empty. redacted_thinking blocks are untouched since they don't have type == "thinking". --- .../prompt_templates/common_utils.py | 18 +- ...llm_core_utils_prompt_templates_factory.py | 176 ++++++++++++++++++ 2 files changed, 188 insertions(+), 6 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 14d47a15c6d..41563d501a9 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1992,11 +1992,14 @@ def is_encrypted_reasoning_block(block: object) -> bool: def is_unsignable_thinking_block(block: object) -> bool: """A thinking block Anthropic cannot accept on input. - Anthropic verifies the thinking signature cryptographically, so a block whose - signature is null, empty, or missing (e.g. from an open-source reasoning model) - is rejected with a 400 and must be dropped rather than blanked or repaired, and - so is a block whose signature or data carries another provider's encrypted - reasoning. A `redacted_thinking` block Anthropic minted is always kept. + Anthropic verifies the signature cryptographically, so a block with a null, + empty, or missing signature (e.g. from an open-source reasoning model) is + rejected with a 400, and so is a block whose signature or data carries + another provider's encrypted reasoning. It also rejects a `thinking` block + whose text is empty or whitespace-only ("each thinking block must contain + thinking"), regardless of signature, e.g. when a `thinking_blocks` history + item from a non-Anthropic reasoning provider is replayed through this path. + `redacted_thinking` blocks carry no signature and are always kept. """ if is_encrypted_reasoning_block(block): return True @@ -2006,7 +2009,10 @@ def is_unsignable_thinking_block(block: object) -> bool: if mapping.get("type") != "thinking": return False signature: Final = mapping.get("signature") - return not (isinstance(signature, str) and len(signature) > 0) + if not (isinstance(signature, str) and len(signature) > 0): + return True + thinking_text: Final = mapping.get("thinking") + return not (isinstance(thinking_text, str) and len(thinking_text.strip()) > 0) def strip_encrypted_reasoning_from_messages(messages: object) -> None: diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 26124ac24de..0c08c5dfd85 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -3895,3 +3895,179 @@ def test_anthropic_messages_pt_drops_a_system_message_with_no_text(): result = anthropic_messages_pt(messages=messages, model="claude-opus-4-8", llm_provider="anthropic") assert [m["role"] for m in result] == ["user", "assistant"] + + +def test_anthropic_messages_pt_drops_empty_but_signed_thinking_block(): + """ + Anthropic rejects a `thinking` block whose `thinking` text is empty, even + when it carries a valid-looking signature, with: + 400 messages.N.content.M.thinking: each thinking block must contain thinking + This shape is reachable via cross-provider replay of a `thinking_blocks` + history item (see PR #36033), so `is_unsignable_thinking_block()` must + also check the thinking text, not just the signature. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "", + "signature": "sig_abc123_looks_valid", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + content_types = [block.get("type") for block in assistant_msg["content"]] + assert "thinking" not in content_types, "empty-text thinking block must be dropped even though it has a signature" + + +def test_anthropic_messages_pt_keeps_non_empty_signed_thinking_block(): + """ + Regression: a real, non-empty, signed thinking block must still pass + through unchanged. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "Let me add these numbers together.", + "signature": "sig_abc123_looks_valid", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + thinking_block = next((b for b in assistant_msg["content"] if b.get("type") == "thinking"), None) + assert thinking_block is not None, "non-empty signed thinking block must be kept" + assert thinking_block["thinking"] == "Let me add these numbers together." + assert thinking_block["signature"] == "sig_abc123_looks_valid" + + +def test_anthropic_messages_pt_keeps_redacted_thinking_block(): + """ + Regression: `redacted_thinking` blocks carry no signature and no plaintext + `thinking` field by design, and must always be kept regardless of the new + emptiness check (which only applies to `type == "thinking"` blocks). + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "redacted_thinking", + "data": "encrypted_opaque_blob", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + content_types = [block.get("type") for block in assistant_msg["content"]] + assert "redacted_thinking" in content_types, "redacted_thinking blocks must always be kept" + + +def test_anthropic_messages_pt_drops_unsigned_thinking_block(): + """ + Regression (pre-existing behaviour): a thinking block with no signature + (or an empty/null one) must still be dropped, independent of whether the + thinking text is populated. + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + {"role": "user", "content": "What's 2+2?"}, + { + "role": "assistant", + "content": "4", + "thinking_blocks": [ + { + "type": "thinking", + "thinking": "Let me add these numbers together.", + "signature": "", + } + ], + }, + ] + + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5-20250929", + llm_provider="anthropic", + ) + + assistant_msg = result[1] + assert isinstance(assistant_msg["content"], list) + content_types = [block.get("type") for block in assistant_msg["content"]] + assert "thinking" not in content_types, "unsigned thinking block must still be dropped" + + +def test_is_unsignable_thinking_block_treats_whitespace_only_as_empty(): + """ + Edge case: a `thinking` field that is present but whitespace-only (e.g. + a single trailing newline forwarded from another provider's empty + reasoning summary) is functionally empty and Anthropic's API will still + reject it with "each thinking block must contain thinking". We treat it + the same as a fully empty string and drop the block. + + The check lives in the shared `is_unsignable_thinking_block` helper, which + `_drop_unsignable_thinking_blocks` calls standalone, so the whitespace-aware + test has to hold there rather than only at the factory call site. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + is_unsignable_thinking_block, + ) + + whitespace_only_block = { + "type": "thinking", + "thinking": " \n\t ", + "signature": "sig_abc123_looks_valid", + } + + assert is_unsignable_thinking_block(whitespace_only_block) is True From 303434d5738b5c1b0bbd226792fb40c1c21b6a16 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 26 Sep 2026 20:26:23 -0700 Subject: [PATCH 20/51] test(e2e): report batch cleanup leftovers as a plain UserWarning (#43405) The leftover warning used a class defined in a test-directory module. The xdist controller cannot import it, so an uncaught leftover warning crashed the whole e2e run. Same change as #43391 on rc/1.103.0 --- tests/e2e/batches/COVERAGE.md | 2 +- tests/e2e/batches/batch_cleanup.py | 8 ++------ tests/e2e/batches/test_batch_cleanup.py | 5 ++--- 3 files changed, 5 insertions(+), 10 deletions(-) diff --git a/tests/e2e/batches/COVERAGE.md b/tests/e2e/batches/COVERAGE.md index cd0fb35165e..2ba49a492ff 100644 --- a/tests/e2e/batches/COVERAGE.md +++ b/tests/e2e/batches/COVERAGE.md @@ -134,7 +134,7 @@ reporting failures as test errors. Already deleted files and batches that are terminal are safe to clean up again. Managed batch cancellation polls for up to two minutes before input deletion. A managed batch still `cancelling` after that is left for the provider to finish, and its input file is left in place because LiteLLM refuses to delete a file a non-terminal -batch references. Both are reported as `BatchCleanupLeftover` warnings naming their ids rather than +batch references. Both are reported as `UserWarning`s naming their ids rather than failing the test. Any other status or error still fails Accepted cancellation may still report validating or in_progress while the provider updates its state. Raw and model-encoded batches are polled until cancelling or diff --git a/tests/e2e/batches/batch_cleanup.py b/tests/e2e/batches/batch_cleanup.py index 5b3baaa624c..f1142a60782 100644 --- a/tests/e2e/batches/batch_cleanup.py +++ b/tests/e2e/batches/batch_cleanup.py @@ -28,10 +28,6 @@ class BatchCleanupClient(Protocol): def cancel_batch(self, batch_id: str, *, key: str, provider: str | None = None) -> Result[BatchObject]: ... -class BatchCleanupLeftover(UserWarning): - pass - - def cleanup_result[R: BaseModel]( action: Callable[[], Result[R]], *, wait: Callable[[float], None] = sleep ) -> Result[R]: @@ -68,7 +64,7 @@ def cleanup_file(client: BatchCleanupClient, file_id: str, *, key: str, provider if isinstance(result, UnknownApiError) and result.status_code == 400 and FILE_IN_USE_REFUSAL in result.body: warnings.warn( f"Left file {file_id} in place: LiteLLM refused to delete it while a batch still references it", - BatchCleanupLeftover, + UserWarning, stacklevel=2, ) return @@ -140,7 +136,7 @@ def cleanup_batch( ) warnings.warn( f"Left batch {batch_id} cancelling after {BATCH_CANCEL_TIMEOUT_SECONDS}s for the provider to finish", - BatchCleanupLeftover, + UserWarning, stacklevel=2, ) return diff --git a/tests/e2e/batches/test_batch_cleanup.py b/tests/e2e/batches/test_batch_cleanup.py index 5e2ac12d300..218e79f37ac 100644 --- a/tests/e2e/batches/test_batch_cleanup.py +++ b/tests/e2e/batches/test_batch_cleanup.py @@ -7,7 +7,6 @@ import pytest from batch_cleanup import ( BATCH_CANCEL_TIMEOUT_SECONDS, CLEANUP_DELAYS, - BatchCleanupLeftover, cleanup_batch, cleanup_file, cleanup_result, @@ -141,7 +140,7 @@ class TestFileCleanup: calls=ExpectedCalls((f"delete None {MANAGED_FILE_ID}",)), files=(UnknownApiError(status_code=400, body=IN_USE_REFUSAL),), ) - with pytest.warns(BatchCleanupLeftover, match=MANAGED_FILE_ID): + with pytest.warns(UserWarning, match=MANAGED_FILE_ID): cleanup_file(client, MANAGED_FILE_ID, key="test-key") client.calls.assert_done() @@ -241,7 +240,7 @@ class TestBatchCancellation: key: Final = manager.key() manager.defer(lambda: cleanup_file(client, MANAGED_FILE_ID, key=key)) manager.defer(lambda: cleanup_batch(client, MANAGED_BATCH_ID, key=key, clock=ticks)) - with pytest.warns(BatchCleanupLeftover) as leftovers: + with pytest.warns(UserWarning, match="^Left ") as leftovers: manager.teardown() client.calls.assert_done() messages: Final = tuple(str(warning.message) for warning in leftovers) From 21055e3fd840375f11a9f90d8b8e1c176284c79b Mon Sep 17 00:00:00 2001 From: Techboy bebop <142545999+kumarpriyanshu09@users.noreply.github.com> Date: Sun, 27 Sep 2026 00:15:12 -0400 Subject: [PATCH 21/51] fix(tools): salvage concatenated JSON tool call arguments (#43260) * fix(tools): salvage concatenated JSON tool call arguments * fix(tools): harden concatenated tool-call salvage for review findings Skip non-dict JSON during split so salvage cannot emit empty tool calls. Collapse srvtoolu_ expansions to the first object so server results stay paired. Allocate __concat_n ids that cannot collide with sibling tool call ids. Propagate cache_control onto every expanded Anthropic tool_use block. Rename the XML invoke loop variable so the key-leak gate no longer flags {args} * test(tools): cover concat id bump and srvtoolu array keep Only collapse srvtoolu_ when concatenated salvage expanded; a valid JSON array argument stays one server tool input * revert(anthropic): drop concat expansion from pass-through adapter Co-authored-by: Techboy bebop * revert(tools): keep concat salvage out of request-side tool converters Co-authored-by: Techboy bebop * fix(tools): expand strictly salvaged concatenated tool arguments in normalized tool calls Co-authored-by: Techboy bebop * fix(tools): retain at most the salvage cap while validating concatenated arguments Co-authored-by: Techboy bebop * test(tools): assert concat sibling ids unique after sanitization A sibling id that only collides after colon-to-underscore sanitization must force the next concat suffix Co-authored-by: Techboy bebop * refactor(tools): drop unused strict mode from split_concatenated_json_objects Strict mode had no production caller. Rejection cases now sit on salvage, and split matches upstream main Co-authored-by: Techboy bebop --------- Co-authored-by: Techboy bebop --- .../prompt_templates/common_utils.py | 55 ++ .../prompt_templates/factory.py | 210 +++++-- ...ore_utils_prompt_templates_common_utils.py | 40 +- ...llm_core_utils_prompt_templates_factory.py | 540 +++++++++--------- 4 files changed, 527 insertions(+), 318 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 41563d501a9..e555d7e8ec0 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -2503,6 +2503,61 @@ def split_concatenated_json_objects(raw: str) -> list[dict[str, object]]: return results +MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: Final = 8 + + +def salvage_concatenated_tool_arguments(raw: str) -> tuple[dict[str, object], ...]: + """Return complete concatenated JSON objects that are safe to expand. + + Identical objects collapse to the first one and are not capped. More than + ``MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS`` objects that are not all identical + returns an empty tuple. Anything that is not a full concatenation of JSON + objects returns an empty tuple. Repeated copies of the first object are not + retained, and once the cap is passed the rest of the string is only checked. + """ + stripped: Final = raw.strip() + if not stripped: + return () + decoder: Final = json.JSONDecoder() + length: Final = len(stripped) + idx = 0 # rebind-ok: cursor walks the concatenated JSON string + count = 0 # rebind-ok: counts complete objects without retaining duplicates + kept = () # rebind-ok: holds at most one object past the salvage cap + exceeded = False # rebind-ok: cap already passed, the tail is only validated + while idx < length: + while idx < length and stripped[idx] in " \t\n\r": + idx += 1 + if idx >= length: + break + try: + obj, end_idx = decoder.raw_decode(stripped, idx) + except json.JSONDecodeError: + return () + if not isinstance(obj, dict): + return () + idx = end_idx + if exceeded: + continue + count += 1 + if not kept: + kept = (obj,) + continue + if obj == kept[0] and len(kept) == 1: + continue + if len(kept) == 1 and count > 2 and count - 1 > MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: + exceeded = True + continue + if len(kept) == 1 and count > 2: + kept = (kept[0],) * (count - 1) + if len(kept) >= MAX_SALVAGED_TOOL_ARGUMENT_OBJECTS: + exceeded = True + continue + kept = (*kept, obj) + if exceeded: + return () + return kept + + def text_completion_prompt_to_messages(prompt: object) -> tuple[AllMessageValues, ...]: """ Wrap an OpenAI ``/v1/completions`` ``prompt`` into Chat Completion messages. diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 7b12d1e939f..c4e242fd360 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1,12 +1,14 @@ import base64 import copy import hashlib +import itertools import json import mimetypes import re import xml.etree.ElementTree as ET from collections.abc import Iterator, Mapping, Sequence from enum import Enum +from types import MappingProxyType from typing import Any, Final, TypeAlias, TypedDict, cast, overload from jinja2.sandbox import ImmutableSandboxedEnvironment @@ -52,6 +54,7 @@ from .common_utils import ( is_non_content_values_set, is_unsignable_thinking_block, parse_tool_call_arguments, + salvage_concatenated_tool_arguments, ) from .image_handling import convert_url_to_base64 @@ -5381,80 +5384,167 @@ class NormalizedToolCall(TypedDict): arguments: dict[str, object] -def _parse_tool_call_arguments(raw: object, tool_name: str | None, context: str) -> dict[str, object]: +_ArgumentObjects: TypeAlias = tuple[dict[str, object], ...] +_ParsedToolCall: TypeAlias = tuple[str | None, str | None, _ArgumentObjects] + + +def _optional_call_id(value: object) -> str | None: + if isinstance(value, str) and value: + return value + return None + + +def _optional_tool_name(value: object) -> str | None: + if isinstance(value, str): + return value + return None + + +def _split_tool_call_ids(calls: Sequence[tuple[str | None, int]]) -> tuple[tuple[str | None, ...], ...]: + taken: Final = frozenset(_sanitize_anthropic_tool_use_id(call_id) for call_id, _ in calls if call_id) + + def fresh(call_id: str) -> Iterator[str]: + return filter( + lambda candidate: _sanitize_anthropic_tool_use_id(candidate) not in taken, + (f"{call_id}__concat_{n}" for n in itertools.count(1)), + ) + + suffixes: Final = MappingProxyType( + {_sanitize_anthropic_tool_use_id(call_id): fresh(call_id) for call_id, count in calls if call_id and count > 1} + ) + return tuple( + ( + call_id, + *(next(suffixes[_sanitize_anthropic_tool_use_id(call_id)]) for _ in range(count - 1)), + ) + if call_id + else (None,) * count + for call_id, count in calls + ) + + +def _parse_tool_call_arguments(raw: object, tool_name: str | None, context: str) -> _ArgumentObjects: # Anthropic's tool_use blocks already carry a parsed dict in "input"; # chat completions and the Responses API carry a JSON string that may be # truncated by the model, so route those through the repair-aware parser. if isinstance(raw, dict): - return raw + return (raw,) if not isinstance(raw, str): - return {} + return ({},) normalized_raw: Final = "{}" if raw == REDACTED_BY_LITELLM else raw - from litellm.litellm_core_utils.prompt_templates.common_utils import ( - parse_tool_call_arguments, - ) - try: parsed: Final = parse_tool_call_arguments(normalized_raw, tool_name=tool_name, context=context) except ValueError as e: + salvaged: Final = salvage_concatenated_tool_arguments(normalized_raw) + if salvaged: + verbose_logger.warning( + "Recovered %d tool call(s) from concatenated JSON arguments for tool '%s' (%s)", + len(salvaged), + tool_name or "", + context, + ) + return salvaged verbose_logger.warning("Failed to parse tool call arguments: %s", e) - return {} - return parsed if isinstance(parsed, dict) else {} + return ({},) + return (parsed,) if isinstance(parsed, dict) else ({},) + + +def _choice_tool_calls(choice: object) -> tuple[object, ...]: + message: Final = get_attribute_or_key(choice, "message", None) + tool_calls: Final = get_attribute_or_key(message, "tool_calls", None) if message is not None else None + if isinstance(tool_calls, list): + return tuple(tool_calls) + return () + + +def _selected_choices(response: object, include_all_choices: bool) -> tuple[object, ...]: + choices: Final = get_attribute_or_key(response, "choices", None) + if not isinstance(choices, list) or not choices: + return () + if include_all_choices: + return tuple(choices) + return (choices[0],) + + +def _parsed_chat_tool_call(tool_call: object) -> _ParsedToolCall | None: + function: Final = get_attribute_or_key(tool_call, "function", None) + if function is None: + return None + name: Final = _optional_tool_name(get_attribute_or_key(function, "name")) + return ( + _optional_call_id(get_attribute_or_key(tool_call, "id")), + name, + _parse_tool_call_arguments( + get_attribute_or_key(function, "arguments", "{}"), + tool_name=name, + context="chat completions", + ), + ) + + +def _parsed_calls_in_choice(choice: object) -> tuple[_ParsedToolCall, ...]: + return tuple( + parsed for tool_call in _choice_tool_calls(choice) if (parsed := _parsed_chat_tool_call(tool_call)) is not None + ) + + +def _parsed_chat_tool_calls(response: object, include_all_choices: bool) -> tuple[_ParsedToolCall, ...]: + grouped: Final = tuple( + _parsed_calls_in_choice(choice) for choice in _selected_choices(response, include_all_choices) + ) + return tuple(itertools.chain.from_iterable(grouped)) + + +def _normalized_tool_calls_for_parse( + name: str | None, + call_ids: tuple[str | None, ...], + arguments: _ArgumentObjects, +) -> tuple[NormalizedToolCall, ...]: + return tuple( + NormalizedToolCall(id=call_id, name=name, arguments=argument) + for call_id, argument in zip(call_ids, arguments, strict=True) + ) + + +def _normalized_tool_calls_from_parses(parses: Sequence[_ParsedToolCall]) -> tuple[NormalizedToolCall, ...]: + id_groups: Final = _split_tool_call_ids(tuple((call_id, len(arguments)) for call_id, _, arguments in parses)) + grouped: Final = tuple( + _normalized_tool_calls_for_parse(name, call_ids, arguments) + for (_, name, arguments), call_ids in zip(parses, id_groups, strict=True) + ) + return tuple(itertools.chain.from_iterable(grouped)) def _tool_calls_from_chat_completion_response( response: object, include_all_choices: bool = False -) -> list[NormalizedToolCall]: - choices: Final = get_attribute_or_key(response, "choices", None) - if not (isinstance(choices, list) and choices): - return [] - tool_calls: Final[list[object]] = [] - for choice in choices if include_all_choices else choices[:1]: - message = get_attribute_or_key(choice, "message", None) - choice_tool_calls = get_attribute_or_key(message, "tool_calls", None) if message else None - if isinstance(choice_tool_calls, list): - tool_calls.extend(choice_tool_calls) - result: Final[list[NormalizedToolCall]] = [] - for tc in tool_calls: - fn = get_attribute_or_key(tc, "function", None) - if fn is None: - continue - name = get_attribute_or_key(fn, "name") - result.append( - NormalizedToolCall( - id=get_attribute_or_key(tc, "id"), - name=name, - arguments=_parse_tool_call_arguments( - get_attribute_or_key(fn, "arguments", "{}"), - tool_name=name, - context="chat completions", - ), - ) - ) - return result +) -> tuple[NormalizedToolCall, ...]: + return _normalized_tool_calls_from_parses(_parsed_chat_tool_calls(response, include_all_choices)) -def _tool_calls_from_responses_api_response(response: object) -> list[NormalizedToolCall]: +def _response_function_calls(response: object) -> tuple[object, ...]: output: Final = get_attribute_or_key(response, "output", None) if not isinstance(output, list): - return [] - result: Final[list[NormalizedToolCall]] = [] - for item in output: - if get_attribute_or_key(item, "type") != "function_call": - continue - name = get_attribute_or_key(item, "name") - result.append( - NormalizedToolCall( - id=get_attribute_or_key(item, "call_id") or get_attribute_or_key(item, "id"), - name=name, - arguments=_parse_tool_call_arguments( - get_attribute_or_key(item, "arguments", "{}"), - tool_name=name, - context="responses API", - ), - ) - ) - return result + return () + return tuple(item for item in output if get_attribute_or_key(item, "type") == "function_call") + + +def _parsed_response_tool_call(item: object) -> _ParsedToolCall: + name: Final = _optional_tool_name(get_attribute_or_key(item, "name")) + raw_id: Final = get_attribute_or_key(item, "call_id") or get_attribute_or_key(item, "id") + return ( + _optional_call_id(raw_id), + name, + _parse_tool_call_arguments( + get_attribute_or_key(item, "arguments", "{}"), + tool_name=name, + context="responses API", + ), + ) + + +def _tool_calls_from_responses_api_response(response: object) -> tuple[NormalizedToolCall, ...]: + parses: Final = tuple(_parsed_response_tool_call(item) for item in _response_function_calls(response)) + return _normalized_tool_calls_from_parses(parses) def _tool_calls_from_anthropic_messages_response(response: object) -> list[NormalizedToolCall]: @@ -5494,16 +5584,18 @@ def get_tool_calls_from_response(response: object, include_all_choices: bool = F Callers that only care about a specific tool should filter the result by ``name`` themselves -- this returns every tool call found. """ - chat_tool_calls = _tool_calls_from_chat_completion_response(response, include_all_choices=include_all_choices) + chat_tool_calls: Final = _tool_calls_from_chat_completion_response( + response, include_all_choices=include_all_choices + ) if chat_tool_calls: - return chat_tool_calls + return list(chat_tool_calls) for extractor in ( _tool_calls_from_responses_api_response, _tool_calls_from_anthropic_messages_response, ): tool_calls = extractor(response) if tool_calls: - return tool_calls + return list(tool_calls) return [] diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 79c50bf2369..45fc93f04c1 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -20,7 +20,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( hoist_images_from_tool_messages, is_encrypted_reasoning_block, merge_consecutive_system_messages, + parse_tool_call_arguments, responses_reasoning_items_from_thinking_blocks, + salvage_concatenated_tool_arguments, split_concatenated_json_objects, strip_encrypted_reasoning_from_messages, system_messages_first, @@ -269,6 +271,40 @@ def test_split_concatenated_json_salvages_prefix_before_truncated_tail(): assert result == [{"a": 1}, {"b": 2}] +def test_parse_tool_call_arguments_rejects_concatenated_json() -> None: + with pytest.raises(ValueError, match="Failed to parse tool call arguments"): + parse_tool_call_arguments('{"a":1}{"b":2}') + + +def _distinct_json_objects(count: int) -> str: + return "".join(json.dumps({"n": index}, separators=(",", ":")) for index in range(count)) + + +@pytest.mark.parametrize( + ("raw", "expected"), + ( + ('{"a":1}{"b":2}', ({"a": 1}, {"b": 2})), + ('{"a":1}{"a":1}{"a":1}', ({"a": 1},)), + ('{"a":1}{"a":1}{"b":2}', ({"a": 1}, {"a": 1}, {"b": 2})), + (_distinct_json_objects(8), tuple({"n": index} for index in range(8))), + (_distinct_json_objects(9), ()), + (_distinct_json_objects(9) + " junk", ()), + ('{"a":1}' * 7 + '{"b":2}', tuple({"a": 1} for _ in range(7)) + ({"b": 2},)), + ('{"a":1}' * 8 + '{"b":2}', ()), + ('{"a":1}' * 5000, ({"a": 1},)), + ('{"a":1}' * 20, ({"a": 1},)), + ('{"a":1}{"b":', ()), + ('0{"x":1}', ()), + ('{"x":1}0', ()), + ('[1]{"x":1}', ()), + ('{"a":1}{"b":2}}', ()), + ('{"a":1} junk', ()), + ), +) +def test_salvage_concatenated_tool_arguments(raw: str, expected: tuple[dict[str, object], ...]) -> None: + assert salvage_concatenated_tool_arguments(raw) == expected + + # --------------------------------------------------------------------------- # Regression tests for non-OpenAI file content blocks. # @@ -1949,6 +1985,8 @@ class TestMergeConsecutiveSystemMessages: assert merged == [{"role": "system", "content": expected_content}, {"role": "user", "content": "Hello"}] def test_keeps_the_first_message_when_no_system_message_in_the_run_has_content(self): - merged = merge_consecutive_system_messages([{"role": "system"}, {"role": "system"}, {"role": "user", "content": "Hi"}]) + merged = merge_consecutive_system_messages( + [{"role": "system"}, {"role": "system"}, {"role": "user", "content": "Hi"}] + ) assert merged == [{"role": "system"}, {"role": "user", "content": "Hi"}] diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 0c08c5dfd85..1e12a973cdb 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -1,4 +1,5 @@ import base64 +import json import logging import os import re @@ -16,10 +17,12 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( _bedrock_tools_pt, _rename_duplicate_bedrock_document_names, _convert_to_bedrock_tool_call_invoke, + _sanitize_anthropic_tool_use_id, _convert_to_bedrock_tool_call_result, anthropic_messages_pt, convert_to_anthropic_tool_result, convert_to_gemini_tool_call_result, + get_tool_calls_from_response, make_valid_bedrock_tool_name, ollama_pt, sanitize_messages_for_tool_calling, @@ -31,9 +34,7 @@ def _get_gemini_function_response_inline_data_parts(result): assert isinstance(result, list), "expected Gemini parts list" assert len(result) == 1, "multimodal function responses should stay in one part" function_response_part = result[0] - assert ( - "inline_data" not in function_response_part - ), "inline_data should be nested under function_response.parts" + assert "inline_data" not in function_response_part, "inline_data should be nested under function_response.parts" function_response = function_response_part["function_response"] nested_parts = function_response["parts"] return [part["inline_data"] for part in nested_parts if "inline_data" in part] @@ -49,7 +50,9 @@ def test_ollama_pt_simple_messages(): result = ollama_pt(model="llama2", messages=messages) - expected_prompt = "### System:\nYou are a helpful assistant\n\n### Assistant:\nHow can I help you?\n\n### User:\nHello\n\n" + expected_prompt = ( + "### System:\nYou are a helpful assistant\n\n### Assistant:\nHow can I help you?\n\n### User:\nHello\n\n" + ) assert isinstance(result, dict) assert result["prompt"] == expected_prompt assert result["images"] == [] @@ -104,10 +107,7 @@ async def test_anthropic_bedrock_thinking_blocks_with_none_content(): # verify the result assert len(result) == 2 - assert ( - result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] - == "This is a test thinking block" - ) + assert result[1]["content"][0]["reasoningContent"]["reasoningText"]["text"] == "This is a test thinking block" def test_bedrock_converse_assistant_with_empty_thinking_block_and_tool_calls(): @@ -175,11 +175,7 @@ def test_bedrock_converse_assistant_with_empty_thinking_block_and_tool_calls(): assert len(assistant_blocks) == 1 for block in assistant_blocks[0]["content"]: if "text" in block: - assert block[ - "text" - ].strip(), ( - f"Bedrock Converse rejects blank-text ContentBlocks; got {block!r}" - ) + assert block["text"].strip(), f"Bedrock Converse rejects blank-text ContentBlocks; got {block!r}" # toolUse blocks must still be present tool_use_blocks = [b for b in assistant_blocks[0]["content"] if "toolUse" in b] assert len(tool_use_blocks) == 2 @@ -220,19 +216,16 @@ def test_anthropic_messages_pt_drops_unsignable_thinking_block(thinking_block): {"role": "user", "content": "Now what is 3+3?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic") assistant = next(m for m in result if m["role"] == "assistant") content = assistant["content"] - assert all( - block.get("type") not in ("thinking", "redacted_thinking") for block in content - ), f"unsignable thinking block must be dropped, got {content!r}" - assert any( - block.get("type") == "text" and block.get("text") == "2+2 equals 4." - for block in content - ), f"assistant answer text must be preserved, got {content!r}" + assert all(block.get("type") not in ("thinking", "redacted_thinking") for block in content), ( + f"unsignable thinking block must be dropped, got {content!r}" + ) + assert any(block.get("type") == "text" and block.get("text") == "2+2 equals 4." for block in content), ( + f"assistant answer text must be preserved, got {content!r}" + ) def test_anthropic_messages_pt_keeps_signed_thinking_block(): @@ -255,9 +248,7 @@ def test_anthropic_messages_pt_keeps_signed_thinking_block(): {"role": "user", "content": "Now what is 3+3?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-6", llm_provider="anthropic") assistant = next(m for m in result if m["role"] == "assistant") thinking_blocks = [b for b in assistant["content"] if b.get("type") == "thinking"] @@ -373,9 +364,7 @@ def test_bedrock_get_document_format_fallback_mimes(): """ # Test DOCX fallback - docx_mime = ( - "application/vnd.openxmlformats-officedocument.wordprocessingml.document" - ) + docx_mime = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" supported_formats = ["pdf", "docx", "xlsx", "csv"] # Mock mimetypes.guess_all_extensions to return empty list (simulating Docker container scenario) @@ -399,15 +388,11 @@ def test_bedrock_get_document_format_mimetypes_success(): """ Test the _get_document_format method when mimetypes.guess_all_extensions works normally. """ - docx_mime = ( - "application/vnd.openxmlformats-officedocument.wordprocessingml.document" - ) + docx_mime = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" supported_formats = ["pdf", "docx", "xlsx", "csv"] # Test normal mimetypes behavior (should not hit fallback) - result = BedrockImageProcessor._get_document_format( - mime_type=docx_mime, supported_doc_formats=supported_formats - ) + result = BedrockImageProcessor._get_document_format(mime_type=docx_mime, supported_doc_formats=supported_formats) assert result == "docx", f"Expected 'docx', got '{result}'" @@ -623,9 +608,7 @@ async def test_bedrock_process_image_async_factory(): image_url = "data:application/pdf; qs=0.001;base64,JVBERi0xLjQKJcOkw7zDtsOfCjIgMCBvYmoKPDwvTGVuZ3RoIDMgMCBSL0ZpbHRlci9GbGF0ZURlY29kZT4" - content_block = await BedrockImageProcessor.process_image_async( - image_url=image_url, format=None - ) + content_block = await BedrockImageProcessor.process_image_async(image_url=image_url, format=None) print(f"content_block: {content_block}") @@ -668,9 +651,7 @@ def test_unpack_defs_resolves_nested_ref_inside_anyof_items(): items_schema = schema["properties"]["vatAmounts"]["anyOf"][0]["items"] # Assertions: items_schema should now be the resolved object, not an empty dict - assert isinstance( - items_schema, dict - ), "Items schema should be a dict after unpacking" + assert isinstance(items_schema, dict), "Items schema should be a dict after unpacking" assert items_schema.get("type") == "object" # Ensure essential properties are present assert set(items_schema.get("properties", {}).keys()) == {"vatRate", "vatAmount"} @@ -861,9 +842,7 @@ def test_convert_gemini_tool_call_result_with_multiple_anthropic_image_blocks(): last_message_with_tool_calls=last_message_with_tool_calls, ) inline_parts = _get_gemini_function_response_inline_data_parts(result) - assert ( - len(inline_parts) == 2 - ), f"expected 2 inline_data parts, got {len(inline_parts)}" + assert len(inline_parts) == 2, f"expected 2 inline_data parts, got {len(inline_parts)}" mime_types = {p["mime_type"] for p in inline_parts} assert mime_types == {"image/png", "image/jpeg"} @@ -899,9 +878,7 @@ def test_convert_gemini_tool_call_result_with_data_url_string(): last_message_with_tool_calls=last_message_with_tool_calls, ) inline_parts = _get_gemini_function_response_inline_data_parts(result) - assert ( - len(inline_parts) == 1 - ), "data-URL image string was not converted to inline_data" + assert len(inline_parts) == 1, "data-URL image string was not converted to inline_data" assert inline_parts[0]["mime_type"] == "image/png" assert inline_parts[0]["data"] == tiny_png_b64 @@ -937,9 +914,9 @@ def test_convert_gemini_tool_call_result_with_data_url_extra_params(): ) inline_parts = _get_gemini_function_response_inline_data_parts(result) assert len(inline_parts) == 1 - assert ( - inline_parts[0]["mime_type"] == "image/png" - ), f"expected clean 'image/png', got '{inline_parts[0]['mime_type']}'" + assert inline_parts[0]["mime_type"] == "image/png", ( + f"expected clean 'image/png', got '{inline_parts[0]['mime_type']}'" + ) def test_bedrock_tools_unpack_defs(): @@ -1036,9 +1013,7 @@ def test_bedrock_tools_pt_strict_parameter(): }, } ] - result = _bedrock_tools_pt( - tools_with_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools_with_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0") assert result[0]["toolSpec"]["strict"] is True assert result[0]["toolSpec"]["inputSchema"]["json"]["additionalProperties"] is False @@ -1060,9 +1035,7 @@ def test_bedrock_tools_pt_strict_parameter(): }, } ] - result = _bedrock_tools_pt( - tools_without_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools_without_strict, model="anthropic.claude-sonnet-4-5-20250929-v1:0") assert "strict" not in result[0]["toolSpec"] assert "additionalProperties" not in result[0]["toolSpec"]["inputSchema"]["json"] @@ -1085,9 +1058,7 @@ def test_bedrock_image_processor_content_type_fallback_url_extension(): # Test with .png URL image_url = "https://example.com/test-image.png" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -1111,9 +1082,7 @@ def test_bedrock_image_processor_content_type_fallback_binary_detection(): # Test with URL without extension image_url = "https://example.com/test-image-without-extension" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/jpeg" assert base64_bytes == base64.b64encode(jpeg_content).decode("utf-8") @@ -1136,9 +1105,7 @@ def test_bedrock_image_processor_content_type_fallback_application_octet_stream( # Test with .gif URL image_url = "https://s3.amazonaws.com/bucket/image.gif" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/gif" assert base64_bytes == base64.b64encode(gif_content).decode("utf-8") @@ -1161,9 +1128,7 @@ def test_bedrock_image_processor_content_type_with_query_params(): # Test with URL containing query parameters (common in S3 signed URLs) image_url = "https://s3.amazonaws.com/bucket/image.webp?AWSAccessKeyId=123&Expires=456&Signature=789" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/webp" assert base64_bytes == base64.b64encode(webp_content).decode("utf-8") @@ -1185,9 +1150,7 @@ def test_bedrock_image_processor_content_type_normal_header(): mock_response.content = png_content image_url = "https://example.com/test-image.png" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert content_type == "image/png" assert base64_bytes == base64.b64encode(png_content).decode("utf-8") @@ -1207,7 +1170,7 @@ def test_bedrock_image_processor_content_type_fallback_failure(): # Test with URL without recognizable extension image_url = "https://example.com/unknown-file" - with pytest.raises(ValueError, match='Unable to determine content type from URL: https') as excinfo: + with pytest.raises(ValueError, match="Unable to determine content type from URL: https") as excinfo: BedrockImageProcessor._post_call_image_processing(mock_response, image_url) assert "Unable to determine content type" in str(excinfo.value) @@ -1227,16 +1190,12 @@ def test_bedrock_image_processor_content_type_jpeg_variants(): # Test with .jpg extension image_url_jpg = "https://example.com/photo.jpg" - _, content_type_jpg = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url_jpg - ) + _, content_type_jpg = BedrockImageProcessor._post_call_image_processing(mock_response, image_url_jpg) assert content_type_jpg == "image/jpeg" # Test with .jpeg extension image_url_jpeg = "https://example.com/photo.jpeg" - _, content_type_jpeg = BedrockImageProcessor._post_call_image_processing( - mock_response, image_url_jpeg - ) + _, content_type_jpeg = BedrockImageProcessor._post_call_image_processing(mock_response, image_url_jpeg) assert content_type_jpeg == "image/jpeg" @@ -1258,9 +1217,7 @@ def test_bedrock_image_processor_content_type_pdf_document(): # Test with .pdf URL pdf_url = "https://s3.amazonaws.com/bucket/document.pdf" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, pdf_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, pdf_url) assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -1293,12 +1250,8 @@ def test_bedrock_image_processor_content_type_document_formats(): ] for url, expected_mime in test_cases: - _, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, url - ) - assert ( - content_type == expected_mime - ), f"Expected {expected_mime} for {url}, got {content_type}" + _, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, url) + assert content_type == expected_mime, f"Expected {expected_mime} for {url}, got {content_type}" def test_bedrock_image_processor_content_type_s3_pdf_with_query(): @@ -1317,9 +1270,7 @@ def test_bedrock_image_processor_content_type_s3_pdf_with_query(): # S3 signed URL with query parameters s3_url = "https://my-bucket.s3.us-east-1.amazonaws.com/documents/report.pdf?AWSAccessKeyId=AKIAIOSFODNN7EXAMPLE&Expires=1234567890&Signature=abcdef123456" - base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing( - mock_response, s3_url - ) + base64_bytes, content_type = BedrockImageProcessor._post_call_image_processing(mock_response, s3_url) assert content_type == "application/pdf" assert base64_bytes == base64.b64encode(pdf_content).decode("utf-8") @@ -1428,12 +1379,8 @@ def test_bedrock_create_bedrock_block_normalized_base64(): base64_content = base64.b64encode(pdf_content).decode("utf-8") # Create versions with different whitespace - base64_with_newlines = "\n".join( - [base64_content[i : i + 64] for i in range(0, len(base64_content), 64)] - ) - base64_with_spaces = " ".join( - [base64_content[i : i + 32] for i in range(0, len(base64_content), 32)] - ) + base64_with_newlines = "\n".join([base64_content[i : i + 64] for i in range(0, len(base64_content), 64)]) + base64_with_spaces = " ".join([base64_content[i : i + 32] for i in range(0, len(base64_content), 32)]) # Create blocks block1 = BedrockImageProcessor._create_bedrock_block( @@ -1565,9 +1512,7 @@ def test_bedrock_create_bedrock_block_document_name_format(): # Check format: DocumentPDFmessages_{16_hex_chars}_{format} pattern = r"^DocumentPDFmessages_[0-9a-f]{16}_pdf$" - assert re.match( - pattern, document_name - ), f"Document name format mismatch: {document_name}" + assert re.match(pattern, document_name), f"Document name format mismatch: {document_name}" def test_bedrock_create_bedrock_block_different_document_formats(): @@ -1620,9 +1565,7 @@ def test_bedrock_nova_web_search_options_mapping(): assert system_tool["name"] == "nova_grounding" # Test with search_context_size (should be ignored for Nova) - result2 = config._map_web_search_options( - {"search_context_size": "high"}, "us.amazon.nova-premier-v1:0" - ) + result2 = config._map_web_search_options({"search_context_size": "high"}, "us.amazon.nova-premier-v1:0") assert result2 is not None system_tool2 = result2.get("systemTool") @@ -1688,9 +1631,7 @@ def test_bedrock_tools_pt_drops_unmappable_responses_builtin_tools(): {"type": "custom", "name": "free_form"}, ] - result = _bedrock_tools_pt( - tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0") names = [block["toolSpec"]["name"] for block in result if "toolSpec" in block] assert names == ["noop"] @@ -1720,9 +1661,7 @@ def test_bedrock_tools_pt_keeps_anthropic_input_schema_tools(): }, ] - result = _bedrock_tools_pt( - tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result = _bedrock_tools_pt(tools=tools, model="anthropic.claude-sonnet-4-5-20250929-v1:0") names = [block["toolSpec"]["name"] for block in result if "toolSpec" in block] assert names == ["lookup"] @@ -1924,9 +1863,7 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): "tool_use_id": "srvtoolu_01ABC123", "content": { "type": "tool_search_tool_search_result", - "tool_references": [ - {"type": "tool_reference", "tool_name": "get_time"} - ], + "tool_references": [{"type": "tool_reference", "tool_name": "get_time"}], }, }, {"type": "text", "text": "I found the time tool. How can I help you?"}, @@ -1954,20 +1891,14 @@ def test_anthropic_messages_pt_server_tool_use_passthrough(): # Verify server_tool_use block is preserved assert "server_tool_use" in content_types - server_tool_use_block = next( - b for b in assistant_msg["content"] if b.get("type") == "server_tool_use" - ) + server_tool_use_block = next(b for b in assistant_msg["content"] if b.get("type") == "server_tool_use") assert server_tool_use_block["id"] == "srvtoolu_01ABC123" assert server_tool_use_block["name"] == "tool_search_tool_regex" assert server_tool_use_block["input"] == {"query": ".*time.*"} # Verify tool_search_tool_result block is preserved assert "tool_search_tool_result" in content_types - tool_result_block = next( - b - for b in assistant_msg["content"] - if b.get("type") == "tool_search_tool_result" - ) + tool_result_block = next(b for b in assistant_msg["content"] if b.get("type") == "tool_search_tool_result") assert tool_result_block["tool_use_id"] == "srvtoolu_01ABC123" assert tool_result_block["content"]["type"] == "tool_search_tool_search_result" assert tool_result_block["content"]["tool_references"][0]["tool_name"] == "get_time" @@ -2019,9 +1950,7 @@ def test_bedrock_tools_unpack_defs_no_oom_with_nested_refs(): "anyOf": [ {"$ref": "#/$defs/Literal"}, {"$ref": "#/$defs/FieldRef"}, - { - "$ref": "#/$defs/Expression" - }, # Circular: Operand -> Expression -> Operand + {"$ref": "#/$defs/Expression"}, # Circular: Operand -> Expression -> Operand ], }, "Literal": { @@ -2155,9 +2084,7 @@ def test_anthropic_messages_pt_file_block_cache_control_with_explicit_provider() file_block = content_blocks[0] assert file_block["type"] == "document" - assert ( - "cache_control" in file_block - ), "cache_control should be preserved on file/document content blocks" + assert "cache_control" in file_block, "cache_control should be preserved on file/document content blocks" assert file_block["cache_control"]["type"] == "ephemeral" text_block = content_blocks[1] @@ -2365,22 +2292,16 @@ def test_bedrock_tool_call_invoke_concatenated_json(): # First block keeps original tool id assert result[0]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN" assert result[0]["toolUse"]["name"] == "shell" - assert result[0]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009", "-m", "10"] - } + assert result[0]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009", "-m", "10"]} # Subsequent blocks get suffixed ids assert result[1]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN_1" assert result[1]["toolUse"]["name"] == "shell" - assert result[1]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009/robots.txt", "-m", "5"] - } + assert result[1]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009/robots.txt", "-m", "5"]} assert result[2]["toolUse"]["toolUseId"] == "tooluse_L7I3TewYAUhoheJZQEuwVN_2" assert result[2]["toolUse"]["name"] == "shell" - assert result[2]["toolUse"]["input"] == { - "command": ["curl", "-i", "http://localhost:9009/sitemap.xml", "-m", "5"] - } + assert result[2]["toolUse"]["input"] == {"command": ["curl", "-i", "http://localhost:9009/sitemap.xml", "-m", "5"]} def test_bedrock_tool_call_invoke_concatenated_json_with_cache_control(): @@ -2535,9 +2456,7 @@ def test_bedrock_tool_call_invoke_unconvertible_raises_non_retryable_bad_request def test_make_valid_bedrock_tool_name_preserves_hyphens(): assert make_valid_bedrock_tool_name("my-tool") == "my-tool" assert ( - make_valid_bedrock_tool_name( - "CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q" - ) + make_valid_bedrock_tool_name("CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q") == "CreateCaseKnowledgeArticle_foTWsqR6yDt-OnSsvR5e6Q" ) @@ -2564,9 +2483,7 @@ def test_bedrock_tool_name_sanitized_consistently_in_tools_and_tool_use(): "function": {"name": raw_name, "arguments": "{}"}, } ] - tool_use_name = _convert_to_bedrock_tool_call_invoke(tool_calls)[0]["toolUse"][ - "name" - ] + tool_use_name = _convert_to_bedrock_tool_call_invoke(tool_calls)[0]["toolUse"]["name"] assert tool_spec_name == "foo_bar" assert tool_use_name == tool_spec_name @@ -2589,15 +2506,8 @@ def test_bedrock_converse_messages_pt_tool_use_matches_tool_spec_hyphen_name(): ], }, ] - translated = _bedrock_converse_messages_pt( - messages=messages, model="", llm_provider="" - ) - tool_use_blocks = [ - block - for msg in translated - for block in msg.get("content", []) - if "toolUse" in block - ] + translated = _bedrock_converse_messages_pt(messages=messages, model="", llm_provider="") + tool_use_blocks = [block for msg in translated for block in msg.get("content", []) if "toolUse" in block] assert len(tool_use_blocks) == 1 assert tool_use_blocks[0]["toolUse"]["name"] == tool_name @@ -2694,11 +2604,7 @@ def test_sanitize_messages_deduplicates_tool_results(): result = sanitize_messages_for_tool_calling(messages) # Count tool messages with this ID — should be exactly 1 - tool_results = [ - m - for m in result - if m.get("role") == "tool" and m.get("tool_call_id") == "call_abc123" - ] + tool_results = [m for m in result if m.get("role") == "tool" and m.get("tool_call_id") == "call_abc123"] assert len(tool_results) == 1 # Should keep the LAST occurrence (most complete) assert tool_results[0]["content"] == '{"temperature": 72, "condition": "sunny"}' @@ -2833,11 +2739,7 @@ def test_sanitize_messages_dedup_scoped_per_turn_preserves_cross_turn(): result = sanitize_messages_for_tool_calling(messages) # Both tool results must survive — one per turn - tool_results = [ - m - for m in result - if m.get("role") == "tool" and m.get("tool_call_id") == "call_X" - ] + tool_results = [m for m in result if m.get("role") == "tool" and m.get("tool_call_id") == "call_X"] assert len(tool_results) == 2, ( f"Expected 2 tool results (one per turn), got {len(tool_results)}. " "Dedup may be global instead of per-turn scoped." @@ -2891,32 +2793,26 @@ def test_sanitize_messages_combined_case_a_and_case_d(): tool_results = [m for m in result if m.get("role") in ("tool", "function")] # Case A: call_missing should have a dummy result injected - missing_results = [ - m for m in tool_results if m.get("tool_call_id") == "call_missing" - ] - assert ( - len(missing_results) == 1 - ), f"Expected 1 dummy result for call_missing (Case A), got {len(missing_results)}" + missing_results = [m for m in tool_results if m.get("tool_call_id") == "call_missing"] + assert len(missing_results) == 1, ( + f"Expected 1 dummy result for call_missing (Case A), got {len(missing_results)}" + ) # Case D: call_duped should have exactly 1 result (the fresh one) - duped_results = [ - m for m in tool_results if m.get("tool_call_id") == "call_duped" - ] - assert ( - len(duped_results) == 1 - ), f"Expected 1 result for call_duped after dedup (Case D), got {len(duped_results)}" - assert ( - duped_results[0]["content"] == "fresh_result" - ), f"Expected last-wins 'fresh_result', got '{duped_results[0]['content']}'" + duped_results = [m for m in tool_results if m.get("tool_call_id") == "call_duped"] + assert len(duped_results) == 1, ( + f"Expected 1 result for call_duped after dedup (Case D), got {len(duped_results)}" + ) + assert duped_results[0]["content"] == "fresh_result", ( + f"Expected last-wins 'fresh_result', got '{duped_results[0]['content']}'" + ) # Verify tool results immediately follow the assistant message asst_idx = next(i for i, m in enumerate(result) if m.get("role") == "assistant") - tool_msgs_after_asst = [ - m for m in result[asst_idx + 1 :] if m.get("role") in ("tool", "function") - ] - assert ( - len(tool_msgs_after_asst) == 2 - ), f"Expected 2 tool results after assistant, got {len(tool_msgs_after_asst)}" + tool_msgs_after_asst = [m for m in result[asst_idx + 1 :] if m.get("role") in ("tool", "function")] + assert len(tool_msgs_after_asst) == 2, ( + f"Expected 2 tool results after assistant, got {len(tool_msgs_after_asst)}" + ) # Both tool_call_ids should be present (order may vary) tool_ids = {m["tool_call_id"] for m in tool_msgs_after_asst} assert tool_ids == { @@ -2958,9 +2854,7 @@ def test_anthropic_messages_pt_file_block_preserves_cache_control(): } ] - result = anthropic_messages_pt( - messages, model="claude-sonnet-4-20250514", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages, model="claude-sonnet-4-20250514", llm_provider="anthropic") content_blocks = result[0]["content"] assert len(content_blocks) == 2 @@ -2968,9 +2862,7 @@ def test_anthropic_messages_pt_file_block_preserves_cache_control(): # Document block (from file) should preserve cache_control doc_block = content_blocks[0] assert doc_block["type"] == "document" - assert ( - "cache_control" in doc_block - ), "cache_control was dropped from file/document block" + assert "cache_control" in doc_block, "cache_control was dropped from file/document block" assert doc_block["cache_control"]["type"] == "ephemeral" # Text block should also preserve cache_control @@ -3013,9 +2905,7 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): } # Claude 4.5 model: ttl should be preserved - result = add_cache_point_tool_block( - tool_with_1h, model="jp.anthropic.claude-opus-4-7" - ) + result = add_cache_point_tool_block(tool_with_1h, model="jp.anthropic.claude-opus-4-7") assert result is not None assert result["cachePoint"]["type"] == "default" assert result["cachePoint"]["ttl"] == "1h" @@ -3024,16 +2914,12 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): tool_with_5m = { "cache_control": {"type": "ephemeral", "ttl": "5m"}, } - result_5m = add_cache_point_tool_block( - tool_with_5m, model="jp.anthropic.claude-opus-4-7" - ) + result_5m = add_cache_point_tool_block(tool_with_5m, model="jp.anthropic.claude-opus-4-7") assert result_5m is not None assert result_5m["cachePoint"]["ttl"] == "5m" # Older model: ttl should be stripped - result_old = add_cache_point_tool_block( - tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0" - ) + result_old = add_cache_point_tool_block(tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0") assert result_old is not None assert result_old["cachePoint"]["type"] == "default" assert "ttl" not in result_old["cachePoint"] @@ -3052,9 +2938,7 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch): # cache_control without ttl: returns default cachePoint (unchanged behavior) tool_no_ttl = {"cache_control": {"type": "ephemeral"}} - result_no_ttl = add_cache_point_tool_block( - tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0" - ) + result_no_ttl = add_cache_point_tool_block(tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0") assert result_no_ttl is not None assert result_no_ttl["cachePoint"]["type"] == "default" assert "ttl" not in result_no_ttl["cachePoint"] @@ -3127,9 +3011,7 @@ def test_bedrock_tools_pt_passes_ttl_for_claude_4_5(monkeypatch): assert cache_blocks[0]["cachePoint"]["ttl"] == "1h" # Older model: cachePoint should not have ttl - result_old = _bedrock_tools_pt( - tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0" - ) + result_old = _bedrock_tools_pt(tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0") cache_blocks_old = [b for b in result_old if "cachePoint" in b] assert len(cache_blocks_old) == 1 assert "ttl" not in cache_blocks_old[0]["cachePoint"] @@ -3204,9 +3086,7 @@ def test_bedrock_converse_messages_pt_document_various_formats(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") doc_block = result[0]["content"][0] assert doc_block["document"]["format"] == expected_format, ( @@ -3233,12 +3113,8 @@ def test_bedrock_converse_messages_pt_document_deterministic_name(): } ] - result1 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) - result2 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result1 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") + result2 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") name1 = result1[0]["content"][0]["document"]["name"] name2 = result2[0]["content"][0]["document"]["name"] @@ -3272,34 +3148,18 @@ def test_bedrock_converse_messages_pt_renames_duplicate_document_names(): }, ] - result1 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) - result2 = _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + result1 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") + result2 = _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") - names1 = [ - block["document"]["name"] - for message in result1 - for block in message["content"] - if "document" in block - ] - names2 = [ - block["document"]["name"] - for message in result2 - for block in message["content"] - if "document" in block - ] + names1 = [block["document"]["name"] for message in result1 for block in message["content"] if "document" in block] + names2 = [block["document"]["name"] for message in result2 for block in message["content"] if "document" in block] assert len(names1) == 2 assert len(set(names1)) == 2 assert names1[1] == f"{names1[0]}_2" assert names1 == names2 - single_turn = _bedrock_converse_messages_pt( - [messages[0]], "anthropic.claude-sonnet-4-6", "bedrock" - ) + single_turn = _bedrock_converse_messages_pt([messages[0]], "anthropic.claude-sonnet-4-6", "bedrock") assert names1[0] == single_turn[0]["content"][0]["document"]["name"] @@ -3321,14 +3181,10 @@ def test_rename_duplicate_bedrock_document_names_skips_organic_suffixes(): def _names(contents): return [block["document"]["name"] for block in contents[0]["content"]] - organic_first = _rename_duplicate_bedrock_document_names( - _contents(["report", "report_2", "report"]) - ) + organic_first = _rename_duplicate_bedrock_document_names(_contents(["report", "report_2", "report"])) assert _names(organic_first) == ["report", "report_2", "report_3"] - organic_last = _rename_duplicate_bedrock_document_names( - _contents(["report", "report", "report_2"]) - ) + organic_last = _rename_duplicate_bedrock_document_names(_contents(["report", "report", "report_2"])) assert _names(organic_last) == ["report", "report_3", "report_2"] @@ -3350,18 +3206,11 @@ def test_bedrock_converse_messages_pt_document_rejects_url_source(): ] with pytest.raises(ValueError, match="only supports base64-encoded"): - _bedrock_converse_messages_pt( - messages, "anthropic.claude-sonnet-4-6", "bedrock" - ) + _bedrock_converse_messages_pt(messages, "anthropic.claude-sonnet-4-6", "bedrock") def _collect_cache_points(blocks): - return [ - block["cachePoint"] - for message in blocks - for block in message["content"] - if "cachePoint" in block - ] + return [block["cachePoint"] for message in blocks for block in message["content"] if "cachePoint" in block] @pytest.mark.parametrize( @@ -3527,6 +3376,189 @@ def test_get_tool_calls_from_response_warns_for_malformed_arguments(caplog): assert "Failed to parse tool call arguments" in caplog.text +def _concatenated_json(*payloads: dict[str, object]) -> str: + return "".join(json.dumps(payload, separators=(",", ":")) for payload in payloads) + + +def _function_tool_call(call_id: str | None, name: str, arguments: str) -> dict[str, object]: + return {"id": call_id, "function": {"name": name, "arguments": arguments}} + + +def _chat_tool_response(*tool_calls: dict[str, object]) -> dict[str, object]: + return {"choices": [{"message": {"tool_calls": list(tool_calls)}}]} + + +def test_get_tool_calls_from_response_expands_distinct_concatenated_arguments(caplog): + raw = '{"flag":true}{"box":"A","limit":50}' + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", raw)) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + tool_calls: Final = get_tool_calls_from_response(response) + + assert tool_calls == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + {"id": "call_move__concat_1", "name": "move", "arguments": {"box": "A", "limit": 50}}, + ] + assert "Recovered 2 tool call(s)" in caplog.text + assert "move" in caplog.text + assert "flag" not in caplog.text + + +def test_get_tool_calls_from_response_expands_responses_api_concatenated_arguments(): + response: Final = { + "output": [ + { + "type": "function_call", + "call_id": "call_move", + "name": "move", + "arguments": '{"flag":true}{"box":"A","limit":50}', + } + ] + } + + assert get_tool_calls_from_response(response) == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + {"id": "call_move__concat_1", "name": "move", "arguments": {"box": "A", "limit": 50}}, + ] + + +def test_get_tool_calls_from_response_collapses_identical_concatenated_arguments(): + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", '{"flag":true}' * 3)) + + assert get_tool_calls_from_response(response) == [ + {"id": "call_move", "name": "move", "arguments": {"flag": True}}, + ] + + +def test_get_tool_calls_from_response_does_not_expand_a_valid_json_array(): + response: Final = _chat_tool_response(_function_tool_call("call_batch", "batch", '[{"a":1},{"b":2}]')) + + assert get_tool_calls_from_response(response) == [ + {"id": "call_batch", "name": "batch", "arguments": {}}, + ] + + +@pytest.mark.parametrize("arguments", ('{"a":1}{"b":', '0{"x":1}')) +def test_get_tool_calls_from_response_drops_partial_concatenated_arguments(arguments: str, caplog): + response: Final = _chat_tool_response(_function_tool_call("call_move", "move", arguments)) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + tool_calls: Final = get_tool_calls_from_response(response) + + assert tool_calls == [{"id": "call_move", "name": "move", "arguments": {}}] + assert "Failed to parse tool call arguments" in caplog.text + + +@pytest.mark.parametrize(("count", "expands"), ((8, True), (9, False))) +def test_get_tool_calls_from_response_caps_distinct_concatenated_arguments(count: int, expands: bool): + raw = _concatenated_json(*({"n": index} for index in range(count))) + response: Final = _chat_tool_response(_function_tool_call("call", "move", raw)) + + tool_calls: Final = get_tool_calls_from_response(response) + + if expands: + assert [call["id"] for call in tool_calls] == ["call", *(f"call__concat_{index}" for index in range(1, count))] + assert [call["arguments"] for call in tool_calls] == [{"n": index} for index in range(count)] + return + assert tool_calls == [{"id": "call", "name": "move", "arguments": {}}] + + +def test_get_tool_calls_from_response_skips_concat_ids_taken_by_a_sibling(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("call", "move", raw), + _function_tool_call("call__concat_1", "look", '{"x":1}'), + ) + + assert [call["id"] for call in get_tool_calls_from_response(response)] == [ + "call", + "call__concat_2", + "call__concat_1", + ] + + +def test_get_tool_calls_from_response_keeps_sanitized_concat_ids_distinct(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("a:b", "move", raw), + _function_tool_call("a_b__concat_1", "look", '{"x":1}'), + ) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + sanitized: Final = [_sanitize_anthropic_tool_use_id(call_id) for call_id in ids if isinstance(call_id, str)] + + assert len(sanitized) == len(set(sanitized)) + assert ids == ["a:b", "a:b__concat_2", "a_b__concat_1"] + + +def test_get_tool_calls_from_response_bumps_suffix_when_sibling_sanitizes_onto_it(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("a_b", "move", raw), + _function_tool_call("a:b__concat_1", "look", '{"x":1}'), + ) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + sanitized: Final = [_sanitize_anthropic_tool_use_id(call_id) for call_id in ids if isinstance(call_id, str)] + + assert len(ids) == len(sanitized) + assert len(sanitized) == len(set(sanitized)) + assert ids == ["a_b", "a_b__concat_2", "a:b__concat_1"] + + +def test_get_tool_calls_from_response_continues_concat_suffixes_per_sanitized_base(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response( + _function_tool_call("x", "move", raw), + _function_tool_call("x", "move", raw), + ) + + assert [call["id"] for call in get_tool_calls_from_response(response)] == [ + "x", + "x__concat_1", + "x", + "x__concat_2", + ] + + +def test_get_tool_calls_from_response_skips_a_run_of_reserved_concat_ids(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + siblings: Final = tuple(_function_tool_call(f"call__concat_{index}", "look", '{"x":1}') for index in range(1, 51)) + response: Final = _chat_tool_response(_function_tool_call("call", "move", raw), *siblings) + + ids: Final = [call["id"] for call in get_tool_calls_from_response(response)] + + assert ids[0] == "call" + assert ids[1] == "call__concat_51" + + +def test_get_tool_calls_from_response_reserves_concat_ids_across_choices(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = { + "choices": [ + {"message": {"tool_calls": [_function_tool_call("call", "move", raw)]}}, + {"message": {"tool_calls": [_function_tool_call("call__concat_1", "look", '{"x":1}')]}}, + ] + } + + assert [call["id"] for call in get_tool_calls_from_response(response, include_all_choices=True)] == [ + "call", + "call__concat_2", + "call__concat_1", + ] + + +def test_get_tool_calls_from_response_does_not_invent_ids_for_a_missing_call_id(): + raw = _concatenated_json({"a": 1}, {"b": 2}) + response: Final = _chat_tool_response(_function_tool_call(None, "move", raw)) + + tool_calls: Final = get_tool_calls_from_response(response) + + assert len(tool_calls) == 2 + assert all(call["id"] is None for call in tool_calls) + assert [call["arguments"] for call in tool_calls] == [{"a": 1}, {"b": 2}] + + def test_group_tool_exchanges_pairs_assistant_with_its_tool_rows(): from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges @@ -3625,9 +3657,7 @@ def test_bedrock_converse_pdf_only_user_message_gets_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert len(result) == 1 assert any("document" in block for block in result[0]["content"]) @@ -3645,9 +3675,7 @@ def test_bedrock_converse_document_with_text_gets_no_extra_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert _text_blocks(result[0]) == ["summarize this"] @@ -3660,9 +3688,7 @@ def test_bedrock_converse_image_only_user_message_gets_no_text_block(): } ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert any("image" in block for block in result[0]["content"]) assert _text_blocks(result[0]) == [] @@ -3705,9 +3731,7 @@ def test_bedrock_converse_tool_round_trip_document_injects_text_before_cache_poi }, ] - result = _bedrock_converse_messages_pt( - messages, "anthropic.claude-haiku-4-5", "bedrock" - ) + result = _bedrock_converse_messages_pt(messages, "anthropic.claude-haiku-4-5", "bedrock") assert _text_blocks(result[0]) == ["read the pdf"] document_message = result[-1] From 829cba1bf18c22593ddf65737e30d1b905651259 Mon Sep 17 00:00:00 2001 From: Thippaluri Yaseen Basha Date: Sun, 27 Sep 2026 09:48:33 +0530 Subject: [PATCH 22/51] fix(gemini): forward seed to the Gemini API instead of rejecting it (#43197) * fix(gemini): forward seed to the Gemini API instead of rejecting it The gemini/ provider left seed out of its supported params, so requests with seed failed with UnsupportedParamsError, or lost the seed silently when drop_params was on. The Gemini API accepts generationConfig.seed and the inherited mapping already translates it, so adding it to the allowlist is enough * test(gemini): assert the forwarded seed without mutating shared state --- litellm/llms/gemini/chat/transformation.py | 1 + ...test_vertex_and_google_ai_studio_gemini.py | 28 +++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/litellm/llms/gemini/chat/transformation.py b/litellm/llms/gemini/chat/transformation.py index 285350aecba..cae0ba49c3f 100644 --- a/litellm/llms/gemini/chat/transformation.py +++ b/litellm/llms/gemini/chat/transformation.py @@ -96,6 +96,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): "logprobs", "frequency_penalty", "presence_penalty", + "seed", "modalities", "parallel_tool_calls", "web_search_options", diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 7548f3c2daa..fd735afb16e 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -3413,6 +3413,34 @@ def test_google_ai_studio_presence_penalty_supported(): assert "presence_penalty" in supported_params +@pytest.mark.asyncio +@pytest.mark.parametrize("drop_params", [False, True]) +async def test_google_ai_studio_forwards_seed_to_generation_config(drop_params: bool): + def echo_seed_sent_upstream(request: httpx.Request) -> httpx.Response: + seed_sent: Final = json.loads(request.content).get("generationConfig", {}).get("seed") + return httpx.Response( + 200, + json={ + "candidates": [ + {"content": {"parts": [{"text": f"seed={seed_sent}"}], "role": "model"}, "finishReason": "STOP"} + ], + "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }, + request=request, + ) + + response: Final = await litellm.acompletion( + model="gemini/gemini-3.8-flash", + messages=[{"role": "user", "content": "hi"}], + seed=42, + drop_params=drop_params, + api_key="fake-gemini-key", + client=AsyncHTTPHandler(transport=httpx.MockTransport(echo_seed_sent_upstream)), + ) + + assert response.choices[0].message.content == "seed=42" + + # ==================== Tool Type Separation Tests ==================== # These tests verify that each Tool object contains exactly one type per Vertex AI API spec # Ref: https://cloud.google.com/vertex-ai/generative-ai/docs/reference/rest/v1beta1/Tool From 2101c860c25546250fa55c5a67653093a9d28873 Mon Sep 17 00:00:00 2001 From: Stewart Park <388348+stewartpark@users.noreply.github.com> Date: Sat, 26 Sep 2026 21:26:38 -0700 Subject: [PATCH 23/51] fix(vertex_ai): make Gemma fake streams work with traced Responses (#43147) * test(vertex_ai): reproduce traced Gemma Responses stream failure * fix(vertex_ai): wrap Gemma fake streams for Responses tracing * test(vertex_ai): cover Gemma traced streams and usage options * test(vertex_ai): inject gemma test deps and assert hidden usage accounting Replace class-level patches in the Vertex AI shard test with the provider's documented dependency-injection seams (httpx.MockTransport client + credential cache), and pin the default/omit-usage trace behavior: LiteLLM still accounts all tokens; ddtrace's metric is absent by design, asserted rather than silent. Mutation-checked: commenting out CustomStreamWrapper chunk accumulation turns the new assertions red; restoring them turns green. * test(vertex_ai): drop explanatory comment from usage-option assertions --- .../vertex_gemma_models/transformation.py | 30 ++++- .../test_vertex_gemma_transformation.py | 93 ++++++++++++++ .../test_vertex_gemma_transformation.py | 117 +++++++++++++++++- 3 files changed, 229 insertions(+), 11 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index ea97f0a0a9a..33922e38674 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -28,8 +28,8 @@ from litellm.types.utils import ModelResponse if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer - from litellm.llms.base_llm.base_model_iterator import MockResponseIterator def parse_vertex_gemma_container_error(predictions: object) -> VertexGemmaContainerError | None: @@ -73,7 +73,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): self, model_response: ModelResponse, stream: bool, - ) -> "ModelResponse | MockResponseIterator": + model: str, + logging_obj: "LiteLLMLoggingObj", + ) -> "ModelResponse | CustomStreamWrapper": """ Helper method to return fake stream iterator if streaming is requested. @@ -82,12 +84,18 @@ class VertexGemmaConfig(OpenAIGPTConfig): stream: Whether streaming was requested Returns: - MockResponseIterator if stream=True, otherwise the model_response + CustomStreamWrapper if stream=True, otherwise the model_response """ if stream: + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator - return MockResponseIterator(model_response=model_response) + return CustomStreamWrapper( + completion_stream=MockResponseIterator(model_response=model_response), + model=model, + custom_llm_provider="vertex_ai", + logging_obj=logging_obj, + ) return model_response def transform_request( @@ -373,7 +381,12 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) async def _async_completion( self, @@ -463,4 +476,9 @@ class VertexGemmaConfig(OpenAIGPTConfig): ) # Return fake stream iterator if streaming was requested - return self._handle_fake_stream_response(model_response=model_response, stream=stream) + return self._handle_fake_stream_response( + model_response=model_response, + stream=stream, + model=model, + logging_obj=logging_obj, + ) diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py new file mode 100644 index 00000000000..294d26b2e58 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -0,0 +1,93 @@ +import json +from collections.abc import AsyncIterator +from types import SimpleNamespace +from typing import Any, cast + +import httpx +import pytest + +import litellm +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.main import vertex_gemma_chat_completion +from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse + +_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" +_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}] +_FAKE_CREDENTIALS = "gemma-test-credentials" + + +def _vertex_response(): + return { + "predictions": { + "id": "chatcmpl-stream-test", + "created": 1759863903, + "model": "google/gemma-3-12b-it", + "object": "chat.completion", + "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}], + "usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15}, + } + } + + +@pytest.fixture(autouse=True) +def _cached_access_token(): + """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" + cache = vertex_gemma_chat_completion._credentials_project_mapping + key = (_FAKE_CREDENTIALS, "test") + cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") + yield + cache.pop(key, None) + + +def test_sync_gemma_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + stream = litellm.completion( + model="vertex_ai/gemma/test-model", + messages=_MESSAGES, + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.Client(transport=httpx.MockTransport(handle)), + ) + + assert isinstance(stream, CustomStreamWrapper) + chunks = list(stream) + + assert "stream" not in captured["body"]["instances"][0] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_async_gemma_responses_stream(): + captured: dict[str, Any] = {} + + def handle(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_vertex_response()) + + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base=_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_CREDENTIALS, + client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), + ) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + + assert "stream" not in captured["body"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage is not None + assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index e5ca31833ce..97f4f290958 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -5,11 +5,18 @@ Maps to: litellm/llms/vertex_ai/vertex_gemma_models/transformation.py """ import json +from collections.abc import AsyncIterator +from typing import cast from unittest.mock import AsyncMock, Mock, patch import pytest import litellm +from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIStreamingResponse, +) @pytest.fixture(autouse=True) @@ -439,8 +446,9 @@ class TestVertexGemmaCompletion: Verifies: 1. Request body does NOT include 'stream' parameter (model doesn't support it) - 2. Response returns a MockResponseIterator that yields chunks + 2. Response wraps a MockResponseIterator and yields chunks """ + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.base_model_iterator import MockResponseIterator # Mock Vertex response @@ -502,8 +510,8 @@ class TestVertexGemmaCompletion: vertex_location="us-central1", ) - # Verify the response is a MockResponseIterator - assert isinstance(response, MockResponseIterator), f"Expected MockResponseIterator, got {type(response)}" + assert isinstance(response, CustomStreamWrapper) + assert isinstance(response.completion_stream, MockResponseIterator) # Verify the request sent to Vertex does NOT include 'stream' call_args = mock_client.post.call_args @@ -520,8 +528,9 @@ class TestVertexGemmaCompletion: async for chunk in response: chunks.append(chunk) - # Should get exactly one chunk (fake streaming) - assert len(chunks) == 1, f"Expected 1 chunk from fake stream, got {len(chunks)}" + assert len(chunks) == 2 + assert chunks[1].choices[0].finish_reason == "stop" + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) # Verify the chunk has the expected content chunk = chunks[0] @@ -529,6 +538,104 @@ class TestVertexGemmaCompletion: assert len(chunk.choices) > 0 assert chunk.choices[0].delta.content == "Streaming test response" + @pytest.mark.asyncio + async def test_aresponses_streams_vertex_gemma_with_llm_tracing(self): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + from ddtrace.llmobs._integrations.base_stream_handler import TracedAsyncStream + + from litellm.responses.litellm_completion_transformation.streaming_iterator import ( + LiteLLMCompletionStreamingIterator, + ) + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock() + client.post = AsyncMock(return_value=reply) + + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + bridge = cast(LiteLLMCompletionStreamingIterator, response) + traced_stream = bridge.litellm_custom_stream_wrapper + assert isinstance(traced_stream, TracedAsyncStream) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + span = traced_stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + finally: + unpatch_litellm() + + assert "stream" not in client.post.call_args.kwargs["json"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage.total_tokens == 114 + + @pytest.mark.asyncio + @pytest.mark.parametrize("stream_options", [None, {"include_usage": False}, {"include_usage": True}]) + async def test_acompletion_stream_respects_usage_option_with_llm_tracing(self, stream_options): + pytest.importorskip("ddtrace") + from ddtrace.contrib.internal.litellm.patch import patch as patch_litellm + from ddtrace.contrib.internal.litellm.patch import unpatch as unpatch_litellm + + reply = Mock(status_code=200) + reply.json.return_value = _make_gemma_vertex_response(content="READY") + client = Mock(post=AsyncMock(return_value=reply)) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client), + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "test-project"), + ), + ): + patch_litellm() + try: + stream = await litellm.acompletion( + model="vertex_ai/gemma/test-model", + messages=[{"role": "user", "content": "Reply exactly READY"}], + stream=True, + **({"stream_options": stream_options} if stream_options is not None else {}), + api_base="https://example.invalid/v1/projects/test-project/locations/us-central1/endpoints/test:predict", + vertex_project="test-project", + vertex_location="us-central1", + ) + chunks = [chunk async for chunk in stream] + span = stream.handler.primary_span + assert span.finished + assert span.get_tag("_dd.llmobs.span_kind") == "llm" + finally: + unpatch_litellm() + + assert len(chunks) == (3 if stream_options and stream_options["include_usage"] else 2) + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + if stream_options and stream_options["include_usage"]: + assert chunks[-1].choices[0].delta.content is None + assert chunks[-1].usage.total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") == 114 + else: + from litellm.litellm_core_utils.streaming_handler import calculate_total_usage + + assert all(getattr(chunk, "usage", None) is None for chunk in chunks) + assert calculate_total_usage(chunks=stream.chunks).total_tokens == 114 + assert span.get_metric("_dd.llmobs.total_tokens") is None + @pytest.mark.asyncio async def test_acompletion_filters_stream_and_stream_options(self): """ From 491d454826342aa8b53aa69edd0242ba3b6f8b4d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 04:51:56 +0000 Subject: [PATCH 24/51] fix(responses): emit the reasoning item on streaming /v1/responses for signature-only thinking (#43414) * fix(responses): emit the reasoning item on streaming /v1/responses for signature-only thinking Anthropic models return thinking blocks with empty text and the reasoning carried in the signature: Claude Fable 5.1 and Claude Opus 5.5 by default, and Bedrock adaptive thinking with or without an effort. On streaming /v1/responses the chat->Responses bridge opened a reasoning output item only on reasoning_content text (LiteLLMCompletionStreamingIterator._ensure_output_item_for_chunk), and ChunkProcessor.get_combined_thinking_content kept an assembled thinking block only when it had thinking text. Such a response emitted no reasoning item mid-stream and none in response.completed, so a streaming Responses client could not replay the reasoning even though the reasoning tokens were billed. Non-streaming /v1/responses was unaffected. Open the reasoning item when the delta carries a signed or redacted thinking block, and keep a signed block through stream assembly even when its thinking text is empty. Unsigned text-only fragments are still dropped. The reasoning-text path is unchanged. (cherry picked from commit bc9b6f8a5c3ac9a2b46e3f9f01f7c2c5f9b688e7) * test(vertex_ai): move orphaned gemma streaming tests into the llm-vertex-ai shard PR #43147 left a copy of the Gemma streaming tests under tests/test_litellm/llms, a tree no CI shard claims, which broke assert-ci-coverage and assert-shard-coverage on main. Fold the two streaming tests into the existing tests/unit/llms/vertex_ai file so the llm-vertex-ai shard runs them Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Chloe Lu Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../streaming_chunk_builder_utils.py | 2 +- .../streaming_iterator.py | 7 +- .../test_vertex_gemma_transformation.py | 93 ------------------- .../test_streaming_chunk_builder_utils.py | 25 +++++ .../test_vertex_gemma_transformation.py | 78 ++++++++++++++++ .../test_streaming_iterator_transformation.py | 40 ++++++++ 6 files changed, 150 insertions(+), 95 deletions(-) delete mode 100644 tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index d975c3551f3..67684a230e3 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -685,7 +685,7 @@ class ChunkProcessor: def _flush_thinking_block() -> None: nonlocal current_thinking_text_parts, current_signature - if len(current_thinking_text_parts) > 0 and current_signature: + if current_signature: thinking_blocks.append( ChatCompletionThinkingBlock( type="thinking", diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 5173cd04a89..21a33c17ab8 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -73,6 +73,11 @@ def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | ) +def _delta_has_signed_thinking_block(delta: object) -> bool: + blocks: Final = getattr(delta, "thinking_blocks", None) or () + return any(isinstance(b, dict) and (b.get("signature") or b.get("data")) for b in blocks) + + class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): """ Async iterator for processing streaming responses from the Responses API. @@ -936,7 +941,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.sent_output_item_added_event = True # Reasoning-first - if hasattr(delta, "reasoning_content") and delta.reasoning_content: + if (hasattr(delta, "reasoning_content") and delta.reasoning_content) or _delta_has_signed_thinking_block(delta): self._reasoning_active = True if self._cached_reasoning_item_id is None: self._cached_reasoning_item_id = f"rs_{uuid.uuid4()}" diff --git a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py deleted file mode 100644 index 294d26b2e58..00000000000 --- a/tests/test_litellm/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ /dev/null @@ -1,93 +0,0 @@ -import json -from collections.abc import AsyncIterator -from types import SimpleNamespace -from typing import Any, cast - -import httpx -import pytest - -import litellm -from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper -from litellm.main import vertex_gemma_chat_completion -from litellm.types.llms.openai import OutputTextDeltaEvent, ResponseCompletedEvent, ResponsesAPIStreamingResponse - -_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" -_MESSAGES = [{"role": "user", "content": "Reply exactly READY"}] -_FAKE_CREDENTIALS = "gemma-test-credentials" - - -def _vertex_response(): - return { - "predictions": { - "id": "chatcmpl-stream-test", - "created": 1759863903, - "model": "google/gemma-3-12b-it", - "object": "chat.completion", - "choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "READY"}}], - "usage": {"prompt_tokens": 14, "completion_tokens": 1, "total_tokens": 15}, - } - } - - -@pytest.fixture(autouse=True) -def _cached_access_token(): - """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" - cache = vertex_gemma_chat_completion._credentials_project_mapping - key = (_FAKE_CREDENTIALS, "test") - cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") - yield - cache.pop(key, None) - - -def test_sync_gemma_stream(): - captured: dict[str, Any] = {} - - def handle(request: httpx.Request) -> httpx.Response: - captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_vertex_response()) - - stream = litellm.completion( - model="vertex_ai/gemma/test-model", - messages=_MESSAGES, - stream=True, - api_base=_VERTEX_URL, - vertex_project="test", - vertex_location="us-central1", - vertex_credentials=_FAKE_CREDENTIALS, - client=httpx.Client(transport=httpx.MockTransport(handle)), - ) - - assert isinstance(stream, CustomStreamWrapper) - chunks = list(stream) - - assert "stream" not in captured["body"]["instances"][0] - assert len(chunks) == 2 - assert chunks[0].choices[0].delta.content == "READY" - assert chunks[1].choices[0].finish_reason == "stop" - - -@pytest.mark.asyncio -async def test_async_gemma_responses_stream(): - captured: dict[str, Any] = {} - - def handle(request: httpx.Request) -> httpx.Response: - captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_vertex_response()) - - response = await litellm.aresponses( - model="vertex_ai/gemma/test-model", - input="Reply exactly READY", - stream=True, - api_base=_VERTEX_URL, - vertex_project="test", - vertex_location="us-central1", - vertex_credentials=_FAKE_CREDENTIALS, - client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), - ) - events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] - - assert "stream" not in captured["body"]["instances"][0] - assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) - assert isinstance(events[-1], ResponseCompletedEvent) - assert events[-1].response.usage is not None - assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py index af763da2d87..aaf877df364 100644 --- a/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -236,6 +236,31 @@ def test_get_combined_thinking_content_preserves_interleaved_blocks(): assert result[2]["signature"] == "sig_block2" +def test_get_combined_thinking_content_keeps_signed_block_without_thinking_text(): + chunks: Final = [ + ModelResponseStream( + id="chatcmpl-123", + object="chat.completion.chunk", + created=1234567890, + model="claude-sonnet-4-20250514", + choices=[ + StreamingChoices( + index=0, + delta=Delta(thinking_blocks=[{"type": "thinking", "thinking": "", "signature": "sig_only"}]), + finish_reason=None, + ) + ], + ) + ] + + result: Final = ChunkProcessor(chunks=chunks).get_combined_thinking_content(chunks) + + assert result is not None + assert [(block["type"], block["thinking"], block["signature"]) for block in result] == [ + ("thinking", "", "sig_only") + ] + + def test_cache_read_input_tokens_retained(): chunk1 = ModelResponseStream( id="chatcmpl-95aabb85-c39f-443d-ae96-0370c404d70c", diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index 97f4f290958..92e684e42c8 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -1303,3 +1303,81 @@ class TestVertexGemmaCompletion: mock_async_post.assert_awaited_once() assert mock_async_post.call_args.kwargs["client"] is None assert response.choices[0].message.content == "default async handler fallback" + + +_GEMMA_VERTEX_URL = "https://example.invalid/v1/projects/test/locations/us-central1/endpoints/test:predict" +_FAKE_GEMMA_CREDENTIALS = "gemma-test-credentials" + + +@pytest.fixture +def _gemma_cached_access_token(): + """Serve a fake token from the handler's credential cache so no auth round-trip runs.""" + from types import SimpleNamespace + + from litellm.main import vertex_gemma_chat_completion + + cache = vertex_gemma_chat_completion._credentials_project_mapping + key = (_FAKE_GEMMA_CREDENTIALS, "test") + cache[key] = (SimpleNamespace(token="fake-token", expired=False), "test") + yield + cache.pop(key, None) + + +def test_sync_gemma_stream(_gemma_cached_access_token): + import httpx + + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + + captured = {} + + def handle(request): + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + + stream = litellm.completion( + model="vertex_ai/gemma/test-model", + messages=[{"role": "user", "content": "Reply exactly READY"}], + stream=True, + api_base=_GEMMA_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_GEMMA_CREDENTIALS, + client=httpx.Client(transport=httpx.MockTransport(handle)), + ) + + assert isinstance(stream, CustomStreamWrapper) + chunks = list(stream) + + assert "stream" not in captured["body"]["instances"][0] + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "READY" + assert chunks[1].choices[0].finish_reason == "stop" + + +@pytest.mark.asyncio +async def test_async_gemma_responses_stream(_gemma_cached_access_token): + import httpx + + captured = {} + + def handle(request): + captured["body"] = json.loads(request.content) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + + response = await litellm.aresponses( + model="vertex_ai/gemma/test-model", + input="Reply exactly READY", + stream=True, + api_base=_GEMMA_VERTEX_URL, + vertex_project="test", + vertex_location="us-central1", + vertex_credentials=_FAKE_GEMMA_CREDENTIALS, + client=httpx.AsyncClient(transport=httpx.MockTransport(handle)), + ) + events = [event async for event in cast(AsyncIterator[ResponsesAPIStreamingResponse], response)] + + assert "stream" not in captured["body"]["instances"][0] + assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) + assert isinstance(events[-1], ResponseCompletedEvent) + assert events[-1].response.usage is not None + assert events[-1].response.usage.total_tokens == 15 diff --git a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py index 8fbba0dbf87..041bcf1b6d7 100644 --- a/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py +++ b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py @@ -978,6 +978,25 @@ def _reasoning_chunk(reasoning: str, finish_reason: str | None = None) -> ModelR ) +def _signature_only_thinking_chunk(signature: str) -> ModelResponseStream: + return ModelResponseStream( + id=CHAT_COMPLETION_ID, + created=1748575031, + model="claude-haiku-4-5", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta( + role="assistant", + thinking_blocks=[{"type": "thinking", "thinking": "", "signature": signature}], + ), + finish_reason=None, + ) + ], + ) + + async def _collect_events( iterator: LiteLLMCompletionStreamingIterator, sync_mode: bool ) -> list[BaseLiteLLMOpenAIResponseObject]: @@ -1015,6 +1034,27 @@ async def test_tool_only_stream_emits_no_message_item_events(sync_mode: bool): assert any(getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED for event in events) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_signature_only_thinking_streams_a_replayable_reasoning_item(sync_mode: bool): + iterator: Final = _build_iterator([_signature_only_thinking_chunk("sig_only"), _chunk("4", finish_reason="stop")]) + + events: Final = await _collect_events(iterator, sync_mode) + + added_item_types: Final = [ + event.item.type + for event in events + if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + ] + completed: Final = next( + event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + ) + reasoning_items: Final = [item for item in completed.response.output if getattr(item, "type", None) == "reasoning"] + assert added_item_types[0] == "reasoning" + assert len(reasoning_items) == 1 + assert json.loads(reasoning_items[0].encrypted_content)[0]["signature"] == "sig_only" + + @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio async def test_reasoning_then_text_announces_message_item_before_text_events(sync_mode: bool): From 9a0ff249d5935ca73216d19597603083e6a0845c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 05:02:35 +0000 Subject: [PATCH 25/51] fix(anthropic): forward the per-turn-control beta to Azure AI Foundry (#43415) Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/anthropic_beta_headers_config.json | 2 +- .../messages/test_anthropic_messages_per_turn_control.py | 8 +++++++- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index 3a28d65e47c..71e7081b440 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -57,7 +57,7 @@ "mcp-servers-2025-12-04": null, "output-128k-2025-02-19": null, "structured-output-2024-03-01": null, - "per-turn-control-2026-07-01": null, + "per-turn-control-2026-07-01": "per-turn-control-2026-07-01", "prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05", "skills-2025-10-02": "skills-2025-10-02", "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py index 557305a945c..4197192e4af 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py @@ -95,13 +95,19 @@ def test_added_per_turn_control_beta_survives_the_anthropic_allowlist(): assert PER_TURN_CONTROL in _betas(filtered) -@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "azure_ai", "databricks"]) +@pytest.mark.parametrize("provider", ["bedrock", "bedrock_converse", "vertex_ai", "databricks"]) def test_per_turn_control_beta_is_dropped_for_providers_without_it(provider): filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider=provider) assert "anthropic-beta" not in filtered +def test_per_turn_control_beta_is_forwarded_for_azure_ai(): + filtered = update_headers_with_filtered_beta(headers={"anthropic-beta": PER_TURN_CONTROL}, provider="azure_ai") + + assert _betas(filtered) == {PER_TURN_CONTROL} + + def test_json_provider_passthrough_adds_per_turn_control_beta(): config = JSONProviderAnthropicMessagesConfig( SimpleProviderConfig( From c1f761eba50bb344259bb7a3ff2ef538ff94600f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 05:02:44 +0000 Subject: [PATCH 26/51] test(vertex_ai): move stray Gemma streaming tests to tests/unit so CI coverage passes (#43422) Co-authored-by: Krrish Dholakia Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../vertex_gemma_models/test_vertex_gemma_transformation.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index 92e684e42c8..efe97ce33a8 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -1332,7 +1332,7 @@ def test_sync_gemma_stream(_gemma_cached_access_token): def handle(request): captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY")) stream = litellm.completion( model="vertex_ai/gemma/test-model", @@ -1362,7 +1362,7 @@ async def test_async_gemma_responses_stream(_gemma_cached_access_token): def handle(request): captured["body"] = json.loads(request.content) - return httpx.Response(200, json=_make_gemma_vertex_response(content="READY", total_tokens=15)) + return httpx.Response(200, json=_make_gemma_vertex_response(content="READY")) response = await litellm.aresponses( model="vertex_ai/gemma/test-model", @@ -1380,4 +1380,4 @@ async def test_async_gemma_responses_stream(_gemma_cached_access_token): assert "READY" in "".join(event.delta for event in events if isinstance(event, OutputTextDeltaEvent)) assert isinstance(events[-1], ResponseCompletedEvent) assert events[-1].response.usage is not None - assert events[-1].response.usage.total_tokens == 15 + assert events[-1].response.usage.total_tokens == 114 From b831e9b4ac8a1f221704663acb9cb542ba40fd11 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 05:10:24 +0000 Subject: [PATCH 27/51] fix(bedrock): keep the provider status code on unprocessable image errors (#43416) Co-authored-by: Krrish Dholakia Co-authored-by: dbalintx Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../exception_mapping_utils.py | 2 +- .../test_exception_mapping_utils.py | 37 ++++++++++++++++++- 2 files changed, 37 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index f09dd9fe75a..0fdfb301291 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -954,7 +954,7 @@ def _map_bedrock_exception( llm_provider="bedrock", response=getattr(original_exception, "response", None), ) - elif "Could not process image" in error_str: + elif "Could not process image" in error_str and getattr(original_exception, "status_code", 500) == 500: raise litellm.InternalServerError( message=f"BedrockException - {error_str}", model=model, diff --git a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py index 9fce0441a58..9de768ea47b 100644 --- a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py @@ -1280,7 +1280,7 @@ def test_bedrock_500_preserves_provider_response_headers(): "bedrock", 400, '{"message":"Could not process image"}', - litellm.InternalServerError, + litellm.BadRequestError, ), ], ) @@ -1313,6 +1313,41 @@ def test_bedrock_classified_errors_preserve_provider_response_headers( assert exc_info.value.response.headers["x-amzn-requestid"] == "req-classified" +@pytest.mark.parametrize( + "status_code, expected_exception", + [ + (400, litellm.BadRequestError), + (503, litellm.ServiceUnavailableError), + (500, litellm.InternalServerError), + ], +) +def test_bedrock_unprocessable_image_keeps_provider_status_code(status_code, expected_exception): + """An unprocessable image maps to the status Bedrock sent, so the 400 it returns stays a client error.""" + provider_message = '{"message":"The model returned the following errors: Could not process image"}' + provider_response = httpx.Response( + status_code=status_code, + text=provider_message, + request=httpx.Request("POST", "https://bedrock-runtime.us-east-1.amazonaws.com/"), + ) + original_exception = BedrockError( + status_code=status_code, + message=provider_message, + headers=provider_response.headers, + response=provider_response, + ) + + with pytest.raises(expected_exception) as exc_info: + exception_type( + model="anthropic.claude-haiku-4-5-20251001-v1:0", + original_exception=original_exception, + custom_llm_provider="bedrock", + completion_kwargs={}, + extra_kwargs={}, + ) + + assert exc_info.value.status_code == status_code + + @pytest.mark.parametrize( "status_code, provider_message", [ From 4274bdda441527c8ab9601c44e4dbbc79a63e707 Mon Sep 17 00:00:00 2001 From: Anmol Jaiswal <68013660+anmolg1997@users.noreply.github.com> Date: Sun, 27 Sep 2026 10:44:53 +0530 Subject: [PATCH 28/51] fix(vertex_ai): stop importing the vertexai SDK in partner-model completion (#42274) completion() imported vertexai only to check that the package exists. Partner models are reached with an authenticated httpx client and never use that SDK, the same reasoning count_tokens in this file already follows (#28084). The import loads all of google-cloud-aiplatform on the first request of every process and made a google-auth-only install fail with a 400 --- .../vertex_ai_partner_models/main.py | 9 +---- .../test_partner_models_credential_reuse.py | 38 +++++++++++++++++++ 2 files changed, 39 insertions(+), 8 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py index 2a36e5cc785..40503edbb9e 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py @@ -109,8 +109,6 @@ class VertexAIPartnerModels(VertexBase): client=None, ): try: - import vertexai - from litellm.llms.anthropic.chat import AnthropicChatCompletion from litellm.llms.codestral.completion.handler import ( CodestralTextCompletion, @@ -119,14 +117,9 @@ class VertexAIPartnerModels(VertexBase): except Exception as e: raise VertexAIError( status_code=400, - message=f"""vertexai import failed please run `pip install -U "google-cloud-aiplatform>=1.38"`. Got error: {e}""", + message=f"Failed to import a partner model handler. Got error: {e}", ) - if not (hasattr(vertexai, "preview") or hasattr(vertexai.preview, "language_models")): - raise VertexAIError( - status_code=400, - message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""", - ) try: access_token, project_id = self._ensure_access_token( credentials=vertex_credentials, diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py index b20442a032e..8e6270e41a1 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/test_partner_models_credential_reuse.py @@ -127,6 +127,44 @@ class TestPartnerModelsCredentialReuse: assert mock_load.call_count == 1 + def test_completion_works_without_the_vertexai_sdk(self): + """completion() reaches the HTTP handler when `import vertexai` raises ImportError.""" + partner = VertexAIPartnerModels() + + with ( + patch.dict(sys.modules, {"vertexai": None}), + patch.object( + partner, + "_ensure_access_token", + return_value=("cached-token", "test-project"), + ), + patch( + "litellm.llms.vertex_ai.vertex_ai_partner_models.main.base_llm_http_handler" + ) as mock_handler, + ): + mock_handler.completion.return_value = "response" + + result = partner.completion( + model="meta/llama-3.1-405b-instruct-maas", + messages=[{"role": "user", "content": "hello"}], + model_response=MagicMock(), + print_verbose=lambda *a, **kw: None, + encoding=MagicMock(), + logging_obj=MagicMock(), + api_base=None, + optional_params={}, + custom_prompt_dict={}, + headers=None, + timeout=30.0, + litellm_params={}, + vertex_project="test-project", + vertex_location="us-central1", + vertex_credentials=None, + ) + + assert result == "response" + mock_handler.completion.assert_called_once() + class TestGemmaModelsCredentialReuse: def test_completion_uses_self_ensure_access_token(self): From 8e6d99d74a63c61e39628baeabea0c07dbeda5f4 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 22:28:31 -0700 Subject: [PATCH 29/51] fix(token_counter): count Gemini function_declarations tools (#43417) * fix(token_counter): count Gemini function_declarations tools Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(token_counter): skip non-dict tools when formatting definitions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/litellm_core_utils/token_counter.py | 78 +++++++++++-------- .../litellm_core_utils/test_token_counter.py | 72 +++++++++++++++++ .../test_vertex_ai_context_caching.py | 9 ++- 3 files changed, 126 insertions(+), 33 deletions(-) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 5d7956059e4..cdd2d0654be 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -951,7 +951,7 @@ def _count_content_list( ) -def _format_function_definitions(tools): +def _format_function_definitions(tools: Sequence[object]) -> str: """Formats tool definitions in the format that OpenAI appears to use. Based on https://github.com/forestwanglin/openai-java/blob/main/jtokkit/src/main/java/xyz/felh/openai/jtokkit/utils/TikTokenUtils.java """ @@ -959,41 +959,57 @@ def _format_function_definitions(tools): lines.append("namespace functions {") lines.append("") for tool in tools: - if not isinstance(tool, dict): + if not isinstance(tool, Mapping): continue - function = tool.get("function") - if not isinstance(function, dict): - # Anthropic tool shape → OpenAI function dict for token counting. - params = tool.get("input_schema") or tool.get("parameters") or {} - if not isinstance(params, dict): - params = {} - function = { - "name": tool.get("name"), - "description": tool.get("description"), - "parameters": params, - } - function_name = function.get("name") - if not function_name: - # Skip malformed tools missing a name to avoid emitting - # ``type None = ...`` which would produce inaccurate token counts. - continue - if function_description := function.get("description"): - lines.append(f"// {function_description}") - parameters = function.get("parameters") or {} - if not isinstance(parameters, dict): - parameters = {} - properties = parameters.get("properties") - if properties and properties.keys(): - lines.append(f"type {function_name} = (_: {{") - lines.append(_format_object_parameters(parameters, 0)) - lines.append("}) => any;") - else: - lines.append(f"type {function_name} = () => any;") - lines.append("") + for function in _function_definitions_for_tool(cast(Mapping[str, object], tool)): + lines.extend(_format_single_function_definition(function)) lines.append("} // namespace functions") return "\n".join(lines) +def _function_definitions_for_tool(tool: Mapping[str, object]) -> Iterable[Mapping[str, object]]: + function: Final = tool.get("function") + if isinstance(function, Mapping): + yield function + return + declarations: Final = tool.get("function_declarations") or tool.get("functionDeclarations") + if isinstance(declarations, list): + for declaration in declarations: + if isinstance(declaration, Mapping): + yield declaration + return + parameters: Final = tool.get("input_schema") or tool.get("parameters") or {} + normalized_parameters: Final = parameters if isinstance(parameters, Mapping) else {} + yield { + "name": tool.get("name"), + "description": tool.get("description"), + "parameters": normalized_parameters, + } + + +def _format_single_function_definition(function: Mapping[str, object]) -> tuple[str, ...]: + function_name: Final = function.get("name") + if not function_name: + return () + function_description: Final = function.get("description") + parameters_value: Final = function.get("parameters") or {} + parameters: Final = parameters_value if isinstance(parameters_value, Mapping) else {} + properties: Final = parameters.get("properties") + if isinstance(properties, Mapping) and properties: + return ( + *((f"// {function_description}",) if function_description else ()), + f"type {function_name} = (_: {{", + _format_object_parameters(parameters, 0), + "}) => any;", + "", + ) + return ( + *((f"// {function_description}",) if function_description else ()), + f"type {function_name} = () => any;", + "", + ) + + def _format_object_parameters(parameters, indent): properties: Final = parameters.get("properties") if not properties: diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index f7ded4f3fa8..c71b1496bdd 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -442,6 +442,78 @@ def test_token_counter_with_tools(message_count_pair): ), f"Expected {expected_tokens} tokens, got {counted_tokens}." +def test_token_counter_counts_gemini_function_declarations(): + openai_tools: Final = [ + { + "type": "function", + "function": { + "name": "lookup_weather", + "description": "Find current weather conditions for a location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string", "description": "City and region"}, + "units": {"type": "string", "enum": ["celsius", "fahrenheit"]}, + }, + "required": ["location"], + }, + }, + } + ] + gemini_tools: Final = litellm.utils.get_optional_params( + model="gemini-2.5-pro", + custom_llm_provider="gemini", + tools=openai_tools, + )["tools"] + camel_case_tools: Final = [{"functionDeclarations": gemini_tools[0]["function_declarations"]}] + + openai_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=[{"role": "user", "content": "What's the weather?"}], + tools=openai_tools, + ) + gemini_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=[{"role": "user", "content": "What's the weather?"}], + tools=gemini_tools, + ) + camel_case_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=[{"role": "user", "content": "What's the weather?"}], + tools=camel_case_tools, + ) + + assert openai_tokens == gemini_tokens == camel_case_tokens + + +def test_token_counter_skips_non_mapping_tools(): + openai_tool: Final = { + "type": "function", + "function": { + "name": "lookup_weather", + "description": "Find current weather conditions for a location", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string", "description": "City and region"}}, + "required": ["location"], + }, + }, + } + messages: Final = [{"role": "user", "content": "What's the weather?"}] + valid_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=messages, + tools=[openai_tool], + ) + mixed_tokens: Final = token_counter_new( + model="gemini-2.5-pro", + messages=messages, + tools=["bad", None, openai_tool], + ) + + assert mixed_tokens == valid_tokens + + class NeedsToleranceUpdateError(Exception): """Custom exception to mark tests that have improved""" diff --git a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 283ed3710d0..67d78d6030e 100644 --- a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1,4 +1,4 @@ -from typing import List +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -1530,7 +1530,7 @@ class TestContextCachingEndpoints: ] all_messages = short_cached_messages + non_cached_messages - large_tools = [ + openai_large_tools: Final = [ { "type": "function", "function": { @@ -1548,6 +1548,11 @@ class TestContextCachingEndpoints: } for i in range(12) ] + large_tools: Final = litellm.utils.get_optional_params( + model="gemini-1.5-pro", + custom_llm_provider="gemini", + tools=openai_large_tools, + )["tools"] optional_params = { **self.sample_optional_params, From f4308bc124eebc783dfc51790ce8db27ed21ae00 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 01:28:02 -0700 Subject: [PATCH 30/51] refactor(types): replace Any with proven types in 5 files (#43304) * refactor(types): replace Any with proven types in 6 files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(types): keep enterprise email import inside try-except for unsafe-import check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(types): keep email_logging_instance annotation as Any pending a guarded alias Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): revert iterator override typing in proxy utils Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/assistants/main.py | 14 +++++++------- litellm/litellm_core_utils/litellm_logging.py | 10 +++++----- litellm/llms/custom_httpx/llm_http_handler.py | 16 +++++++++------- litellm/proxy/common_request_processing.py | 8 +++++--- litellm/utils.py | 2 +- 5 files changed, 27 insertions(+), 23 deletions(-) diff --git a/litellm/assistants/main.py b/litellm/assistants/main.py index 1ce40e94320..c14c4aec093 100644 --- a/litellm/assistants/main.py +++ b/litellm/assistants/main.py @@ -3,7 +3,7 @@ import asyncio import contextvars import os -from collections.abc import Coroutine, Iterable +from collections.abc import Coroutine, Iterable, Mapping, Sequence from functools import partial from typing import Any, Final, Literal @@ -233,8 +233,8 @@ def create_assistants( name: str | None = None, description: str | None = None, instructions: str | None = None, - tools: list[dict[str, Any]] | None = None, - tool_resources: dict[str, Any] | None = None, + tools: Sequence[Mapping[str, object]] | None = None, + tool_resources: Mapping[str, object] | None = None, metadata: dict[str, str] | None = None, temperature: float | None = None, top_p: float | None = None, @@ -244,7 +244,7 @@ def create_assistants( api_base: str | None = None, api_version: str | None = None, **kwargs, -) -> Assistant | Coroutine[Any, Any, Assistant]: +) -> Assistant | Coroutine[None, None, Assistant]: async_create_assistants: Final[bool | None] = kwargs.pop("async_create_assistants", None) if async_create_assistants is not None and not isinstance(async_create_assistants, bool): raise ValueError("Invalid value passed in for async_create_assistants. Only bool or None allowed") @@ -283,7 +283,7 @@ def create_assistants( # only send params that are not None create_assistant_data = {k: v for k, v in create_assistant_data.items() if v is not None} - response: Coroutine[Any, Any, Assistant] | Assistant | None = None + response: Coroutine[None, None, Assistant] | Assistant | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -415,7 +415,7 @@ def delete_assistant( api_base: str | None = None, api_version: str | None = None, **kwargs, -) -> AssistantDeleted | Coroutine[Any, Any, AssistantDeleted]: +) -> AssistantDeleted | Coroutine[None, None, AssistantDeleted]: optional_params: Final = GenericLiteLLMParams(api_key=api_key, api_base=api_base, api_version=api_version, **kwargs) litellm_params_dict: Final = get_litellm_params(**kwargs) @@ -440,7 +440,7 @@ def delete_assistant( elif timeout is None: timeout = 600.0 - response: AssistantDeleted | Coroutine[Any, Any, AssistantDeleted] | None = None + response: AssistantDeleted | Coroutine[None, None, AssistantDeleted] | None = None if custom_llm_provider == "openai": api_base = ( optional_params.api_base diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 152fd54e55d..5b4187846ff 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -640,8 +640,8 @@ class Logging(LiteLLMLoggingBaseClass): self._own_session_id: str = session_id_var.get() self.function_id = function_id - self.streaming_chunks: list[Any] = [] # for generating complete stream response - self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response + self.streaming_chunks: list[object] = [] # for generating complete stream response + self.sync_streaming_chunks: list[object] = [] # for generating complete stream response self.log_raw_request_response = log_raw_request_response self.raw_request_only = raw_request_only @@ -693,7 +693,7 @@ class Logging(LiteLLMLoggingBaseClass): self.response_timing_metrics: Mapping[str, float] = {} # mutable-ok: kept deep-copyable # Passthrough endpoint guardrails config for field targeting - self.passthrough_guardrails_config: dict[str, Any] | None = None + self.passthrough_guardrails_config: dict[str, object] | None = None self.model_call_details: dict[str, Any] = { "litellm_trace_id": self.litellm_trace_id, @@ -4479,7 +4479,7 @@ def set_callbacks(callback_list, function_id=None): def _init_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, internal_usage_cache: DualCache | None, - llm_router: Any | None, # expect litellm.Router, but typing errors due to circular import + llm_router: object, # expect litellm.Router, but typing errors due to circular import custom_logger_init_args: dict | None = {}, ) -> CustomLogger | None: """ @@ -6439,7 +6439,7 @@ def _autorouter_savings_for_payload( def get_standard_logging_object_payload( kwargs: dict | None, - init_response_obj: Any | BaseModel | dict, + init_response_obj: object, start_time: dt_object, end_time: dt_object, logging_obj: Logging, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 8aa38ff3341..ce9f7a2ea54 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6375,7 +6375,7 @@ class BaseLLMHTTPHandler: custom_llm_provider: str | None = None, first_message: str | None = None, request_defaults: ResponsesWebSocketRequestDefaults | None = None, - **kwargs: Any, + **kwargs: object, ) -> Exception | None: """ Handles Responses API WebSocket mode. @@ -10378,13 +10378,14 @@ class BaseLLMHTTPHandler: encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url: Final = f"{api_base}/{encoded_vector_store_id}" - request_body: Final[dict[str, Any]] = dict(vector_store_update_optional_params) + request_body: Final[dict[str, object]] = dict(vector_store_update_optional_params) + metadata: Final = vector_store_update_optional_params.get("metadata") # Clean metadata to only include string values (OpenAI requirement) - if "metadata" in request_body and request_body["metadata"] is not None: + if metadata is not None: from litellm.utils import add_openai_metadata - request_body["metadata"] = add_openai_metadata(request_body["metadata"]) + request_body["metadata"] = add_openai_metadata(metadata) if extra_body: request_body.update(extra_body) @@ -10456,13 +10457,14 @@ class BaseLLMHTTPHandler: encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id") url: Final = f"{api_base}/{encoded_vector_store_id}" - request_body: Final[dict[str, Any]] = dict(vector_store_update_optional_params) + request_body: Final[dict[str, object]] = dict(vector_store_update_optional_params) + metadata: Final = vector_store_update_optional_params.get("metadata") # Clean metadata to only include string values (OpenAI requirement) - if "metadata" in request_body and request_body["metadata"] is not None: + if metadata is not None: from litellm.utils import add_openai_metadata - request_body["metadata"] = add_openai_metadata(request_body["metadata"]) + request_body["metadata"] = add_openai_metadata(metadata) if extra_body: request_body.update(extra_body) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 64b0c6c1967..15610da9aec 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -739,7 +739,7 @@ async def _parse_event_data_for_error(event_line: str | bytes) -> int | None: if not json_str or json_str == "[DONE]": # handle empty data or [DONE] message return None try: - data: Final = orjson.loads(json_str) + data: Final[object] = orjson.loads(json_str) if isinstance(data, dict) and "error" in data and isinstance(data["error"], dict): error_code_raw: Final = data["error"].get("code") error_code: int | None = None @@ -792,7 +792,7 @@ def _extract_error_from_sse_chunk(event_line: str | bytes) -> dict: return default_error try: - data: Final = orjson.loads(json_str) + data: Final[object] = orjson.loads(json_str) if isinstance(data, dict) and "error" in data: error_obj: Final = data["error"] if isinstance(error_obj, dict): @@ -4131,7 +4131,9 @@ class ProxyBaseLLMRequestProcessing: if stripped_ln.startswith("data:"): json_part = stripped_ln.split("data:", 1)[1].strip() if json_part and json_part != "[DONE]": - obj = json.loads(json_part) + obj: object = json.loads(json_part) + if not isinstance(obj, dict): + return None maybe_modified = ProxyBaseLLMRequestProcessing._inject_cost_into_usage_dict( obj, model_name, litellm_logging_obj ) diff --git a/litellm/utils.py b/litellm/utils.py index 7ce412e818c..d45b29c0f16 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1469,7 +1469,7 @@ async def async_pre_call_deployment_hook(kwargs: dict[str, Any], call_type: str) async def async_post_call_success_deployment_hook( request_data: dict, response: object, call_type: CallTypes | None -) -> Any | None: +) -> object: """ Allow modifying / reviewing the response just after it's received from the deployment. """ From ff462f7a77a5af4da86129265692530dbf04fe69 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 08:56:50 -0700 Subject: [PATCH 31/51] chore(cost-map): update azure_ai/grok-4.6 input price from Azure pricing page (#43440) --- litellm/model_prices_and_context_window_backup.json | 4 ++-- model_prices_and_context_window.json | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 09fc442e5a7..5362b1b042c 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -12257,7 +12257,7 @@ "azure_ai/grok-4.6": { "cache_read_input_token_cost": 5e-07, "cache_read_input_token_cost_above_200k_tokens": 1e-06, - "input_cost_per_token": 2e-06, + "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "azure_ai", "max_input_tokens": 200000, @@ -12266,7 +12266,7 @@ "mode": "chat", "output_cost_per_token": 6e-06, "output_cost_per_token_above_200k_tokens": 1.2e-05, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 09fc442e5a7..5362b1b042c 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -12257,7 +12257,7 @@ "azure_ai/grok-4.6": { "cache_read_input_token_cost": 5e-07, "cache_read_input_token_cost_above_200k_tokens": 1e-06, - "input_cost_per_token": 2e-06, + "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 4e-06, "litellm_provider": "azure_ai", "max_input_tokens": 200000, @@ -12266,7 +12266,7 @@ "mode": "chat", "output_cost_per_token": 6e-06, "output_cost_per_token_above_200k_tokens": 1.2e-05, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/grok/", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, From 22b36cbcf6583e2d6b552cc0e87ae6ab82c46341 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 09:14:33 -0700 Subject: [PATCH 32/51] chore(cost-map): update azure_ai/grok-4.6 input price and add azure_ai/MAI-Cyber-1-Flash (#43446) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- .../model_prices_and_context_window_backup.json | 15 +++++++++++++++ model_prices_and_context_window.json | 15 +++++++++++++++ 2 files changed, 30 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 5362b1b042c..b36d84ea027 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -77895,5 +77895,20 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_tool_choice": false + }, + "azure_ai/MAI-Cyber-1-Flash": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 6e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 256000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 3.5e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/microsoft/", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 5362b1b042c..b36d84ea027 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -77895,5 +77895,20 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_tool_choice": false + }, + "azure_ai/MAI-Cyber-1-Flash": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 6e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 256000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 3.5e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/microsoft/", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_tool_choice": true } } From 268e8bb735b6871bfed8e593be1b0b53e277d949 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 14:53:12 -0700 Subject: [PATCH 33/51] refactor(rust): share anthropic types, request helpers, and streaming contracts across crates (#43426) * refactor(rust): standardize Azure Messages module path * docs(rust): define shared types crate boundaries * refactor(rust): share request helpers and type Anthropic blocks * docs(rust): format shared type invariants as bullets * test(rust): parameterize repeated cases with rstest * refactor(rust): move Responses transform result into llms * fix(anthropic): validate chat and batch responses * docs(rust): clarify API format ownership boundaries * docs: clarify Rust error message construction * refactor(auth): keep shared Rust errors provider-neutral * refactor(rust): separate format contracts from provider policy * fix(rust): type Anthropic chat response text collection * fix(rust): pass audio secret sources through hosts * fix(rust): unblock batch lint and OCR error assertions * test(rust): assert response failures at the adapter boundary * refactor(rust): declare error messages with typed context * wip * fix(rust): adapt Bedrock error details * style(rust): cargo fmt bedrock audio transcription Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): adapt tests and dead code to typed error details Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): keep converse error contracts and read env secrets without litellm Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci(rust): raise the native wheel size gate to 45 MB Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): tolerate missing usage in converse responses on the transcription route Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/scripts/verify_linux_native_wheel.py | 2 +- litellm-rust/AGENTS.md | 2 + litellm-rust/Cargo.lock | 10 +- litellm-rust/crates/auth-azure/src/native.rs | 74 ++-- litellm-rust/crates/auth-azure/src/resolve.rs | 70 +++- litellm-rust/crates/auth-azure/src/types.rs | 7 +- litellm-rust/crates/auth-gcp/Cargo.toml | 3 + litellm-rust/crates/auth-gcp/src/lib.rs | 51 +-- litellm-rust/crates/auth-types/Cargo.toml | 1 + .../crates/auth-types/src/credential.rs | 16 +- litellm-rust/crates/auth-types/src/error.rs | 173 +++++----- litellm-rust/crates/auth-types/src/http.rs | 8 +- litellm-rust/crates/auth-types/src/lib.rs | 2 +- litellm-rust/crates/auth-types/src/policy.rs | 10 +- litellm-rust/crates/auth-types/tests/error.rs | 74 ++++ .../crates/core-utils/src/call_arguments.rs | 34 +- .../crates/core-utils/src/core_helpers.rs | 34 +- .../crates/core-utils/src/serde_compat.rs | 29 +- .../crates/core-utils/src/settings.rs | 16 + .../crates/core-utils/src/url_utils.rs | 24 +- .../crates/core-utils/tests/settings.rs | 33 ++ litellm-rust/crates/core/AGENTS.md | 2 + litellm-rust/crates/core/Cargo.toml | 2 - .../core/src/audio_transcription/handler.rs | 10 +- .../core/src/audio_transcription/mod.rs | 4 +- .../core/src/audio_transcription/prepare.rs | 13 +- .../core/src/audio_transcription/types.rs | 2 + .../core/src/chat_completions/handler.rs | 13 +- .../crates/core/src/chat_completions/mod.rs | 6 +- .../core/src/chat_completions/prepare.rs | 39 ++- .../crates/core/src/chat_completions/types.rs | 2 + litellm-rust/crates/core/src/error.rs | 35 +- .../crates/core/src/messages/AGENTS.md | 7 + .../crates/core/src/messages/common_utils.rs | 4 +- .../crates/core/src/messages/handler.rs | 19 +- .../crates/core/src/messages/prepare.rs | 17 +- .../crates/core/src/messages/route.rs | 2 +- .../crates/core/src/messages/types.rs | 17 +- litellm-rust/crates/core/src/ocr/document.rs | 34 +- litellm-rust/crates/core/src/ocr/route.rs | 2 +- .../crates/core/src/responses/websocket.rs | 82 +---- .../crates/core/tests/audio_transcription.rs | 71 +++- .../crates/core/tests/chat_completions.rs | 169 +++++++++- .../crates/core/tests/messages/host.rs | 6 +- .../crates/core/tests/messages/request.rs | 126 +++++-- .../crates/core/tests/messages/response.rs | 2 +- .../crates/core/tests/messages/secrets.rs | 8 +- .../crates/core/tests/ocr/azure_ai.rs | 4 +- litellm-rust/crates/cost/Cargo.toml | 1 + litellm-rust/crates/cost/tests/calculation.rs | 28 +- .../src/audio_transcription.rs | 1 + .../gateway-inference/src/chat_completions.rs | 1 + litellm-rust/crates/host/src/machine/auth.rs | 2 +- litellm-rust/crates/http/AGENTS.md | 6 + litellm-rust/crates/http/Cargo.toml | 3 + litellm-rust/crates/http/src/lib.rs | 1 + litellm-rust/crates/http/src/media.rs | 48 +-- litellm-rust/crates/http/src/request.rs | 34 +- litellm-rust/crates/http/src/websocket.rs | 61 ++++ litellm-rust/crates/http/tests/request.rs | 69 ++++ litellm-rust/crates/http/tests/websocket.rs | 59 ++++ litellm-rust/crates/llms/AGENTS.md | 22 +- .../crates/llms/src/anthropic/AGENTS.md | 8 + .../src/anthropic/batches/transformation.rs | 64 +++- .../crates/llms/src/anthropic/chat/handler.rs | 23 +- .../llms/src/anthropic/chat/transformation.rs | 97 ++++-- .../crates/llms/src/anthropic/common_utils.rs | 319 +++++++----------- .../llms/src/anthropic/messages/AGENTS.md | 10 +- .../llms/src/anthropic/messages/handler.rs | 27 +- .../crates/llms/src/anthropic/messages/mod.rs | 1 - .../llms/src/anthropic/messages/thinking.rs | 205 +++++------ .../src/anthropic/messages/transformation.rs | 160 ++++----- .../crates/llms/src/azure_ai/anthropic/mod.rs | 1 - .../crates/llms/src/azure_ai/common_utils.rs | 25 ++ .../llms/src/azure_ai/messages/AGENTS.md | 3 + .../messages}/mod.rs | 1 - .../transformation.rs} | 191 ++++------- litellm-rust/crates/llms/src/azure_ai/mod.rs | 3 +- .../llms/src/azure_ai/ocr/common_utils.rs | 7 +- .../llms/src/azure_ai/ocr/transformation.rs | 6 +- .../audio_transcription/transformation.rs | 18 +- litellm-rust/crates/llms/src/base_llm/auth.rs | 45 +-- .../llms/src/base_llm/chat/transformation.rs | 2 + .../llms/src/base_llm/messages/AGENTS.md | 5 + .../llms/src/base_llm/messages/context.rs | 134 ++++++++ .../crates/llms/src/base_llm/messages/mod.rs | 4 + .../src/base_llm/messages/normalization.rs | 50 +++ .../streaming.rs | 47 ++- .../transformation.rs | 50 ++- litellm-rust/crates/llms/src/base_llm/mod.rs | 2 +- .../src/base_llm/responses/transformation.rs | 101 +----- .../src/bedrock/audio_transcription/mod.rs | 105 ++++-- .../bedrock/chat/converse_transformation.rs | 170 ++++++++-- .../llms/src/bedrock/chat/invoke_handler.rs | 66 ++-- .../llms/src/bedrock/messages/AGENTS.md | 3 + .../anthropic_claude3_transformation.rs | 83 ++--- litellm-rust/crates/llms/src/error.rs | 126 ++++++- litellm-rust/crates/llms/src/lib.rs | 2 +- .../src/openai/responses/transformation.rs | 94 +++++- .../src/openai_like/chat/transformation.rs | 4 + .../llms/src/openai_like/common_utils.rs | 2 +- .../llms/src/vertex_ai/ocr/common_utils.rs | 5 +- .../tests/anthropic_chat_transformation.rs | 61 ++-- .../tests/bedrock_converse_transformation.rs | 78 +++-- .../llms/tests/messages_normalization.rs | 48 +++ .../tests/openai_like_chat_transformation.rs | 2 +- .../crates/python-bridge/src/coercion.rs | 49 ++- .../crates/python-bridge/src/credentials.rs | 45 +-- .../crates/python-bridge/src/errors.rs | 4 +- .../src/routes/audio_transcription.rs | 20 +- .../src/routes/chat_completions.rs | 6 + .../python-bridge/src/routes/messages/host.rs | 4 +- .../python-bridge/src/secrets/config.rs | 10 +- .../crates/python-bridge/src/secrets/mod.rs | 8 +- .../src/secret_manager/client.rs | 25 +- .../crates/token-counter-fast/src/error.rs | 16 +- .../crates/token-counter-fast/src/lib.rs | 2 +- .../crates/token-counter-fast/src/tiktoken.rs | 31 +- .../token-counter-huggingface/Cargo.toml | 3 + .../token-counter-huggingface/src/lib.rs | 18 +- .../crates/token-counter-tiktoken/Cargo.toml | 3 + .../crates/token-counter-tiktoken/src/lib.rs | 49 ++- .../token-counter-tiktoken/src/ranks.rs | 19 +- .../crates/token-counter/src/error.rs | 2 +- litellm-rust/crates/token-counter/src/fast.rs | 2 +- .../crates/token-counter/src/tiktoken.rs | 2 +- litellm-rust/crates/types/AGENTS.md | 60 ++++ .../crates/types/src/audio_transcription.rs | 15 + litellm-rust/crates/types/src/lib.rs | 2 + .../anthropic_messages/anthropic_request.rs | 53 ++- .../crates/types/src/messages/AGENTS.md | 5 + litellm-rust/crates/types/src/messages/mod.rs | 1 + .../src/messages/streaming.rs} | 30 +- .../src/responses/streaming_websocket.rs | 79 ++--- .../crates/types/tests/anthropic_request.rs | 49 +++ .../crates/types/tests/messages_streaming.rs | 37 ++ 136 files changed, 3199 insertions(+), 1615 deletions(-) create mode 100644 litellm-rust/crates/auth-types/tests/error.rs create mode 100644 litellm-rust/crates/core-utils/tests/settings.rs create mode 100644 litellm-rust/crates/core/src/messages/AGENTS.md create mode 100644 litellm-rust/crates/http/src/websocket.rs create mode 100644 litellm-rust/crates/http/tests/request.rs create mode 100644 litellm-rust/crates/http/tests/websocket.rs create mode 100644 litellm-rust/crates/llms/src/anthropic/AGENTS.md delete mode 100644 litellm-rust/crates/llms/src/azure_ai/anthropic/mod.rs create mode 100644 litellm-rust/crates/llms/src/azure_ai/common_utils.rs create mode 100644 litellm-rust/crates/llms/src/azure_ai/messages/AGENTS.md rename litellm-rust/crates/llms/src/{base_llm/anthropic_messages => azure_ai/messages}/mod.rs (55%) rename litellm-rust/crates/llms/src/azure_ai/{anthropic/messages_transformation.rs => messages/transformation.rs} (80%) create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/AGENTS.md create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/context.rs create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/mod.rs create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/normalization.rs rename litellm-rust/crates/llms/src/base_llm/{anthropic_messages => messages}/streaming.rs (67%) rename litellm-rust/crates/llms/src/base_llm/{anthropic_messages => messages}/transformation.rs (75%) create mode 100644 litellm-rust/crates/llms/src/bedrock/messages/AGENTS.md create mode 100644 litellm-rust/crates/llms/tests/messages_normalization.rs create mode 100644 litellm-rust/crates/types/AGENTS.md create mode 100644 litellm-rust/crates/types/src/audio_transcription.rs create mode 100644 litellm-rust/crates/types/src/messages/AGENTS.md create mode 100644 litellm-rust/crates/types/src/messages/mod.rs rename litellm-rust/crates/{llms/src/anthropic/messages/streaming_iterator.rs => types/src/messages/streaming.rs} (87%) create mode 100644 litellm-rust/crates/types/tests/anthropic_request.rs create mode 100644 litellm-rust/crates/types/tests/messages_streaming.rs diff --git a/.github/scripts/verify_linux_native_wheel.py b/.github/scripts/verify_linux_native_wheel.py index 6b7fcd57bbc..465918f5a81 100644 --- a/.github/scripts/verify_linux_native_wheel.py +++ b/.github/scripts/verify_linux_native_wheel.py @@ -214,7 +214,7 @@ def main( native_module: Final = load_native_module(native_path) native_module_loads: Final = native_module is not None panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test") - native_size_limit: Final = 40_000_000 + native_size_limit: Final = 45_000_000 native_size_within_limit: Final = native_member.file_size <= native_size_limit validations: Final = ( (f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG), diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index bc6a2552e4c..b1dc35d3698 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -16,7 +16,9 @@ Use [`#[rstest]`](https://docs.rs/rstest/latest/rstest/attr.rstest.html) for new ## Error definitions - A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` +- Put message templates in the variant's `#[error(...)]` declaration. Callers pass only the small typed arguments needed to fill them, never `Error::Variant(format!(...))` or a preformatted message. Keep the smallest set of neutral variants that callers need to distinguish; different wording or providers do not justify new variants - Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each. A failure mode is something a caller handles differently (phase, status code, retry, a message Python parity pins exactly); failures no caller tells apart share one variant and differ only in its message +- Keep shared error enums minimal and provider-neutral. Provider names, credential types, configuration fields, and setup guidance belong in caller-supplied data, not dedicated variants or hardcoded shared messages. Reuse a variant for the same failure mode across providers, such as `MissingApiBase { provider: "Azure", guidance: "..." }`. An exact parity message does not justify a provider-specific variant when caller-supplied context can preserve it - Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string - Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return - Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index d67623feffd..bee4421f1e9 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2891,6 +2891,7 @@ dependencies = [ "http 1.4.2", "litellm-auth-types", "moka", + "rstest", "serde_json", "sha2 0.10.9", "tokio", @@ -2900,6 +2901,7 @@ dependencies = [ name = "litellm-auth-types" version = "0.1.0" dependencies = [ + "rstest", "serde", "subtle", "thiserror 2.0.19", @@ -3146,8 +3148,6 @@ dependencies = [ "reqwest 0.12.28", "rstest", "rstest_reuse", - "rustls 0.23.42", - "rustls-native-certs", "serde", "serde_json", "sha2 0.10.9", @@ -3194,6 +3194,7 @@ version = "0.1.0" dependencies = [ "criterion", "proptest", + "rstest", ] [[package]] @@ -3305,6 +3306,7 @@ dependencies = [ name = "litellm-http" version = "0.1.0" dependencies = [ + "futures-util", "http 1.4.2", "hyper-util", "litellm-core-utils", @@ -3312,11 +3314,13 @@ dependencies = [ "reqwest 0.12.28", "rstest", "rustls 0.23.42", + "rustls-native-certs", "serde", "serde_json", "tempfile", "thiserror 2.0.19", "tokio", + "tokio-tungstenite", "veil", "webpki-roots", ] @@ -3661,6 +3665,7 @@ dependencies = [ name = "litellm-token-counter-huggingface" version = "0.1.0" dependencies = [ + "rstest", "serde_json", "thiserror 2.0.19", "tokenizers", @@ -3672,6 +3677,7 @@ version = "0.1.0" dependencies = [ "base64 0.22.1", "once_cell", + "rstest", "rustc-hash", "thiserror 2.0.19", "tiktoken-rs", diff --git a/litellm-rust/crates/auth-azure/src/native.rs b/litellm-rust/crates/auth-azure/src/native.rs index d635e559641..64752162384 100644 --- a/litellm-rust/crates/auth-azure/src/native.rs +++ b/litellm-rust/crates/auth-azure/src/native.rs @@ -133,7 +133,7 @@ impl NativeAzureTokenAcquirer { let token = credential .get_token(&[scope.as_str()], None) .await - .map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?; + .map_err(|error| Error::CredentialAcquisition(error.to_string().into()))?; let expires_on = u64::try_from(token.expires_on.unix_timestamp()) .ok() .map(|seconds| UNIX_EPOCH + Duration::from_secs(seconds)); @@ -250,7 +250,12 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> { let Some(authority) = authority else { return Ok(()); }; - let url = url::Url::parse(authority.value()).map_err(|_| Error::InvalidAzureAuthority)?; + let url = url::Url::parse(authority.value()).map_err(|_| { + Error::InvalidConfiguration( + "Azure authority must be an HTTPS origin without credentials, query, or fragment" + .into(), + ) + })?; if url.scheme() != "https" || url.host_str().is_none() || !url.username().is_empty() @@ -259,7 +264,10 @@ fn validate_authority(request: &NativeAzureRequest) -> Result<(), Error> { || url.fragment().is_some() || !matches!(url.path(), "" | "/") { - return Err(Error::InvalidAzureAuthority); + return Err(Error::InvalidConfiguration( + "Azure authority must be an HTTPS origin without credentials, query, or fragment" + .into(), + )); } Ok(()) } @@ -368,7 +376,9 @@ fn trusted_source(sources: &[InputSource]) -> InputSource { } fn mixed_sources() -> Result { - Err(Error::MixedAzureCredentialSources) + Err(Error::InvalidConfiguration( + "request-controlled Azure auth inputs cannot be combined with host credentials".into(), + )) } fn build_credential( @@ -433,7 +443,12 @@ fn build_credential( NativeAzureRequest::DeveloperTools { .. } => DeveloperToolsCredential::new(None) .map(|credential| credential as Arc), } - .map_err(|error| Error::AzureCredentialInitialization(error.to_string())) + .map_err(|error| { + Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed( + "Azure credential initialization", + error, + )) + }) } fn client_options( @@ -638,7 +653,7 @@ mod tests { assert_eq!(transport.requests.lock().unwrap().len(), 6); } - #[test] + #[rstest::rstest] fn request_authority_requires_request_owned_client_secret_identity() { let error = ValidatedAzureRequest::new(sourced_client_secret( InputSource::Deployment, @@ -647,10 +662,13 @@ mod tests { )) .unwrap_err(); - assert!(matches!( + assert_eq!( error, - litellm_auth_types::Error::MixedAzureCredentialSources - )); + litellm_auth_types::Error::InvalidConfiguration( + "request-controlled Azure auth inputs cannot be combined with host credentials" + .into() + ) + ); } #[test] @@ -665,24 +683,24 @@ mod tests { assert_eq!(request.credential_source(), InputSource::Request); } - #[test] - fn authority_is_restricted_to_an_https_origin() { - for authority in [ - "http://login.example", - "https://user@login.example", - "https://login.example/tenant", - "https://login.example?target=other", - ] { - let error = ValidatedAzureRequest::new(sourced_client_secret( - InputSource::Deployment, - InputSource::Deployment, - authority, - )) - .unwrap_err(); - assert!(matches!( - error, - litellm_auth_types::Error::InvalidAzureAuthority - )); - } + #[rstest::rstest] + #[case::http("http://login.example")] + #[case::userinfo("https://user@login.example")] + #[case::path("https://login.example/tenant")] + #[case::query("https://login.example?target=other")] + fn authority_is_restricted_to_an_https_origin(#[case] authority: &str) { + let error = ValidatedAzureRequest::new(sourced_client_secret( + InputSource::Deployment, + InputSource::Deployment, + authority, + )) + .unwrap_err(); + assert_eq!( + error, + litellm_auth_types::Error::InvalidConfiguration( + "Azure authority must be an HTTPS origin without credentials, query, or fragment" + .into() + ) + ); } } diff --git a/litellm-rust/crates/auth-azure/src/resolve.rs b/litellm-rust/crates/auth-azure/src/resolve.rs index 9a7afe645db..2142e22db50 100644 --- a/litellm-rust/crates/auth-azure/src/resolve.rs +++ b/litellm-rust/crates/auth-azure/src/resolve.rs @@ -91,7 +91,9 @@ impl AzureAuthService { AzureCredentialPlan::Caller(caller) => { let credential = caller.acquire().await?; if credential.secret().expose().is_empty() { - return Err(Error::EmptyAzureToken); + return Err(Error::EmptyCallerCredential( + "Azure AD token provider returned an empty token", + )); } Ok(Some(Sourced::new(credential, InputSource::Deployment))) } @@ -104,7 +106,11 @@ impl AzureAuthService { } => { let assertion = resolve_reference(inputs, env_lookup, reference.value()) .await? - .ok_or(Error::UnresolvedOidcReference)?; + .ok_or_else(|| { + Error::CredentialAcquisition( + "Azure OIDC reference did not resolve to a value".into(), + ) + })?; let request = ValidatedAzureRequest::new(NativeAzureRequest::ClientAssertion { tenant_id, client_id, @@ -167,7 +173,7 @@ pub(crate) fn select_auth_plan( .map(|selector| Sourced::new(selector, value.source())) }) .transpose() - .map_err(|_| Error::InvalidAzureSelector)?; + .map_err(|_| Error::InvalidConfiguration("invalid Azure credential selector".into()))?; let federated_token_file = configured_string( &inputs.federated_token_file, AZURE_FEDERATED_TOKEN_FILE_ENV, @@ -257,7 +263,9 @@ fn select_native_plan( let selection_source = selected.source(); match selected.into_value() { - AzureCredentialType::ClientSecretCredential => Err(Error::MissingClientSecretFields), + AzureCredentialType::ClientSecretCredential => Err(Error::InvalidConfiguration( + "ClientSecretCredential requires tenant_id, client_id, and client_secret".into(), + )), AzureCredentialType::WorkloadIdentityCredential => { Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new( workload_request(tenant_id, client_id, federated_token_file, scope, authority)?, @@ -341,9 +349,17 @@ fn workload_request( authority: Option>, ) -> Result { Ok(NativeAzureRequest::WorkloadIdentity { - tenant_id: tenant_id.ok_or(Error::MissingWorkloadTenant)?, - client_id: client_id.ok_or(Error::MissingWorkloadClient)?, - token_file_path: token_file_path.ok_or(Error::MissingWorkloadTokenFile)?, + tenant_id: tenant_id.ok_or_else(|| { + Error::InvalidConfiguration("WorkloadIdentityCredential requires tenant_id".into()) + })?, + client_id: client_id.ok_or_else(|| { + Error::InvalidConfiguration("WorkloadIdentityCredential requires client_id".into()) + })?, + token_file_path: token_file_path.ok_or_else(|| { + Error::InvalidConfiguration( + "WorkloadIdentityCredential requires azure_federated_token_file".into(), + ) + })?, scope, authority, }) @@ -394,10 +410,11 @@ async fn resolve_reference( .map_or(CredentialLookup::Missing, CredentialLookup::Found), CredentialRef::None => return Ok(None), CredentialRef::File(_) | CredentialRef::Request(_) | CredentialRef::Host(_) => { - let resolver = inputs - .credential_resolver - .as_ref() - .ok_or(Error::MissingHostResolver)?; + let resolver = inputs.credential_resolver.as_ref().ok_or_else(|| { + Error::InvalidConfiguration( + "credential reference requires a host credential resolver".into(), + ) + })?; resolver.resolve(reference).await? } }; @@ -415,7 +432,9 @@ fn oidc_reference( }; let value = token.value().expose(); if token.source() == InputSource::Request && value.starts_with("oidc/") { - return Err(Error::RequestAzureCredentialReference); + return Err(Error::InvalidConfiguration( + "request-controlled Azure credential references are not allowed".into(), + )); } if let Some(name) = value.strip_prefix("oidc/env/") { return non_empty_reference(name, "OIDC environment reference") @@ -437,14 +456,20 @@ fn oidc_reference( ))); } if value.starts_with("oidc/") { - return Err(Error::UnsupportedOidcReference); + return Err(Error::InvalidConfiguration( + "unsupported OIDC reference".into(), + )); } Ok(None) } fn non_empty_reference(value: &str, kind: &str) -> Result { if value.is_empty() { - return Err(Error::EmptyReference(kind.to_string())); + return Err(Error::InvalidConfiguration( + litellm_auth_types::ErrorDetail::Empty { + subject: kind.into(), + }, + )); } Ok(value.to_string()) } @@ -493,7 +518,7 @@ mod tests { expires_on: None, }) } else { - Err(Error::AzureTokenAcquisition(format!("{kind} failed"))) + Err(Error::CredentialAcquisition(kind.into())) } }) } @@ -602,7 +627,7 @@ mod tests { assert!(error.to_string().contains("unsupported OIDC reference")); } - #[test] + #[rstest::rstest] fn request_oidc_reference_is_rejected_before_lookup() { let params = json!({ "azure_ad_token": "oidc/env/ASSERTION", @@ -624,7 +649,12 @@ mod tests { }) .unwrap_err(); - assert!(matches!(error, Error::RequestAzureCredentialReference)); + assert_eq!( + error, + Error::InvalidConfiguration( + "request-controlled Azure credential references are not allowed".into() + ) + ); } #[tokio::test] @@ -723,6 +753,7 @@ mod tests { assert_eq!(credential.value().secret().expose(), "caller-token"); } + #[rstest::rstest] #[tokio::test] async fn empty_caller_token_is_rejected() { let error = AzureAuthService::default() @@ -730,6 +761,9 @@ mod tests { .await .unwrap_err(); - assert!(matches!(error, Error::EmptyAzureToken)); + assert_eq!( + error, + Error::EmptyCallerCredential("Azure AD token provider returned an empty token") + ); } } diff --git a/litellm-rust/crates/auth-azure/src/types.rs b/litellm-rust/crates/auth-azure/src/types.rs index a3a898f000f..a042937a047 100644 --- a/litellm-rust/crates/auth-azure/src/types.rs +++ b/litellm-rust/crates/auth-azure/src/types.rs @@ -117,7 +117,12 @@ fn string_config( None => Ok(ConfigValue::Absent), Some(Value::Null) => Ok(ConfigValue::ExplicitNone(source)), Some(Value::String(value)) => Ok(ConfigValue::Value(Sourced::new(value.clone(), source))), - Some(_) => Err(Error::InvalidFieldType(name.to_string())), + Some(_) => Err(Error::InvalidConfiguration( + litellm_auth_types::ErrorDetail::InvalidType { + field: name.into(), + expected: "a string or null", + }, + )), } } diff --git a/litellm-rust/crates/auth-gcp/Cargo.toml b/litellm-rust/crates/auth-gcp/Cargo.toml index 0c6258a193c..8a3598234e1 100644 --- a/litellm-rust/crates/auth-gcp/Cargo.toml +++ b/litellm-rust/crates/auth-gcp/Cargo.toml @@ -19,3 +19,6 @@ tokio.workspace = true gcp_auth = "0.12.7" google-cloud-auth = { workspace = true, optional = true } http = { workspace = true, optional = true } + +[dev-dependencies] +rstest.workspace = true diff --git a/litellm-rust/crates/auth-gcp/src/lib.rs b/litellm-rust/crates/auth-gcp/src/lib.rs index 4374dff95aa..97bc2c482c3 100644 --- a/litellm-rust/crates/auth-gcp/src/lib.rs +++ b/litellm-rust/crates/auth-gcp/src/lib.rs @@ -299,7 +299,7 @@ fn validate_request_credentials(configured: &str) -> Result<&str, Error> { .map(str::to_string) }); if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) { - return Err(Error::RequestVertexTokenEndpoint); + return Err(Error::InvalidConfiguration("request-controlled Vertex credentials must use the canonical Google OAuth token endpoint".into())); } Ok(configured) } @@ -376,10 +376,20 @@ fn optional_credentials( .map(SecretValue::new) .map(|value| Sourced::new(value, source)) .map(Some) - .map_err(|error| Error::InvalidFieldType(format!("{}: {error}", names[0]))); + .map_err(|error| { + Error::InvalidConfiguration(litellm_auth_types::ErrorDetail::failed( + "credential serialization", + error, + )) + }); } Some(_) => { - return Err(Error::InvalidFieldType(names[0].to_string())); + return Err(Error::InvalidConfiguration( + litellm_auth_types::ErrorDetail::InvalidType { + field: names[0].into(), + expected: "a string or null", + }, + )); } } } @@ -397,7 +407,12 @@ fn optional_string(params: &Map, names: &[&str]) -> Result