diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index d56e29fb627..3f4f5620176 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -5,6 +5,7 @@ flag="${1:?usage: unit_selection.sh }" legacy_flags=( caching-local + core-utils enterprise-package enterprise-routing integrations @@ -32,6 +33,7 @@ legacy_flags=( legacy_paths() { case "$1" in caching-local) echo tests/unit/caching ;; + core-utils) echo tests/unit/litellm_core_utils ;; enterprise-package) echo tests/unit/enterprise/integrations echo tests/unit/enterprise/proxy/auth @@ -42,6 +44,8 @@ legacy_paths() { echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;; enterprise-routing) echo tests/unit/google_genai + echo tests/unit/router_strategy + echo tests/unit/router_utils 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 @@ -77,6 +81,7 @@ legacy_paths() { echo tests/unit/messages echo tests/unit/rag echo tests/unit/rerank_api + echo tests/unit/rust_bridge echo tests/unit/secret_managers echo tests/unit/vector_stores echo tests/unit/videos ;; @@ -142,7 +147,9 @@ legacy_paths() { proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;; proxy-extras) echo tests/unit/litellm_proxy_extras ;; proxy-infra) echo tests/unit/gateway ;; - responses-caching-types) echo tests/unit/types ;; + responses-caching-types) + find tests/unit/responses -name 'test_*.py' -not -path 'tests/unit/responses/mcp/*' + echo tests/unit/types ;; *) echo "unit_selection.sh: unknown flag $1" >&2; exit 1 ;; esac } diff --git a/.circleci/tests.yml b/.circleci/tests.yml index 41e9f11cefa..a9cd21bad5e 100644 --- a/.circleci/tests.yml +++ b/.circleci/tests.yml @@ -369,6 +369,13 @@ workflows: reruns: 2 base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> + - unit: + name: unit-core-utils + flag: core-utils + shards: 2 + reruns: 1 + base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >> + pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >> - unit: name: unit-integrations flag: integrations diff --git a/.github/merge-smoke-tests.json b/.github/merge-smoke-tests.json index a563424c230..727733fa954 100644 --- a/.github/merge-smoke-tests.json +++ b/.github/merge-smoke-tests.json @@ -7,9 +7,9 @@ "MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]", "COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones", "COST-ZERO": "tests/unit/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero", - "LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on", - "LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off", - "CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger", - "CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger" + "LOG-CONTENT-ON": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on", + "LOG-CONTENT-OFF": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off", + "CALLBACK-SUCCESS": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger", + "CALLBACK-FAILURE": "tests/unit/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger" } } diff --git a/.github/workflows/test-redis-compat.yml b/.github/workflows/test-redis-compat.yml index 2f5ce4d441a..0423b014ec5 100644 --- a/.github/workflows/test-redis-compat.yml +++ b/.github/workflows/test-redis-compat.yml @@ -12,9 +12,9 @@ on: - "litellm/caching/evicted_client_closer.py" - "tests/unit/test_redis.py" - "tests/local_testing/test_caching.py" - - "tests/test_litellm/caching/test_redis_connection_pool.py" - - "tests/test_litellm/caching/test_redis_cluster_cache.py" - - "tests/test_litellm/caching/test_evicted_client_closer.py" + - "tests/unit/caching/test_redis_connection_pool.py" + - "tests/unit/caching/test_redis_cluster_cache.py" + - "tests/unit/caching/test_evicted_client_closer.py" - ".github/workflows/test-redis-compat.yml" - "pyproject.toml" - "uv.lock" @@ -85,9 +85,9 @@ jobs: redis-server --version uv run --no-sync pytest \ tests/unit/test_redis.py \ - tests/test_litellm/caching/test_redis_connection_pool.py \ - tests/test_litellm/caching/test_redis_cluster_cache.py \ - tests/test_litellm/caching/test_evicted_client_closer.py \ + tests/unit/caching/test_redis_connection_pool.py \ + tests/unit/caching/test_redis_cluster_cache.py \ + tests/unit/caching/test_evicted_client_closer.py \ tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \ tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \ --tb=short -vv \ diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 1f3b5c4d97c..808bb2afd08 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -24,7 +24,7 @@ on: - ".github/actions/setup-uv-with-retries/**" - ".github/scripts/smoke_test_native_wheel.py" - ".github/scripts/verify_linux_native_wheel.py" - - "tests/test_litellm/rust_bridge/native_route_wheel_test.py" + - "tests/unit/rust_bridge/native_route_wheel_test.py" - ".github/workflows/test-rust.yml" pull_request: branches: @@ -52,7 +52,7 @@ on: - ".github/actions/setup-uv-with-retries/**" - ".github/scripts/smoke_test_native_wheel.py" - ".github/scripts/verify_linux_native_wheel.py" - - "tests/test_litellm/rust_bridge/native_route_wheel_test.py" + - "tests/unit/rust_bridge/native_route_wheel_test.py" - ".github/workflows/test-rust.yml" permissions: @@ -171,7 +171,7 @@ jobs: env: RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }} - - run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl + - run: python tests/unit/rust_bridge/native_route_wheel_test.py dist/*.whl - name: Run pytest tests/test_litellm_rust with the compiled extension run: make test-rust-extension diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 4dca8075440..d75213d37ea 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -62,6 +62,7 @@ jobs: - shard: core-utils artifact-name: core-utils test-path: "tests/test_litellm/litellm_core_utils" + unit-flag: core-utils workers: 2 reruns: 1 timeout-minutes: 20 @@ -69,9 +70,7 @@ jobs: - shard: enterprise-routing artifact-name: enterprise-routing - test-path: >- - tests/test_litellm/router_utils - tests/test_litellm/router_strategy + test-path: "" unit-flag: enterprise-routing workers: 2 reruns: 2 @@ -111,7 +110,6 @@ jobs: tests/test_litellm/interactions tests/test_litellm/ocr tests/test_litellm/passthrough - tests/test_litellm/rust_bridge tests/test_litellm/test_*.py unit-flag: misc workers: 2 @@ -228,9 +226,7 @@ jobs: - shard: responses-caching-types artifact-name: responses-caching-types - test-path: >- - tests/test_litellm/responses - tests/test_litellm/caching + test-path: "" unit-flag: responses-caching-types workers: 2 reruns: 2 diff --git a/Makefile b/Makefile index f27525b58ff..311a7daef92 100644 --- a/Makefile +++ b/Makefile @@ -301,7 +301,7 @@ test-rust-extension: UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \ $(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \ "$$temporary/venv/bin/python" -I -m mypy.stubtest \ - --mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \ + --mypy-config-file tests/unit/rust_bridge/stubtest.ini \ litellm.rust_bridge._native && \ LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \ "$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust @@ -329,10 +329,10 @@ test-unit-integrations: install-test-deps $(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20 test-unit-core-utils: install-test-deps - $(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20 + $(UV_RUN) pytest tests/unit/litellm_core_utils --tb=short -vv -n 2 --durations=20 test-unit-other: install-test-deps - $(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20 + $(UV_RUN) pytest tests/unit/caching tests/unit/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/unit/router_strategy tests/unit/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20 test-unit-root: install-test-deps $(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20 diff --git a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs index 030bf03d4ba..69f72fbc177 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs @@ -33,7 +33,7 @@ mod test_support { use crate::{LegacyLogging, LegacySurface, PublicCall}; /// The parameters of every `callbacks_legacy_python` function, as the real module declares them. - /// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python + /// `tests/unit/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python /// signatures, and [`namespace`] binds every fake call against it. pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json"); diff --git a/litellm-rust/crates/secrets/README.md b/litellm-rust/crates/secrets/README.md index c8b01fe9b3a..10619613516 100644 --- a/litellm-rust/crates/secrets/README.md +++ b/litellm-rust/crates/secrets/README.md @@ -30,7 +30,7 @@ The HashiCorp Vault backend is enabled with the `hashicorp` feature and reads KV Native backends consistently distinguish absence from failure instead of swallowing provider errors. Python-compatible resolution maps these results back to the Python handler contract before applying fallback -`hosted_keys` excludes a name for every backend. Python's handler recognizes Azure `SecretClient` and Google `KeyManagementServiceClient` instances before the `local` branch, allowing excluded names to reach those providers. Rust treats that as a routing bug. `test_rust_hosted_keys_exclude_azure_sdk_clients_too` in `tests/test_litellm/rust_bridge/ocr/test_secrets.py` pins this behavior +`hosted_keys` excludes a name for every backend. Python's handler recognizes Azure `SecretClient` and Google `KeyManagementServiceClient` instances before the `local` branch, allowing excluded names to reach those providers. Rust treats that as a routing bug. `test_rust_hosted_keys_exclude_azure_sdk_clients_too` in `tests/unit/rust_bridge/ocr/test_secrets.py` pins this behavior Google rejects malformed base64 and mismatched CRC32C values instead of accepting corrupted payloads. Python currently ignores the checksum and uses permissive base64 decoding. Rust follows [RFC 4648](https://www.rfc-editor.org/rfc/rfc4648#section-3.3) and [Google's integrity guidance](https://docs.cloud.google.com/secret-manager/docs/data-integrity); `failed_or_missing_reads_are_not_cached` covers rejection and recovery diff --git a/litellm/__init__.py b/litellm/__init__.py index e334fbe8ca8..676c735b9e8 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -50,6 +50,7 @@ from litellm.types.integrations.datadog import DatadogInitParams from litellm.types.integrations.newrelic import NewRelicInitParams from litellm.litellm_core_utils.core_helpers import drop_params_env_flag from litellm.types.integrations.pointfive import PointFiveInitParams +from litellm.types.integrations.zerobus import ZerobusInitParams from litellm._logging import ( set_verbose, _turn_on_debug, @@ -157,6 +158,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "deepeval", "s3_v2", "pointfive", + "zerobus", "aws_sqs", "vector_store_pre_call_hook", "dotprompt", @@ -442,6 +444,7 @@ datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]] datadog_params: Optional[Union[DatadogInitParams, Dict]] = None newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None pointfive_params: Optional[Union[PointFiveInitParams, Mapping[str, object]]] = None +zerobus_params: Optional[Union[ZerobusInitParams, Mapping[str, object]]] = None aws_sqs_callback_params: Optional[Dict] = None generic_logger_headers: Optional[Dict] = None default_key_generate_params: Optional[Dict] = None diff --git a/litellm/anthropic_interface/exceptions/__init__.py b/litellm/anthropic_interface/exceptions/__init__.py index 7f2de0e60dc..7c3cea0a28a 100644 --- a/litellm/anthropic_interface/exceptions/__init__.py +++ b/litellm/anthropic_interface/exceptions/__init__.py @@ -2,7 +2,9 @@ from .exception_mapping_utils import ( ANTHROPIC_ERROR_TYPE_MAP, + AnthropicErrorSseFrame, AnthropicExceptionMapping, + anthropic_error_sse_frame, ) from .exceptions import ( AnthropicErrorDetail, @@ -14,6 +16,8 @@ __all__ = [ "ANTHROPIC_ERROR_TYPE_MAP", "AnthropicErrorDetail", "AnthropicErrorResponse", + "AnthropicErrorSseFrame", "AnthropicErrorType", "AnthropicExceptionMapping", + "anthropic_error_sse_frame", ] diff --git a/litellm/anthropic_interface/exceptions/exception_mapping_utils.py b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py index d9c9925275b..eb3ec8aaee2 100644 --- a/litellm/anthropic_interface/exceptions/exception_mapping_utils.py +++ b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py @@ -4,11 +4,12 @@ Utilities for mapping exceptions to Anthropic error format. Similar to litellm/litellm_core_utils/exception_mapping_utils.py but for Anthropic response format. """ +import json from typing import Final from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from .exceptions import AnthropicErrorResponse, AnthropicErrorType +from .exceptions import AnthropicErrorDetail, AnthropicErrorResponse, AnthropicErrorType # HTTP status code -> Anthropic error type # Source: https://docs.anthropic.com/en/api/errors @@ -166,3 +167,36 @@ class AnthropicExceptionMapping: message=message, request_id=request_id, ) + + +class AnthropicErrorSseFrame(str): + """One `event: error` frame, for a stream that fails once the response headers are out. + + Anthropic clients pick stream events by the `event:` name, so a frame carrying only a `data:` + line is skipped and the failure never reaches the caller. The frame remembers the status and + body it was built from, so a stream that fails before its first byte can still answer as a + JSON error with that exact status instead of a 200 that only says `api_error` + """ + + status_code: int + error_response: AnthropicErrorResponse + + def __new__(cls, status_code: int, error_response: AnthropicErrorResponse) -> "AnthropicErrorSseFrame": + frame: Final = super().__new__(cls, f"event: error\ndata: {json.dumps(error_response)}\n\n") + frame.status_code = status_code + frame.error_response = error_response + return frame + + def json_body(self, call_id: str | None) -> AnthropicErrorResponse: + if call_id is None: + return self.error_response + detail: Final[AnthropicErrorDetail] = {**self.error_response["error"], "litellm_call_id": call_id} + body: Final[AnthropicErrorResponse] = {**self.error_response, "error": detail} + return body + + +def anthropic_error_sse_frame(status_code: int, raw_message: str) -> AnthropicErrorSseFrame: + return AnthropicErrorSseFrame( + status_code, + AnthropicExceptionMapping.transform_to_anthropic_error(status_code=status_code, raw_message=raw_message), + ) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 819a279a43c..246ac4fd369 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -706,6 +706,10 @@ def _get_batch_job_usage_from_response_body( if ResponseAPILoggingUtils._is_response_api_usage(_usage_dict): return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_usage_dict) usage: Final[Usage] = Usage(**_usage_dict) + if custom_llm_provider == "xai": + from litellm.llms.xai.chat.transformation import XAIChatConfig + + XAIChatConfig.fold_reasoning_tokens_into_completion(usage) return usage diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 76b6c73b375..f977fc03891 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -31,6 +31,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.llms.openai.openai import OpenAIBatchesAPI from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction +from litellm.llms.xai.batches.handler import XAIBatchesHandler from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( CancelBatchRequest, @@ -59,6 +60,7 @@ openai_batches_instance: Final = OpenAIBatchesAPI() azure_batches_instance: Final = AzureBatchesAPI() vertex_ai_batches_instance: Final = VertexAIBatchPrediction(gcs_bucket_name="") anthropic_batches_instance: Final = AnthropicBatchesHandler() +xai_batches_instance: Final = XAIBatchesHandler() base_llm_http_handler = BaseLLMHTTPHandler() ################################################# @@ -105,10 +107,22 @@ def _resolve_timeout( @client async def acreate_batch( completion_window: Literal["24h"], - endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"], + endpoint: Literal[ + "/v1/chat/completions", + "/v1/embeddings", + "/v1/completions", + "/v1/responses", + "/v1/ocr", + "/v1/images/generations", + "/v1/images/edits", + "/v1/videos/generations", + "/v1/videos", + "/v1/videos/edits", + "/v1/videos/extensions", + ], input_file_id: str, custom_llm_provider: Literal[ - "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral" + "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "xai" ] = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, @@ -157,10 +171,22 @@ async def acreate_batch( @client def create_batch( completion_window: Literal["24h"], - endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"], + endpoint: Literal[ + "/v1/chat/completions", + "/v1/embeddings", + "/v1/completions", + "/v1/responses", + "/v1/ocr", + "/v1/images/generations", + "/v1/images/edits", + "/v1/videos/generations", + "/v1/videos", + "/v1/videos/edits", + "/v1/videos/extensions", + ], input_file_id: str, custom_llm_provider: Literal[ - "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral" + "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "xai" ] = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, @@ -243,6 +269,14 @@ def create_batch( model=model, ) return response + if custom_llm_provider == LlmProviders.XAI.value: + return xai_batches_instance.create_batch( + _is_async=_is_async, + create_batch_data=_create_batch_request, + api_base=optional_params.api_base, + api_key=optional_params.api_key, + timeout=timeout, + ) api_base: str | None = None if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -345,7 +379,7 @@ def create_batch( async def aretrieve_batch( batch_id: str, custom_llm_provider: Literal[ - "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral" + "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai" ] = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, @@ -393,10 +427,18 @@ def _handle_retrieve_batch_providers_without_provider_config( _retrieve_batch_request: RetrieveBatchRequest, _is_async: bool, custom_llm_provider: Literal[ - "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral" + "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai" ] = "openai", logging_obj: LiteLLMLoggingObj | None = None, ): + if custom_llm_provider == LlmProviders.XAI.value: + return xai_batches_instance.retrieve_batch( + _is_async=_is_async, + batch_id=batch_id, + api_base=optional_params.api_base, + api_key=optional_params.api_key, + timeout=timeout, + ) api_base: str | None = None if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there @@ -518,7 +560,7 @@ def _handle_retrieve_batch_providers_without_provider_config( def retrieve_batch( batch_id: str, custom_llm_provider: Literal[ - "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral" + "openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai" ] = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, @@ -741,6 +783,15 @@ def list_batches( timeout = 600.0 _is_async: Final = kwargs.pop("alist_batches", False) is True + if custom_llm_provider == LlmProviders.XAI.value: + return xai_batches_instance.list_batches( + _is_async=_is_async, + api_base=optional_params.api_base, + api_key=optional_params.api_key, + timeout=timeout, + after=after, + limit=limit, + ) if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: # for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there api_base = ( @@ -837,7 +888,7 @@ def list_batches( async def acancel_batch( batch_id: str, model: str | None = None, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] = "openai", + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, extra_body: dict[str, str] | None = None, @@ -883,7 +934,7 @@ async def acancel_batch( def cancel_batch( batch_id: str, model: str | None = None, - custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] | str = "openai", + custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] | str = "openai", metadata: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, extra_body: dict[str, str] | None = None, @@ -933,6 +984,14 @@ def cancel_batch( ) _is_async: Final = kwargs.pop("acancel_batch", False) is True + if custom_llm_provider == LlmProviders.XAI.value: + return xai_batches_instance.cancel_batch( + _is_async=_is_async, + batch_id=batch_id, + api_base=optional_params.api_base, + api_key=optional_params.api_key, + timeout=timeout, + ) api_base: str | None = None if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: api_base = ( diff --git a/litellm/constants.py b/litellm/constants.py index e7ba1f6b07f..8316761c95b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1952,6 +1952,15 @@ SENTRY_DENYLIST: Final = [ "auth_token", "jwt_token", "private_key", + "authorization", + "api-key", + "x-api-key", + "x-goog-api-key", + "ocp-apim-subscription-key", + "x-litellm-api-key", + "x-mcp-auth", + "cookie", + "set-cookie", "SLACK_WEBHOOK_URL", "ALERTING_WEBHOOK_URL", "webhook_url", @@ -1974,6 +1983,12 @@ SENTRY_DENYLIST: Final = [ ] SENTRY_PII_DENYLIST: Final = [ "user_id", + "user_email", + "end_user_id", + "user_api_key_hash", + "user_api_key_user_id", + "user_api_key_user_email", + "user_api_key_end_user_id", "email", "phone", "address", diff --git a/litellm/files/main.py b/litellm/files/main.py index 72832aeccc9..723784795b0 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -28,12 +28,15 @@ FileCreateProvider = Literal[ "manus", "anthropic", "mistral", + "xai", ] FileRetrieveProvider = Literal[ - "openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral" + "openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral", "xai" ] -FileDeleteProvider = Literal["openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral"] -FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic", "mistral"] +FileDeleteProvider = Literal[ + "openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral", "xai" +] +FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic", "mistral", "xai"] import litellm from litellm import get_secret_str from litellm.files.streaming import FileContentStreamingResponse @@ -49,6 +52,8 @@ from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.llms.openai.common_utils import get_openai_credentials from litellm.llms.openai.openai import FileDeleted, FileObject, OpenAIFilesAPI from litellm.llms.vertex_ai.files.handler import VertexAIFilesHandler +from litellm.llms.xai.batches.handler import XAIBatchesHandler +from litellm.llms.xai.batches.transformation import is_xai_batch_results_id from litellm.types.llms.openai import ( CreateFileRequest, FileContentRequest, @@ -103,6 +108,7 @@ openai_files_instance: Final = OpenAIFilesAPI() azure_files_instance: Final = AzureOpenAIFilesAPI() vertex_ai_files_instance: Final = VertexAIFilesHandler() bedrock_files_instance: Final = BedrockFilesHandler() +xai_batch_results_instance: Final = XAIBatchesHandler() ################################################# @@ -920,6 +926,15 @@ def file_content( client=client, ) + if custom_llm_provider == LlmProviders.XAI.value and is_xai_batch_results_id(file_id): + return xai_batch_results_instance.batch_results_content( + _is_async=_is_async, + batch_id=file_id, + api_base=optional_params.api_base, + api_key=optional_params.api_key, + timeout=timeout, + ) + # Check if provider has a custom files config (e.g., Anthropic, Manus) provider_config: Final = ProviderConfigManager.get_provider_files_config( model="", diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json index 5bd8aca55fa..4e72075dc5c 100644 --- a/litellm/integrations/callback_configs.json +++ b/litellm/integrations/callback_configs.json @@ -406,6 +406,45 @@ }, "description": "PointFive Logging Integration" }, + { + "id": "zerobus", + "displayName": "Databricks Zerobus", + "logo": "databricks.svg", + "supports_key_team_logging": false, + "dynamic_params": { + "ZEROBUS_WORKSPACE_URL": { + "type": "text", + "ui_name": "Workspace URL", + "description": "Databricks workspace URL, e.g. https://dbc-a1b2c3d4-e5f6.cloud.databricks.com", + "required": true + }, + "ZEROBUS_SERVER_ENDPOINT": { + "type": "text", + "ui_name": "Zerobus Endpoint", + "description": "Zerobus ingest endpoint, e.g. https://.zerobus..cloud.databricks.com", + "required": true + }, + "ZEROBUS_CLIENT_ID": { + "type": "text", + "ui_name": "Service Principal Client ID", + "description": "OAuth client id of a service principal with USE CATALOG, USE SCHEMA, SELECT and MODIFY on the table", + "required": true + }, + "ZEROBUS_CLIENT_SECRET": { + "type": "password", + "ui_name": "Service Principal Client Secret", + "description": "OAuth client secret of the service principal", + "required": true + }, + "ZEROBUS_TABLE_NAME": { + "type": "text", + "ui_name": "Table", + "description": "Fully qualified Unity Catalog table, catalog.schema.table, created with the LiteLLM trace schema", + "required": true + } + }, + "description": "Databricks Zerobus Ingest Logging Integration" + }, { "id": "s3", "displayName": "S3", diff --git a/litellm/integrations/zerobus/__init__.py b/litellm/integrations/zerobus/__init__.py new file mode 100644 index 00000000000..b1f5bc2ca40 --- /dev/null +++ b/litellm/integrations/zerobus/__init__.py @@ -0,0 +1,5 @@ +"""Databricks Zerobus logging integration for LiteLLM.""" + +from litellm.integrations.zerobus.logger import ZerobusLogger + +__all__ = ("ZerobusLogger",) diff --git a/litellm/integrations/zerobus/client.py b/litellm/integrations/zerobus/client.py new file mode 100644 index 00000000000..bf3a9e3e269 --- /dev/null +++ b/litellm/integrations/zerobus/client.py @@ -0,0 +1,161 @@ +""" +Writes rows to a Unity Catalog table through the Zerobus Ingest REST API. + +Zerobus only accepts a Databricks OAuth token minted for its own resource and scoped to +the target table's privileges, so the client mints that token itself with the service +principal's client credentials and reuses it until shortly before it expires. +""" + +import asyncio +import base64 +import json +import time +from collections.abc import Callable, Mapping, Sequence +from typing import Final + +import httpx +from pydantic import BaseModel, ValidationError + +import litellm +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.types.integrations.zerobus import ( + RETRYABLE_INGEST_STATUS_CODES, + TOKEN_REFRESH_LEEWAY_SECONDS, + ZerobusAccessToken, + ZerobusConnection, + ZerobusIngestFailure, +) + +TOKEN_PATH: Final = "/oidc/v1/token" +OAUTH_SCOPE: Final = "all-apis" + + +class _TokenResponse(BaseModel): + access_token: str + expires_in: float = 3600 + + +class ZerobusIngestError(Exception): + """A batch could not be written and the failure is worth retrying.""" + + +def zerobus_resource(workspace_id: str) -> str: + return f"api://databricks/workspaces/{workspace_id}/zerobusDirectWriteApi" + + +def authorization_details(table_name: str) -> str: + """The Unity Catalog privileges Zerobus requires the token to carry, as the token endpoint expects them.""" + catalog, schema, _table = table_name.split(".", 2) + return json.dumps( + ( + { + "type": "unity_catalog_privileges", + "privileges": ("USE CATALOG",), + "object_type": "CATALOG", + "object_full_path": catalog, + }, + { + "type": "unity_catalog_privileges", + "privileges": ("USE SCHEMA",), + "object_type": "SCHEMA", + "object_full_path": f"{catalog}.{schema}", + }, + { + "type": "unity_catalog_privileges", + "privileges": ("SELECT", "MODIFY"), + "object_type": "TABLE", + "object_full_path": table_name, + }, + ) + ) + + +def insert_url(connection: ZerobusConnection) -> str: + return f"{connection.server_endpoint.rstrip('/')}/zerobus/v1/tables/{connection.table_name}/insert" + + +def token_url(connection: ZerobusConnection) -> str: + return f"{connection.workspace_url.rstrip('/')}{TOKEN_PATH}" + + +def _basic_auth(client_id: str, client_secret: str) -> str: + return "Basic " + base64.b64encode(f"{client_id}:{client_secret}".encode()).decode() + + +def _status_failure(what: str, error: httpx.HTTPStatusError) -> ZerobusIngestFailure: + status: Final = error.response.status_code + return ZerobusIngestFailure( + detail=f"{what} returned {status}: {error.response.text}"[:500], + retryable=status in RETRYABLE_INGEST_STATUS_CODES, + ) + + +class ZerobusIngestClient: + def __init__( + self, + connection: ZerobusConnection, + http_client: AsyncHTTPHandler, + clock: Callable[[], float] = time.time, + ) -> None: + self.connection: Final = connection + self.http_client: Final = http_client + self.clock: Final = clock + self._token: ZerobusAccessToken | None = None + self._token_lock: Final = asyncio.Lock() + + async def insert(self, rows: Sequence[Mapping[str, object]]) -> ZerobusIngestFailure | None: + """Write ``rows`` as one request. ``None`` means Zerobus accepted every row.""" + token: Final = await self.access_token() + if isinstance(token, ZerobusIngestFailure): + return token + try: + await self.http_client.post( + insert_url(self.connection), + content=json.dumps([dict(row) for row in rows]).encode(), + headers={"Content-Type": "application/json", "Authorization": f"Bearer {token.value}"}, + ) + except httpx.HTTPStatusError as error: + if error.response.status_code == 401: + self._token = None + return ZerobusIngestFailure(detail="insert returned 401, token discarded", retryable=True) + return _status_failure("insert", error) + except (httpx.HTTPError, litellm.Timeout) as error: + return ZerobusIngestFailure(detail=f"insert failed: {error}", retryable=True) + return None + + async def access_token(self) -> ZerobusAccessToken | ZerobusIngestFailure: + """The cached token while it has more than the leeway left, otherwise a fresh one.""" + async with self._token_lock: + cached: Final = self._token + if cached is not None and cached.expires_at - self.clock() > TOKEN_REFRESH_LEEWAY_SECONDS: + return cached + minted: Final = await self._mint_token() + if isinstance(minted, ZerobusAccessToken): + self._token = minted + return minted + + async def _mint_token(self) -> ZerobusAccessToken | ZerobusIngestFailure: + connection: Final = self.connection + try: + response: Final = await self.http_client.post( + token_url(connection), + data={ + "grant_type": "client_credentials", + "scope": OAUTH_SCOPE, + "resource": zerobus_resource(connection.workspace_id), + "authorization_details": authorization_details(connection.table_name), + }, + headers={ + "Content-Type": "application/x-www-form-urlencoded", + "Authorization": _basic_auth(connection.client_id, connection.client_secret), + }, + ) + except httpx.HTTPStatusError as error: + return _status_failure("token request", error) + except (httpx.HTTPError, litellm.Timeout) as error: + return ZerobusIngestFailure(detail=f"token request failed: {error}", retryable=True) + try: + parsed: Final = _TokenResponse.model_validate_json(response.text) + except ValidationError as error: + return ZerobusIngestFailure(detail=f"token response was not understood: {error}", retryable=False) + return ZerobusAccessToken(value=parsed.access_token, expires_at=self.clock() + parsed.expires_in) diff --git a/litellm/integrations/zerobus/logger.py b/litellm/integrations/zerobus/logger.py new file mode 100644 index 00000000000..e2007218c8e --- /dev/null +++ b/litellm/integrations/zerobus/logger.py @@ -0,0 +1,230 @@ +"""Databricks Zerobus logging integration.""" + +import asyncio +from collections.abc import Mapping +from datetime import datetime +from typing import Final +from urllib.parse import urlsplit + +import litellm +from litellm._logging import verbose_logger +from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.integrations.zerobus.client import ZerobusIngestClient, ZerobusIngestError +from litellm.integrations.zerobus.row import trace_row +from litellm.litellm_core_utils.redact_messages import ( + redacted_standard_logging_payload, + should_redact_message_logging, +) +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, httpxSpecialProvider +from litellm.secret_managers.main import get_secret_str +from litellm.types.integrations.zerobus import ZerobusConnection, ZerobusInitParams + +_ENV_REFERENCE_PREFIX: Final = "os.environ/" + + +def _resolved_secret(value: str | None) -> str | None: + """Resolve a config value that may name a secret; an unset ``os.environ/NAME`` stays unresolved.""" + if value is None: + return None + resolved: Final = get_secret_str(value) + if resolved: + return resolved + return None if value.startswith(_ENV_REFERENCE_PREFIX) else value + + +def _configured_params() -> ZerobusInitParams: + configured: Final = litellm.zerobus_params + if isinstance(configured, ZerobusInitParams): + return configured + if isinstance(configured, Mapping): + return ZerobusInitParams.model_validate(configured) + return ZerobusInitParams() + + +def _setting(configured: str | None, env_var: str) -> str: + """Prefer the configured value, falling back to the environment the proxy UI writes.""" + value: Final = _resolved_secret(configured) or get_secret_str(env_var) + if not value: + raise ValueError( + f"zerobus logging requires {env_var}. Set it in the environment, or " + f"litellm_settings.zerobus_params.{env_var.removeprefix('ZEROBUS_').lower()} in config.yaml" + ) + return value + + +def _workspace_id(server_endpoint: str) -> str: + """The Zerobus endpoint is ``https://.zerobus..``, so the id is its first label.""" + host: Final = urlsplit(server_endpoint).hostname or "" + workspace_id: Final = host.split(".", 1)[0] + if not workspace_id.isdigit(): + raise ValueError( + f"ZEROBUS_SERVER_ENDPOINT {server_endpoint!r} does not look like " + "https://.zerobus..cloud.databricks.com" + ) + return workspace_id + + +def _table_name(configured: str | None) -> str: + table_name: Final = _setting(configured, "ZEROBUS_TABLE_NAME") + if table_name.count(".") != 2: + raise ValueError(f"ZEROBUS_TABLE_NAME {table_name!r} must be fully qualified as catalog.schema.table") + return table_name + + +def connection_for(params: ZerobusInitParams) -> ZerobusConnection: + """The connection configured right now, so a UI edit takes effect without a restart.""" + server_endpoint: Final = _setting(params.server_endpoint, "ZEROBUS_SERVER_ENDPOINT") + return ZerobusConnection( + workspace_url=_setting(params.workspace_url, "ZEROBUS_WORKSPACE_URL"), + workspace_id=_workspace_id(server_endpoint), + server_endpoint=server_endpoint, + client_id=_setting(params.client_id, "ZEROBUS_CLIENT_ID"), + client_secret=_setting(params.client_secret, "ZEROBUS_CLIENT_SECRET"), + table_name=_table_name(params.table_name), + ) + + +class ZerobusLogger(CustomBatchLogger): + preserve_events_added_during_flush = True + + def __init__( + self, + params: ZerobusInitParams | None = None, + client: ZerobusIngestClient | None = None, + start_periodic_flush: bool = True, + ) -> None: + resolved: Final = params if params is not None else _configured_params() + self.params: Final = resolved + self.given_client: Final = client + self._cached_client: ZerobusIngestClient | None = None + if client is None: + connection_for(resolved) + super().__init__( + flush_lock=asyncio.Lock(), + batch_size=resolved.batch_size, + flush_interval=resolved.flush_interval, + turn_off_message_logging=bool(resolved.turn_off_message_logging), + ) + self._flushing: bool = False + self._batch_flush_task: asyncio.Task[None] | None = None + self._periodic_flush_task: asyncio.Task[None] | None = ( + self._start_periodic_flush_task() if start_periodic_flush else None + ) + + @property + def client(self) -> ZerobusIngestClient: + """A client for the current connection, kept while the connection is unchanged so its token is reused.""" + if self.given_client is not None: + return self.given_client + connection: Final = connection_for(self.params) + cached: Final = self._cached_client + if cached is not None and cached.connection == connection: + return cached + fresh: Final = ZerobusIngestClient( + connection=connection, + http_client=get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback), + ) + self._cached_client = fresh + return fresh + + def _start_periodic_flush_task(self) -> asyncio.Task[None] | None: + try: + loop: Final = asyncio.get_running_loop() + except RuntimeError: + return None + return loop.create_task(self.periodic_flush()) + + def _start_batch_flush_task(self) -> None: + if self._batch_flush_task is not None and not self._batch_flush_task.done(): + return + try: + loop: Final = asyncio.get_running_loop() + except RuntimeError: + return + self._batch_flush_task = loop.create_task(self.flush_queue(skip_if_flushing=True)) + + def _flush_task_is_alive(self) -> bool: + task: Final = self._periodic_flush_task + return task is not None and not task.done() and not task.get_loop().is_closed() + + async def async_log_success_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: datetime, + end_time: datetime, + ) -> None: + await self._enqueue(kwargs) + + async def async_log_failure_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: datetime, + end_time: datetime, + ) -> None: + await self._enqueue(kwargs) + + async def _enqueue(self, kwargs: Mapping[str, object]) -> None: + try: + if not self._flush_task_is_alive(): + self._periodic_flush_task = self._start_periodic_flush_task() + + payload: Final = self._payload_for(kwargs) + if payload is None: + verbose_logger.debug("zerobus: event carried no standard_logging_object, skipping") + return + + if self._flushing and len(self.log_queue) >= self.max_queue_size: + verbose_logger.warning("zerobus: queue at %s rows during a flush, dropped a row", self.max_queue_size) + return + + self.log_queue.append(trace_row(payload)) + self._drop_overflow() + if len(self.log_queue) >= self.batch_size: + self._start_batch_flush_task() + except Exception: # noqa: BLE001 # logging must never break the request path + verbose_logger.exception("zerobus: failed to queue an event") + + def _payload_for(self, kwargs: Mapping[str, object]) -> Mapping[str, object] | None: + """The payload to buffer, redacted the way the framework redacts the success path.""" + details: Final = self.redact_standard_logging_payload_from_model_call_details( + dict(kwargs) # mutable-ok: both framework helpers take the call details as a dict + ) + payload: Final = details.get("standard_logging_object") + if not isinstance(payload, dict): + return None + if should_redact_message_logging(details): + return redacted_standard_logging_payload(payload) + return payload + + def _drop_overflow(self) -> None: + """Trim the oldest rows, except mid flush when the in-flight batch is the head of the queue.""" + if self._flushing: + return + overflow: Final = len(self.log_queue) - self.max_queue_size + if overflow <= 0: + return + del self.log_queue[:overflow] + verbose_logger.warning("zerobus: queue over %s rows, dropped %s oldest", self.max_queue_size, overflow) + + async def flush_queue(self, skip_if_flushing: bool = False) -> None: + if skip_if_flushing and self._flushing: + return + self._flushing = True + try: + await super().flush_queue() + finally: + self._flushing = False + + async def async_send_batch(self) -> None: + """A retryable failure propagates so the rows are kept; a permanent one drops them so the queue moves on.""" + rows: Final = tuple(self.log_queue) + if not rows: + return + failure: Final = await self.client.insert(rows) + if failure is None: + return + if failure.retryable: + raise ZerobusIngestError(failure.detail) + verbose_logger.error("zerobus: dropping %s rows, %s", len(rows), failure.detail) diff --git a/litellm/integrations/zerobus/row.py b/litellm/integrations/zerobus/row.py new file mode 100644 index 00000000000..c4da7975c44 --- /dev/null +++ b/litellm/integrations/zerobus/row.py @@ -0,0 +1,156 @@ +""" +Shape of one Delta table row per LiteLLM request. + +Zerobus validates every record against the target table and rejects unknown columns, so +the row is a fixed set of scalar columns for filtering plus JSON-encoded ``VARIANT`` +columns for anything nested. ``create_table_sql`` renders the matching DDL. +""" + +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + +TRACE_TABLE_COLUMNS: Final[Mapping[str, str]] = MappingProxyType( + { + "id": "STRING", + "trace_id": "STRING", + "session_id": "STRING", + "litellm_call_id": "STRING", + "call_type": "STRING", + "status": "STRING", + "model": "STRING", + "model_group": "STRING", + "model_id": "STRING", + "custom_llm_provider": "STRING", + "api_base": "STRING", + "stream": "BOOLEAN", + "cache_hit": "BOOLEAN", + "start_time": "TIMESTAMP", + "end_time": "TIMESTAMP", + "completion_start_time": "TIMESTAMP", + "response_time": "DOUBLE", + "prompt_tokens": "LONG", + "completion_tokens": "LONG", + "total_tokens": "LONG", + "response_cost": "DOUBLE", + "saved_cache_cost": "DOUBLE", + "api_key_hash": "STRING", + "api_key_alias": "STRING", + "team_id": "STRING", + "team_alias": "STRING", + "user_id": "STRING", + "org_id": "STRING", + "end_user": "STRING", + "requester_ip_address": "STRING", + "user_agent": "STRING", + "request_tags": "VARIANT", + "messages": "VARIANT", + "response": "VARIANT", + "error_str": "STRING", + "error_information": "VARIANT", + "metadata": "VARIANT", + "model_parameters": "VARIANT", + "hidden_params": "VARIANT", + "guardrail_information": "VARIANT", + "cost_breakdown": "VARIANT", + } +) + +_MICROSECONDS: Final = 1_000_000 + + +def create_table_sql(table_name: str) -> str: + columns: Final = ",\n".join(f" {name} {delta_type}" for name, delta_type in TRACE_TABLE_COLUMNS.items()) + return f"CREATE TABLE {table_name} (\n{columns}\n);" + + +def _text(payload: Mapping[str, object], key: str) -> str | None: + value: Final = payload.get(key) + return value if isinstance(value, str) else None + + +def _flag(payload: Mapping[str, object], key: str) -> bool | None: + value: Final = payload.get(key) + return value if isinstance(value, bool) else None + + +def _number(payload: Mapping[str, object], key: str) -> float | None: + value: Final = payload.get(key) + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + return float(value) + + +def _count(payload: Mapping[str, object], key: str) -> int | None: + value: Final = _number(payload, key) + return None if value is None else int(value) + + +def _timestamp_micros(payload: Mapping[str, object], key: str) -> int | None: + """Delta ``TIMESTAMP`` over Zerobus is epoch microseconds; LiteLLM keeps epoch seconds.""" + seconds: Final = _number(payload, key) + if seconds is None or seconds <= 0: + return None + return int(seconds * _MICROSECONDS) + + +def _json(payload: Mapping[str, object], key: str) -> str | None: + value: Final = payload.get(key) + return None if value is None else safe_dumps(value) + + +def _metadata(payload: Mapping[str, object]) -> Mapping[str, object]: + value: Final = payload.get("metadata") + return value if isinstance(value, Mapping) else MappingProxyType({}) + + +def trace_row(payload: Mapping[str, object]) -> Mapping[str, object]: + """One ``TRACE_TABLE_COLUMNS`` row for a ``StandardLoggingPayload``.""" + metadata: Final = _metadata(payload) + return MappingProxyType( + { + "id": _text(payload, "id"), + "trace_id": _text(payload, "trace_id"), + "session_id": _text(payload, "session_id"), + "litellm_call_id": _text(payload, "litellm_call_id"), + "call_type": _text(payload, "call_type"), + "status": _text(payload, "status"), + "model": _text(payload, "model"), + "model_group": _text(payload, "model_group"), + "model_id": _text(payload, "model_id"), + "custom_llm_provider": _text(payload, "custom_llm_provider"), + "api_base": _text(payload, "api_base"), + "stream": _flag(payload, "stream"), + "cache_hit": _flag(payload, "cache_hit"), + "start_time": _timestamp_micros(payload, "startTime"), + "end_time": _timestamp_micros(payload, "endTime"), + "completion_start_time": _timestamp_micros(payload, "completionStartTime"), + "response_time": _number(payload, "response_time"), + "prompt_tokens": _count(payload, "prompt_tokens"), + "completion_tokens": _count(payload, "completion_tokens"), + "total_tokens": _count(payload, "total_tokens"), + "response_cost": _number(payload, "response_cost"), + "saved_cache_cost": _number(payload, "saved_cache_cost"), + "api_key_hash": _text(metadata, "user_api_key_hash"), + "api_key_alias": _text(metadata, "user_api_key_alias"), + "team_id": _text(metadata, "user_api_key_team_id"), + "team_alias": _text(metadata, "user_api_key_team_alias"), + "user_id": _text(metadata, "user_api_key_user_id"), + "org_id": _text(metadata, "user_api_key_org_id"), + "end_user": _text(payload, "end_user"), + "requester_ip_address": _text(payload, "requester_ip_address"), + "user_agent": _text(payload, "user_agent"), + "request_tags": _json(payload, "request_tags"), + "messages": _json(payload, "messages"), + "response": _json(payload, "response"), + "error_str": _text(payload, "error_str"), + "error_information": _json(payload, "error_information"), + "metadata": _json(payload, "metadata"), + "model_parameters": _json(payload, "model_parameters"), + "hidden_params": _json(payload, "hidden_params"), + "guardrail_information": _json(payload, "guardrail_information"), + "cost_breakdown": _json(payload, "cost_breakdown"), + } + ) diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 6294f3bc577..7049fdd1f39 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -52,6 +52,7 @@ from litellm.integrations.vantage.vantage_logger import VantageLogger from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( VectorStorePreCallHook, ) +from litellm.integrations.zerobus import ZerobusLogger from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler from litellm.proxy.hooks.dynamic_rate_limiter_v3 import _PROXY_DynamicRateLimitHandlerV3 @@ -97,6 +98,7 @@ class CustomLoggerRegistry: "deepeval": DeepEvalLogger, "s3_v2": S3Logger, "pointfive": PointFiveLogger, + "zerobus": ZerobusLogger, "aws_sqs": SQSLogger, "dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler, "dynamic_rate_limiter_v3": _PROXY_DynamicRateLimitHandlerV3, diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 22b8d850c83..a0f027cd58f 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -114,10 +114,9 @@ class HealthCheckHelpers: """ Health check for batch mode. - Calls list_batches for providers that support it (openai, hosted_vllm, azure, - vertex_ai). For all other providers (e.g. bedrock) the batch API surface doesn't - include list_batches, so we fall back to acompletion to verify connectivity and - credential validity instead. + Calls list_batches for providers that support it. For all other providers (e.g. bedrock) + the batch API surface doesn't include list_batches, so we fall back to acompletion to + verify connectivity and credential validity instead. """ import litellm @@ -132,10 +131,9 @@ class HealthCheckHelpers: litellm_params={"api_base": api_base} if api_base else None, ) - if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS: - return await litellm.alist_batches(**filtered_model_params) - else: + if custom_llm_provider not in LIST_BATCHES_SUPPORTED_PROVIDERS: return await litellm.acompletion(**model_params) + return await litellm.alist_batches(**{**filtered_model_params, "custom_llm_provider": custom_llm_provider}) @staticmethod async def _image_edit_health_check(edit_request: Callable[[], Awaitable["ImageResponse"]]) -> "ImageResponse": diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 83ab2bc11a2..28d72702f3e 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -20,11 +20,7 @@ from httpx import Response from pydantic import BaseModel, JsonValue import litellm -from litellm import ( - _custom_logger_compatible_callbacks_literal, - json_logs, - turn_off_message_logging, -) +from litellm import _custom_logger_compatible_callbacks_literal from litellm._logging import ( _is_debugging_on, _redact_string, @@ -43,8 +39,7 @@ from litellm.constants import ( DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT, EMPTY_MAPPING, PROVIDER_REQUEST_ID_HEADERS, - SENTRY_DENYLIST, - SENTRY_PII_DENYLIST, + REDACTED_BY_LITELLM, ) from litellm.cost_calculator import ( RealtimeAPITokenUsageProcessor, @@ -213,6 +208,7 @@ from ..integrations.s3 import S3Logger from ..integrations.s3_v2 import S3Logger as S3V2Logger from ..integrations.supabase import Supabase from ..integrations.traceloop import TraceloopLogger +from ..integrations.zerobus import ZerobusLogger from .exception_mapping_utils import _get_response_headers from .initialize_dynamic_callback_params import ( get_trusted_callback_params, @@ -380,9 +376,12 @@ _DEPLOYMENT_PRICING_KEYS: Final = ( "output_cost_per_token", "input_cost_per_token_batches", "output_cost_per_token_batches", + "input_cost_per_token_above_200k_tokens_batches", "input_cost_per_token_above_272k_tokens_batches", + "output_cost_per_token_above_200k_tokens_batches", "output_cost_per_token_above_272k_tokens_batches", "cache_read_input_token_cost_batches", + "cache_read_input_token_cost_above_200k_tokens_batches", "cache_read_input_token_cost_above_272k_tokens_batches", "cache_creation_input_token_cost_batches", "cache_creation_input_token_cost_above_272k_tokens_batches", @@ -1356,10 +1355,19 @@ class Logging(LiteLLMLoggingBaseClass): _litellm_params: Final = self.model_call_details.get("litellm_params", {}) _metadata: Final = _litellm_params.get("metadata", {}) or {} try: - # [Non-blocking Extra Debug Information in metadata] - if turn_off_message_logging is True: - _metadata["raw_request"] = "redacted by litellm. \ - 'litellm.turn_off_message_logging=True'" + self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict( + raw_request_api_base=self._get_masked_api_base(str(additional_args.get("api_base") or "")), + raw_request_body=self._get_raw_request_body(additional_args.get("complete_input_dict", {})), + # NOTE: setting ignore_sensitive_headers to True will cause + # the Authorization header to be leaked when calls to the health + # endpoint are made and fail. + raw_request_headers=self._get_masked_headers( + additional_args.get("headers", {}) or {}, + ), + error=None, + ) + if should_redact_message_logging(self.model_call_details): + _metadata["raw_request"] = REDACTED_BY_LITELLM else: curl_command: Final = self._get_request_curl_command( api_base=additional_args.get("api_base", ""), @@ -1367,20 +1375,7 @@ class Logging(LiteLLMLoggingBaseClass): additional_args=additional_args, data=additional_args.get("complete_input_dict", {}), ) - _metadata["raw_request"] = _redact_string(str(curl_command)) - # split up, so it's easier to parse in the UI - self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict( - raw_request_api_base=self._get_masked_api_base(str(additional_args.get("api_base") or "")), - raw_request_body=self._get_raw_request_body(additional_args.get("complete_input_dict", {})), - # NOTE: setting ignore_sensitive_headers to True will cause - # the Authorization header to be leaked when calls to the health - # endpoint are made and fail. - raw_request_headers=self._get_masked_headers( - additional_args.get("headers", {}) or {}, - ), - error=None, - ) except Exception as e: self.model_call_details["raw_request_typed_dict"] = RawRequestTypedDict( error=str(e), @@ -1474,7 +1469,7 @@ class Logging(LiteLLMLoggingBaseClass): def _print_llm_call_debugging_log( self, api_base: str, - headers: dict, + headers: dict | None, additional_args: dict, ): """ @@ -1483,8 +1478,8 @@ class Logging(LiteLLMLoggingBaseClass): Prints the RAW curl command sent from LiteLLM """ if _is_debugging_on() or self.litellm_request_debug: - if json_logs: - masked_headers: Final = self._get_masked_headers(headers) + if litellm.json_logs: + masked_headers: Final = self._get_masked_headers(headers or {}) masked_api_base: Final = self._get_masked_api_base(str(api_base or "")) if self.litellm_request_debug: verbose_logger.warning( # .warning ensures this shows up in all environments @@ -1561,20 +1556,12 @@ class Logging(LiteLLMLoggingBaseClass): else: attr = "debug" - if json_logs: - callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug - callattr( - "RAW RESPONSE:\n{}\n\n".format( - self.model_call_details.get("original_response", self.model_call_details) - ), - ) - else: - callattr = verbose_logger.warning if attr == "warning" else verbose_logger.debug - callattr( - "RAW RESPONSE:\n{}\n\n".format( - self.model_call_details.get("original_response", self.model_call_details) - ) + callattr: Final = verbose_logger.warning if attr == "warning" else verbose_logger.debug + callattr( + "RAW RESPONSE:\n{}\n\n".format( + self.model_call_details.get("original_response", self.model_call_details) ) + ) if getattr(self, "logger_fn", None) and callable(self.logger_fn): try: self.logger_fn( @@ -4215,6 +4202,9 @@ class Logging(LiteLLMLoggingBaseClass): json_mode=False, litellm_params={}, ) + elif result is None: + verbose_logger.warning("LiteLLM: the anthropic_messages stream assembled no response, logging an empty one") + return litellm.ModelResponse(model=self.model) else: from litellm.types.llms.anthropic import AnthropicResponse @@ -4423,21 +4413,10 @@ def set_callbacks(callback_list, function_id=None): print_verbose("Package 'sentry_sdk' is missing. Installing it...") subprocess.check_call([sys.executable, "-m", "pip", "install", "sentry_sdk"]) import sentry_sdk - from sentry_sdk.scrubber import EventScrubber + from litellm.litellm_core_utils.sentry_scrubbing import build_sentry_init_options sentry_sdk_instance = sentry_sdk - sentry_trace_rate = os.environ.get("SENTRY_API_TRACE_RATE", "1.0") - sentry_sample_rate = ( - os.environ.get("SENTRY_API_SAMPLE_RATE") if "SENTRY_API_SAMPLE_RATE" in os.environ else "1.0" - ) - sentry_sdk_instance.init( - dsn=os.environ.get("SENTRY_DSN"), - traces_sample_rate=float(sentry_trace_rate), - sample_rate=float(sentry_sample_rate if sentry_sample_rate else 1.0), - send_default_pii=False, # Prevent sending Personal Identifiable Information - event_scrubber=EventScrubber(denylist=SENTRY_DENYLIST, pii_denylist=SENTRY_PII_DENYLIST), - environment=os.environ.get("SENTRY_ENVIRONMENT", "production"), - ) + sentry_sdk_instance.init(**build_sentry_init_options(os.environ)) capture_exception = sentry_sdk_instance.capture_exception add_breadcrumb = sentry_sdk_instance.add_breadcrumb elif callback == "slack": @@ -4660,6 +4639,14 @@ def _init_custom_logger_compatible_class( _pointfive_logger: Final = PointFiveLogger() _in_memory_loggers.append(_pointfive_logger) return _pointfive_logger + elif logging_integration == "zerobus": + for callback in _in_memory_loggers: + if isinstance(callback, ZerobusLogger): + return callback + + _zerobus_logger: Final = ZerobusLogger() + _in_memory_loggers.append(_zerobus_logger) + return _zerobus_logger elif logging_integration == "aws_sqs": for callback in _in_memory_loggers: if isinstance(callback, SQSLogger): @@ -5352,6 +5339,10 @@ def get_custom_logger_compatible_class( for callback in _in_memory_loggers: if isinstance(callback, PointFiveLogger): return callback + elif logging_integration == "zerobus": + for callback in _in_memory_loggers: + if isinstance(callback, ZerobusLogger): + return callback elif logging_integration == "aws_sqs": for callback in _in_memory_loggers: if isinstance(callback, SQSLogger): diff --git a/litellm/litellm_core_utils/sentry_scrubbing.py b/litellm/litellm_core_utils/sentry_scrubbing.py new file mode 100644 index 00000000000..4c14cabc2ab --- /dev/null +++ b/litellm/litellm_core_utils/sentry_scrubbing.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +import re +from collections.abc import Callable, Mapping, Sequence +from functools import reduce +from typing import TYPE_CHECKING, Final, TypeAlias, cast + +from pydantic import JsonValue +from sentry_sdk.scrubber import DEFAULT_DENYLIST, DEFAULT_PII_DENYLIST, EventScrubber +from typing_extensions import ReadOnly, TypedDict + +from litellm.constants import ( + LENGTH_OF_LITELLM_GENERATED_KEY, + MINIMUM_CUSTOM_KEY_LENGTH, + SENTRY_DENYLIST, + SENTRY_PII_DENYLIST, +) +from litellm.secret_managers.main import str_to_bool + +if TYPE_CHECKING: + from sentry_sdk.types import Event, Hint + +EventScrubFn: TypeAlias = "Callable[[Event, Hint], Event]" +JsonPath: TypeAlias = tuple[str, ...] + +FILTERED: Final = "[Filtered]" +SEND_DEFAULT_PII_ENV: Final = "SENTRY_SEND_DEFAULT_PII" +SECRET_FIELD_NAMES: Final = tuple(DEFAULT_DENYLIST) + tuple(SENTRY_DENYLIST) +PII_FIELD_NAMES: Final = tuple(DEFAULT_PII_DENYLIST) + tuple(SENTRY_PII_DENYLIST) + +KEY_PREFIX: Final = "sk-" + + +def build_key_pattern(custom_key_minimum: int, generated_key_bytes: int) -> re.Pattern[str]: + generated_suffix_length: Final = (generated_key_bytes * 4 + 2) // 3 + floor: Final = min(custom_key_minimum - len(KEY_PREFIX), generated_suffix_length) + return re.compile(rf"{KEY_PREFIX}[A-Za-z0-9_-]{{{floor},}}") + + +LITELLM_KEY_PATTERN: Final = build_key_pattern(MINIMUM_CUSTOM_KEY_LENGTH, LENGTH_OF_LITELLM_GENERATED_KEY) +SOURCE_CONTEXT_KEYS: Final = frozenset({"pre_context", "context_line", "post_context"}) +STACK_FRAME_PATHS: Final = frozenset( + { + ("exception", "values", "*", "stacktrace", "frames", "*"), + ("threads", "values", "*", "stacktrace", "frames", "*"), + ("stacktrace", "frames", "*"), + } +) +MAX_SCRUB_DEPTH: Final = 64 +EMAIL_PATTERN: Final = re.compile(r"[A-Za-z0-9._%+-]+@[A-Za-z0-9-]+(?:\.[A-Za-z0-9-]+)*\.[A-Za-z]{2,}") +SHA256_HEX_PATTERN: Final = re.compile(r"(? re.Pattern[str]: + names: Final = "|".join(re.escape(name) for name in field_names) + return re.compile( + rf"(?P(?{QUOTED_VALUE}|{BRACKETED_VALUE}|{BARE_VALUE})", + re.IGNORECASE, + ) + + +def build_string_scrubber(send_default_pii: bool) -> Callable[[str], str]: + field_names: Final = SECRET_FIELD_NAMES if send_default_pii else SECRET_FIELD_NAMES + PII_FIELD_NAMES + field_pattern: Final = build_repr_field_pattern(field_names) + value_patterns: Final = ( + (LITELLM_KEY_PATTERN,) if send_default_pii else (LITELLM_KEY_PATTERN, EMAIL_PATTERN, SHA256_HEX_PATTERN) + ) + + def scrub(text: str) -> str: + fields_scrubbed: Final = field_pattern.sub(_filtered_field, text) + return _substitute_all(value_patterns, fields_scrubbed) + + return scrub + + +def _filtered_field(match: re.Match[str]) -> str: + quote: Final = '"' if match.group("value").startswith('"') else "'" + return f"{match.group('field')}{quote}{FILTERED}{quote}" + + +def _substitute_all(patterns: Sequence[re.Pattern[str]], text: str) -> str: + return reduce(lambda scrubbed, pattern: pattern.sub(FILTERED, scrubbed), patterns, text) + + +def scrub_json_strings(value: JsonValue, scrub: Callable[[str], str], path: JsonPath = ()) -> JsonValue: + if len(path) > MAX_SCRUB_DEPTH: + return FILTERED + if isinstance(value, str): + return scrub(value) + if isinstance(value, dict): + unscrubbed_keys: Final = SOURCE_CONTEXT_KEYS if path in STACK_FRAME_PATHS else frozenset[str]() + return { # mutable-ok: JSON object + key: item if key in unscrubbed_keys else scrub_json_strings(item, scrub, (*path, key)) + for key, item in value.items() + } + if isinstance(value, list): + return [scrub_json_strings(item, scrub, (*path, "*")) for item in value] # mutable-ok: JSON array + return value + + +def build_event_scrubber(send_default_pii: bool) -> EventScrubFn: + scrub: Final = build_string_scrubber(send_default_pii) + + def scrub_event(event: Event, _hint: Hint) -> Event: + json_event: Final = cast("JsonValue", event) # cast-ok: [LIT006] the SDK serialized the event to JSON already + return cast("Event", scrub_json_strings(json_event, scrub)) # cast-ok: [LIT006] same JSON shape going back + + return scrub_event + + +def send_default_pii_from_env(env: Mapping[str, str]) -> bool: + return str_to_bool(env.get(SEND_DEFAULT_PII_ENV)) is True + + +def build_sentry_init_options(env: Mapping[str, str]) -> SentryInitOptions: + send_default_pii: Final = send_default_pii_from_env(env) + scrub_event: Final = build_event_scrubber(send_default_pii) + return SentryInitOptions( + dsn=env.get("SENTRY_DSN"), + traces_sample_rate=float(env.get("SENTRY_API_TRACE_RATE") or "1.0"), + sample_rate=float(env.get("SENTRY_API_SAMPLE_RATE") or "1.0"), + send_default_pii=send_default_pii, + event_scrubber=EventScrubber( + denylist=list(SECRET_FIELD_NAMES), # mutable-ok: EventScrubber appends pii_denylist onto denylist in place + pii_denylist=list(PII_FIELD_NAMES), # mutable-ok: EventScrubber takes List[str] + recursive=True, + send_default_pii=send_default_pii, + ), + before_send=scrub_event, + before_send_transaction=scrub_event, + environment=env.get("SENTRY_ENVIRONMENT", "production"), + ) diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index 6c87ef4a3de..43e16599bf6 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -77,13 +77,13 @@ class _CallerHeadersView(TypedDict): headers: ReadOnly[dict[str, str]] -# Globally-routable IPs that are cloud-internal. Everything else -# non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented by -# Python's ``ipaddress`` module). This list only holds IPs that are -# publicly routable *and* point to cloud-fabric services reachable from -# inside a VM via special in-fabric routing. +# Cloud-internal IPs that ``ip.is_global`` can report as public. Everything +# else non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented +# by Python's ``ipaddress`` module). Older Python patch releases (3.12.2, for +# one) treat most of 192.0.0.0/24 as global, so it is listed to block it everywhere. _CLOUD_METADATA_EXCEPTIONS: Final = [ ip_network("168.63.129.16/32"), # Azure Wire Server + ip_network("192.0.0.0/24"), ] _ALLOWED_SCHEMES: Final = ("http", "https") diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index 20753afee5c..24d5b7f366e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -68,15 +68,11 @@ def _error_status_and_message(exc: Exception) -> tuple[int, str]: def _mid_stream_error_sse_event(exc: Exception) -> bytes: from litellm.anthropic_interface.exceptions.exception_mapping_utils import ( - AnthropicExceptionMapping, + anthropic_error_sse_frame, ) status_code, message = _error_status_and_message(exc) - error_response = AnthropicExceptionMapping.transform_to_anthropic_error( - status_code=status_code, - raw_message=message, - ) - return f"event: error\ndata: {json.dumps(error_response)}\n\n".encode() + return anthropic_error_sse_frame(status_code=status_code, raw_message=message).encode() def _delta_payload_field(delta_type: StreamingContentBlockDeltaType) -> str: diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index bd358805743..65a34f72167 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -69,7 +69,9 @@ def make_sync_call( completion_stream: Any = MockResponseIterator(model_response=model_response, json_mode=json_mode) else: decoder: Final = AWSEventStreamDecoder(model=model, json_mode=json_mode) - completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) + completion_stream = decoder.iter_bytes( + response.iter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers + ) # LOGGING logging_obj.post_call( diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index c7b4018b80b..93804e20041 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -1,6 +1,6 @@ import types -from collections.abc import AsyncIterator, Iterator -from typing import Final, cast +from collections.abc import AsyncIterator, Iterator, Mapping +from typing import TYPE_CHECKING, Final, cast import httpx from pydantic import TypeAdapter @@ -51,7 +51,11 @@ from ..common_utils import ( bedrock_tool_name_mappings: Final[InMemoryCache] = InMemoryCache(max_size_in_memory=50, default_ttl=600) from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig +if TYPE_CHECKING: + from botocore.eventstream import EventStreamMessage + converse_config: Final = AmazonConverseConfig() +_STREAM_HEAD_BYTES: Final = 200 NOVA_INVOKE_STREAM_EVENT_TYPES: Final = ( "messageStart", "contentBlockStart", @@ -162,6 +166,22 @@ class AmazonCohereChatConfig: return optional_params +def _stream_decoder( + bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None, + *, + model: str, + json_mode: bool | None, + sync_stream: bool, +) -> "AWSEventStreamDecoder": + if bedrock_invoke_provider == "anthropic": + return AmazonAnthropicClaudeStreamDecoder(model=model, sync_stream=sync_stream, json_mode=json_mode) + if bedrock_invoke_provider == "deepseek_r1": + return AmazonDeepSeekR1StreamDecoder(model=model, sync_stream=sync_stream) + if bedrock_invoke_provider == "moonshot": + return AmazonOpenAICompatibleStreamDecoder(model=model, sync_stream=sync_stream) + return AWSEventStreamDecoder(model=model, json_mode=json_mode) + + async def make_call( client: AsyncHTTPHandler | None, api_base: str, @@ -218,28 +238,13 @@ async def make_call( completion_stream: MockResponseIterator | AsyncIterator[GChunk | ModelResponseStream | dict] = ( MockResponseIterator(model_response=model_response, json_mode=json_mode) ) - elif bedrock_invoke_provider == "anthropic": - decoder: AWSEventStreamDecoder = AmazonAnthropicClaudeStreamDecoder( - model=model, - sync_stream=False, - json_mode=json_mode, - ) - completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size)) - elif bedrock_invoke_provider == "deepseek_r1": - decoder = AmazonDeepSeekR1StreamDecoder( - model=model, - sync_stream=False, - ) - completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size)) - elif bedrock_invoke_provider == "moonshot": - decoder = AmazonOpenAICompatibleStreamDecoder( - model=model, - sync_stream=False, - ) - completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size)) else: - decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode) - completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size)) + decoder: Final = _stream_decoder( + bedrock_invoke_provider, model=model, json_mode=json_mode, sync_stream=False + ) + completion_stream = decoder.aiter_bytes( + response.aiter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers + ) # LOGGING logging_obj.post_call( @@ -322,28 +327,13 @@ def make_sync_call( completion_stream: MockResponseIterator | Iterator[GChunk | ModelResponseStream | dict] = ( MockResponseIterator(model_response=model_response, json_mode=json_mode) ) - elif bedrock_invoke_provider == "anthropic": - decoder: AWSEventStreamDecoder = AmazonAnthropicClaudeStreamDecoder( - model=model, - sync_stream=True, - json_mode=json_mode, - ) - completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) - elif bedrock_invoke_provider == "deepseek_r1": - decoder = AmazonDeepSeekR1StreamDecoder( - model=model, - sync_stream=True, - ) - completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) - elif bedrock_invoke_provider == "moonshot": - decoder = AmazonOpenAICompatibleStreamDecoder( - model=model, - sync_stream=True, - ) - completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) else: - decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode) - completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) + decoder: Final = _stream_decoder( + bedrock_invoke_provider, model=model, json_mode=json_mode, sync_stream=True + ) + completion_stream = decoder.iter_bytes( + response.iter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers + ) # LOGGING logging_obj.post_call( @@ -370,6 +360,49 @@ def make_sync_call( raise BedrockError(status_code=500, message=str(e)) +def _response_header(response_headers: Mapping[str, str] | None, name: str) -> str | None: + return None if response_headers is None else response_headers.get(name) + + +class _EventStreamTally: + def __init__(self) -> None: + self.bytes_received = 0 + self.bytes_decoded = 0 + self.events = 0 + self.head = b"" + + def add_chunk(self, chunk: bytes) -> None: + self.bytes_received += len(chunk) + if len(self.head) < _STREAM_HEAD_BYTES: + self.head = (self.head + chunk)[:_STREAM_HEAD_BYTES] + + def add_event(self, event: "EventStreamMessage") -> None: + self.events += 1 + self.bytes_decoded += event.prelude.total_length + + def undecoded_stream_error(self, response_headers: Mapping[str, str] | None) -> BedrockError | None: + undecoded: Final = self.bytes_received - self.bytes_decoded + if self.events and not undecoded: + return None + detail: Final = ( + f"content-type={_response_header(response_headers, 'content-type')!r}, " + f"x-amzn-requestid={_response_header(response_headers, 'x-amzn-requestid')!r}, " + f"{self.bytes_received} bytes received" + ) + if not self.events: + return BedrockError( + status_code=502, + message=( + "Bedrock answered the stream with HTTP 200 but its body decoded to no events " + f"({detail}, first bytes={self.head!r})" + ), + ) + return BedrockError( + status_code=502, + message=f"Bedrock stream ended with {undecoded} undecoded bytes after {self.events} events ({detail})", + ) + + class AWSEventStreamDecoder: def __init__(self, model: str, json_mode: bool | None = False) -> None: from botocore.parsers import EventStreamJSONParser @@ -709,32 +742,48 @@ class AWSEventStreamDecoder: tool_use=None, ) - def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[GChunk | ModelResponseStream | dict]: + def iter_bytes( + self, iterator: Iterator[bytes], *, response_headers: Mapping[str, str] | None = None + ) -> Iterator[GChunk | ModelResponseStream | dict]: """Given an iterator that yields lines, iterate over it & yield every event encountered""" from botocore.eventstream import EventStreamBuffer event_stream_buffer: Final = EventStreamBuffer() + tally: Final = _EventStreamTally() for chunk in iterator: event_stream_buffer.add_data(chunk) + tally.add_chunk(chunk) for event in event_stream_buffer: + tally.add_event(event) message = self._parse_message_from_event(event) if message: # sse_event = ServerSentEvent(data=message, event="completion") _data = json.loads(message) yield self._chunk_parser(chunk_data=_data) + undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers) + if undecoded_stream_error is not None: + raise undecoded_stream_error - async def aiter_bytes(self, iterator: AsyncIterator[bytes]) -> AsyncIterator[GChunk | ModelResponseStream | dict]: + async def aiter_bytes( + self, iterator: AsyncIterator[bytes], *, response_headers: Mapping[str, str] | None = None + ) -> AsyncIterator[GChunk | ModelResponseStream | dict]: """Given an async iterator that yields lines, iterate over it & yield every event encountered""" from botocore.eventstream import EventStreamBuffer event_stream_buffer: Final = EventStreamBuffer() + tally: Final = _EventStreamTally() async for chunk in iterator: event_stream_buffer.add_data(chunk) + tally.add_chunk(chunk) for event in event_stream_buffer: + tally.add_event(event) message = self._parse_message_from_event(event) if message: _data = json.loads(message) yield self._chunk_parser(chunk_data=_data) + undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers) + if undecoded_stream_error is not None: + raise undecoded_stream_error def _parse_message_from_event(self, event) -> str | None: response_stream_shape: Final = get_bedrock_response_stream_shape() diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 14bd2bee6cf..cefc8afed25 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -770,7 +770,9 @@ class AmazonAnthropicClaudeMessagesConfig( aws_decoder: Final = AmazonAnthropicClaudeMessagesStreamDecoder( model=model, ) - completion_stream: Final = aws_decoder.aiter_bytes(httpx_response.aiter_bytes()) + completion_stream: Final = aws_decoder.aiter_bytes( + httpx_response.aiter_bytes(), response_headers=httpx_response.headers + ) # Convert decoded Bedrock events to Server-Sent Events expected by Anthropic clients. return self.bedrock_sse_wrapper( completion_stream=completion_stream, diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index ac95d1348f9..f2da04a7db7 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -53,14 +53,18 @@ class VertexAIFilesHandler(GCSBucketBase): Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` / ``bucket_name`` and ``vertex_credentials``), mirroring the write path in - ``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global - ``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch - run entirely at the model-group level, so output written to a per-model bucket is - readable without setting the global env vars. + ``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the + ``GCS_BATCH_BUCKET_NAME`` then ``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` + env vars. This lets Vertex batch run entirely at the model-group level, so output + written to a per-model bucket is readable without setting the global env vars. """ params: Final[Mapping[str, object]] = litellm_params or {} bucket_candidate: Final = params.get("gcs_bucket_name") or params.get("bucket_name") - configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME") + configured_bucket_name = ( + bucket_candidate + if isinstance(bucket_candidate, str) + else os.getenv("GCS_BATCH_BUCKET_NAME") or os.getenv("GCS_BUCKET_NAME") + ) credentials: Final = params.get("vertex_credentials") or vertex_credentials if isinstance(credentials, dict): diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index dbb41b57348..2b0694697a4 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -961,7 +961,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def _get_configured_bucket_name(self, litellm_params: dict) -> str: bucket_name: Final = ( - litellm_params.get("gcs_bucket_name") or litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME") + litellm_params.get("gcs_bucket_name") + or litellm_params.get("bucket_name") + or os.getenv("GCS_BATCH_BUCKET_NAME") + or os.getenv("GCS_BUCKET_NAME") ) if not bucket_name: raise ValueError("GCS bucket_name is required") diff --git a/litellm/llms/vertex_ai/rag_engine/ingestion.py b/litellm/llms/vertex_ai/rag_engine/ingestion.py index d9916209a14..c10bac595b6 100644 --- a/litellm/llms/vertex_ai/rag_engine/ingestion.py +++ b/litellm/llms/vertex_ai/rag_engine/ingestion.py @@ -122,41 +122,26 @@ class VertexAIRAGIngestion(BaseRAGIngestion): """ import litellm - # Set GCS_BUCKET_NAME env var for litellm.files.create_file - # The handler uses this to determine where to upload - original_bucket: Final = os.environ.get("GCS_BUCKET_NAME") - if self.gcs_bucket: - os.environ["GCS_BUCKET_NAME"] = self.gcs_bucket + file_tuple: Final = (filename, file_content, content_type) - try: - # Create file tuple for litellm.files.acreate_file - file_tuple: Final = (filename, file_content, content_type) + verbose_logger.debug( + "Uploading file to GCS via litellm.files.acreate_file: %s (bucket: %s)", filename, self.gcs_bucket + ) - verbose_logger.debug( - "Uploading file to GCS via litellm.files.acreate_file: %s (bucket: %s)", filename, self.gcs_bucket - ) + response: Final = await litellm.acreate_file( + file=file_tuple, + purpose="assistants", + custom_llm_provider="vertex_ai", + gcs_bucket_name=self.gcs_bucket, + vertex_project=self.vertex_project, + vertex_location=self.vertex_location, + vertex_credentials=self.vertex_credentials, + ) - # Upload to GCS using LiteLLM's file upload - response: Final = await litellm.acreate_file( - file=file_tuple, - purpose="assistants", # Purpose for file storage - custom_llm_provider="vertex_ai", - vertex_project=self.vertex_project, - vertex_location=self.vertex_location, - vertex_credentials=self.vertex_credentials, - ) + gcs_uri: Final = response.id + verbose_logger.info("Uploaded file to GCS: %s", gcs_uri) - # The response.id should be the GCS URI - gcs_uri: Final = response.id - verbose_logger.info("Uploaded file to GCS: %s", gcs_uri) - - return gcs_uri - finally: - # Restore original env var - if original_bucket is not None: - os.environ["GCS_BUCKET_NAME"] = original_bucket - elif "GCS_BUCKET_NAME" in os.environ: - del os.environ["GCS_BUCKET_NAME"] + return gcs_uri async def _import_file_to_corpus_via_sdk( self, @@ -259,6 +244,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion): content_type: str | None, chunks: list[str], embeddings: list[list[float]] | None, + existing_file_id: str | None = None, ) -> tuple[str | None, str | None]: """ Store content in Vertex AI RAG corpus. @@ -274,6 +260,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion): content_type: MIME type chunks: Ignored - Vertex AI handles chunking embeddings: Ignored - Vertex AI handles embedding + existing_file_id: Existing provider file ID, unsupported for Vertex AI RAG Engine Returns: Tuple of (corpus_id, gcs_uri) diff --git a/tests/test_litellm/litellm_core_utils/audio_utils/__init__.py b/litellm/llms/xai/batches/__init__.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/audio_utils/__init__.py rename to litellm/llms/xai/batches/__init__.py diff --git a/litellm/llms/xai/batches/handler.py b/litellm/llms/xai/batches/handler.py new file mode 100644 index 00000000000..62db1c4833a --- /dev/null +++ b/litellm/llms/xai/batches/handler.py @@ -0,0 +1,195 @@ +from collections.abc import Coroutine +from itertools import chain +from typing import Final + +import httpx +from typing_extensions import NotRequired, ReadOnly, TypedDict + +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, + get_async_httpx_client, +) +from litellm.types.llms.openai import CreateBatchRequest, HttpxBinaryResponseContent +from litellm.types.utils import LiteLLMBatch, LlmProviders + +from .transformation import ( + XAI_RESULTS_PAGE_SIZE, + OpenAIBatchListResponse, + XAIBatch, + XAIBatchList, + XAIBatchResult, + XAIBatchResultsPage, + get_xai_auth_headers, + raise_for_xai_status, + results_to_openai_jsonl, + to_create_batch_body, + to_litellm_batch, + to_openai_batch_list, + xai_batches_url, +) + +_JSONL_CONTENT_TYPE: Final = ("content-type", "application/jsonl") + + +class _PageParams(TypedDict): + limit: ReadOnly[int] + pagination_token: NotRequired[ReadOnly[str]] + + +def _results_params(after: str | None, limit: int | None) -> dict[str, object]: # mutable-ok: httpx params + if after is None: + return dict(_PageParams(limit=limit or XAI_RESULTS_PAGE_SIZE)) # mutable-ok: httpx params + return dict(_PageParams(limit=limit or XAI_RESULTS_PAGE_SIZE, pagination_token=after)) # mutable-ok: httpx params + + +def _flatten(pages: list[XAIBatchResultsPage]) -> tuple[XAIBatchResult, ...]: + return tuple(chain.from_iterable(page.results for page in pages)) + + +def _jsonl_response(url: str, results: tuple[XAIBatchResult, ...]) -> HttpxBinaryResponseContent: + return HttpxBinaryResponseContent( + response=httpx.Response( + status_code=200, + content=results_to_openai_jsonl(results), + headers=(_JSONL_CONTENT_TYPE,), + request=httpx.Request(method="GET", url=url), + ) + ) + + +class XAIBatchesHandler: + def __init__(self, sync_client: HTTPHandler | None = None, async_client: AsyncHTTPHandler | None = None) -> None: + self._sync_client = sync_client + self._async_client = async_client + + def _sync(self, timeout: float | httpx.Timeout) -> HTTPHandler: + return self._sync_client or HTTPHandler(timeout=timeout) + + def _async(self, timeout: float | httpx.Timeout) -> AsyncHTTPHandler: + return self._async_client or get_async_httpx_client( + llm_provider=LlmProviders.XAI, + params={"timeout": timeout}, # mutable-ok: get_async_httpx_client takes a dict + ) + + def create_batch( + self, + _is_async: bool, + create_batch_data: CreateBatchRequest, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + ) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]: + url: Final = xai_batches_url(api_base) + headers: Final = get_xai_auth_headers(api_key=api_key) + body: Final = dict(to_create_batch_body(create_batch_data)) # mutable-ok: httpx json body + endpoint: Final = create_batch_data.get("endpoint") or "/v1/chat/completions" + if _is_async: + + async def _acreate() -> LiteLLMBatch: + response: Final = await self._async(timeout).post(url, json=body, headers=headers, timeout=timeout) + return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()), endpoint) + + return _acreate() + response: Final = self._sync(timeout).post(url, json=body, headers=headers, timeout=timeout) + return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()), endpoint) + + def retrieve_batch( + self, + _is_async: bool, + batch_id: str, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + ) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]: + url: Final = xai_batches_url(api_base, batch_id) + headers: Final = get_xai_auth_headers(api_key=api_key) + if _is_async: + + async def _aretrieve() -> LiteLLMBatch: + response: Final = await self._async(timeout).get(url, headers=headers, timeout=timeout) + return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json())) + + return _aretrieve() + response: Final = self._sync(timeout).get(url, headers=headers, timeout=timeout) + return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json())) + + def cancel_batch( + self, + _is_async: bool, + batch_id: str, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + ) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]: + url: Final = xai_batches_url(api_base, batch_id, suffix=":cancel") + headers: Final = get_xai_auth_headers(api_key=api_key) + if _is_async: + + async def _acancel() -> LiteLLMBatch: + response: Final = await self._async(timeout).post(url, headers=headers, timeout=timeout) + return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json())) + + return _acancel() + response: Final = self._sync(timeout).post(url, headers=headers, timeout=timeout) + return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json())) + + def list_batches( + self, + _is_async: bool, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + after: str | None = None, + limit: int | None = None, + ) -> OpenAIBatchListResponse | Coroutine[None, None, OpenAIBatchListResponse]: + url: Final = xai_batches_url(api_base) + headers: Final = get_xai_auth_headers(api_key=api_key) + params: Final = _results_params(after, limit) + if _is_async: + + async def _alist() -> OpenAIBatchListResponse: + response: Final = await self._async(timeout).get(url, params=params, headers=headers, timeout=timeout) + return to_openai_batch_list(XAIBatchList.model_validate(raise_for_xai_status(response).json())) + + return _alist() + response: Final = self._sync(timeout).get(url, params=params, headers=headers, timeout=timeout) + return to_openai_batch_list(XAIBatchList.model_validate(raise_for_xai_status(response).json())) + + def batch_results_content( + self, + _is_async: bool, + batch_id: str, + api_base: str | None, + api_key: str | None, + timeout: float | httpx.Timeout, + ) -> HttpxBinaryResponseContent | Coroutine[None, None, HttpxBinaryResponseContent]: + url: Final = xai_batches_url(api_base, batch_id, suffix="/results") + headers: Final = get_xai_auth_headers(api_key=api_key) + if _is_async: + + async def _aresults() -> HttpxBinaryResponseContent: + client: Final = self._async(timeout) + + async def _page(after: str | None) -> XAIBatchResultsPage: + response: Final = await client.get( + url, params=_results_params(after, None), headers=headers, timeout=timeout + ) + return XAIBatchResultsPage.model_validate(raise_for_xai_status(response).json()) + + pages = [await _page(None)] # mutable-ok: page walk terminates on the cursor, not on a fixed count + while pages[-1].pagination_token and pages[-1].results: + pages.append(await _page(pages[-1].pagination_token)) + return _jsonl_response(url, _flatten(pages)) + + return _aresults() + client: Final = self._sync(timeout) + + def _page(after: str | None) -> XAIBatchResultsPage: + response: Final = client.get(url, params=_results_params(after, None), headers=headers, timeout=timeout) + return XAIBatchResultsPage.model_validate(raise_for_xai_status(response).json()) + + pages = [_page(None)] # mutable-ok: page walk terminates on the cursor, not on a fixed count + while pages[-1].pagination_token and pages[-1].results: + pages.append(_page(pages[-1].pagination_token)) + return _jsonl_response(url, _flatten(pages)) diff --git a/litellm/llms/xai/batches/transformation.py b/litellm/llms/xai/batches/transformation.py new file mode 100644 index 00000000000..8f305b8c203 --- /dev/null +++ b/litellm/llms/xai/batches/transformation.py @@ -0,0 +1,278 @@ +""" +xAI Batch API reference: https://docs.x.ai/developers/advanced-api-usage/batch-api + +xAI batches carry request counters, not a status, and no output file: results are paged from +``GET /v1/batches/{id}/results``, so LiteLLM hands back the batch id as ``output_file_id``. +""" + +import json +from collections.abc import Mapping, Sequence +from datetime import datetime, timezone +from types import MappingProxyType +from typing import Final, Literal, TypeAlias + +import httpx +from openai.types.batch import BatchRequestCounts +from openai.types.batch import Errors as BatchErrors +from openai.types.batch_error import BatchError +from pydantic import BaseModel, ConfigDict +from typing_extensions import NotRequired, ReadOnly, TypedDict + +from litellm.constants import XAI_API_BASE +from litellm.litellm_core_utils.url_utils import encode_url_path_segment +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.xai.common_utils import XAIModelInfo +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import CreateBatchRequest +from litellm.types.utils import LiteLLMBatch + +OpenAIBatchStatus: TypeAlias = Literal[ + "validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled" +] + +XAI_BATCH_ID_PREFIX: Final = "batch_" +XAI_RESULTS_PAGE_SIZE: Final = 1000 +DEFAULT_BATCH_NAME: Final = "litellm-batch" +DEFAULT_BATCH_ENDPOINT: Final = "/v1/chat/completions" +_EMPTY_HEADERS: Final[Mapping[str, str]] = MappingProxyType({}) + + +class XAIBatchesError(BaseLLMException): + pass + + +def xai_batches_error( + error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers +) -> XAIBatchesError: + return XAIBatchesError( + status_code=status_code, + message=error_message, + headers=headers if isinstance(headers, httpx.Headers) else httpx.Headers(tuple(headers.items())), + ) + + +def raise_for_xai_status(response: httpx.Response) -> httpx.Response: + if response.status_code >= 400: + raise xai_batches_error(response.text, response.status_code, response.headers) + return response + + +def get_xai_api_base(api_base: str | None) -> str: + resolved: Final = (api_base or get_secret_str("XAI_API_BASE") or XAI_API_BASE).rstrip("/") + return resolved.removesuffix("/v1") + + +def get_xai_auth_headers( + headers: Mapping[str, str] = _EMPTY_HEADERS, api_key: str | None = None +) -> dict[str, str]: # mutable-ok: BaseConfig.validate_environment contract returns dict + resolved_key: Final = XAIModelInfo.get_api_key(api_key) + if resolved_key is None: + raise xai_batches_error( + "Missing xAI API Key. Pass api_key, set litellm.xai_key or XAI_API_KEY", 401, _EMPTY_HEADERS + ) + return dict(headers, Authorization=f"Bearer {resolved_key}") # mutable-ok: BaseConfig contract returns dict + + +def xai_batches_url(api_base: str | None, batch_id: str | None = None, suffix: str = "") -> str: + base: Final = f"{get_xai_api_base(api_base)}/v1/batches" + if batch_id is None: + return base + return f"{base}/{encode_url_path_segment(batch_id, field_name='batch_id')}{suffix}" + + +def is_xai_batch_results_id(file_id: str) -> bool: + return file_id.startswith(XAI_BATCH_ID_PREFIX) + + +class XAICreateBatchRequest(TypedDict): + name: ReadOnly[str] + input_file_id: NotRequired[ReadOnly[str]] + + +class XAIBatchState(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + num_requests: int = 0 + num_pending: int = 0 + num_success: int = 0 + num_error: int = 0 + num_cancelled: int = 0 + + +class XAIBatch(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + batch_id: str + name: str = "" + create_time: str | None = None + expire_time: str | None = None + cancel_time: str | None = None + cancel_by_xai_message: str | None = None + state: XAIBatchState = XAIBatchState() + input_file_id: str | None = None + + +class XAIBatchList(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + batches: tuple[XAIBatch, ...] = () + pagination_token: str | None = None + + +class XAIBatchResultError(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + code: int | str | None = None + message: str = "" + + +class XAIBatchResultData(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + response: Mapping[str, Mapping[str, object]] | None = None + error: XAIBatchResultError | None = None + + +class XAIBatchResult(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + batch_request_id: str + batch_result: XAIBatchResultData = XAIBatchResultData() + + +class XAIBatchResultsPage(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + results: tuple[XAIBatchResult, ...] = () + pagination_token: str | None = None + + +def _to_unix_timestamp(value: str | None) -> int | None: + """xAI returns RFC 3339 timestamps over gRPC but a bare ``YYYY-MM-DD`` over REST.""" + if value is None: + return None + try: + parsed: Final = datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError: + return None + return int((parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=timezone.utc)).timestamp()) + + +def xai_batch_status(batch: XAIBatch) -> OpenAIBatchStatus: + """xAI exposes counters, not a status. A batch xAI itself cancelled (input validation failed) is a failure, + a caller-cancelled batch is cancelled, an empty batch is still validating its input file, and a batch + with nothing pending has completed.""" + if batch.cancel_time is not None: + return "failed" if batch.cancel_by_xai_message else "cancelled" + if batch.state.num_requests == 0: + return "validating" + if batch.state.num_pending > 0: + return "in_progress" + return "completed" + + +def to_litellm_batch(batch: XAIBatch, endpoint: str = DEFAULT_BATCH_ENDPOINT) -> LiteLLMBatch: + status: Final = xai_batch_status(batch) + created_at: Final = _to_unix_timestamp(batch.create_time) + cancelled_at: Final = _to_unix_timestamp(batch.cancel_time) + errors: Final = ( + BatchErrors(object="list", data=[BatchError(message=batch.cancel_by_xai_message)]) # mutable-ok: openai type + if batch.cancel_by_xai_message + else None + ) + return LiteLLMBatch( + id=batch.batch_id, + object="batch", + endpoint=endpoint, + input_file_id=batch.input_file_id or "", + completion_window="24h", + status=status, + created_at=created_at if created_at is not None else 0, + expires_at=_to_unix_timestamp(batch.expire_time), + failed_at=cancelled_at if status == "failed" else None, + cancelled_at=cancelled_at if status == "cancelled" else None, + output_file_id=batch.batch_id if status == "completed" else None, + errors=errors, + request_counts=BatchRequestCounts( + total=batch.state.num_requests, + completed=batch.state.num_success, + failed=batch.state.num_error + batch.state.num_cancelled, + ), + metadata={"name": batch.name} if batch.name else None, # mutable-ok: LiteLLMBatch.metadata is a dict + ) + + +class OpenAIBatchListResponse(BaseModel): + model_config = ConfigDict(frozen=True) + + object: Literal["list"] = "list" + data: tuple[LiteLLMBatch, ...] + first_id: str | None + last_id: str | None + has_more: bool + next_page_token: str | None = None + + +def to_openai_batch_list(page: XAIBatchList) -> OpenAIBatchListResponse: + data: Final = tuple(to_litellm_batch(b) for b in page.batches) + return OpenAIBatchListResponse( + data=data, + first_id=data[0].id if data else None, + last_id=data[-1].id if data else None, + has_more=bool(page.pagination_token), + next_page_token=page.pagination_token or None, + ) + + +def to_create_batch_body(create_batch_data: CreateBatchRequest) -> XAICreateBatchRequest: + input_file_id: Final = create_batch_data.get("input_file_id") + if not input_file_id: + raise xai_batches_error("input_file_id is required to create an xAI batch", 400, _EMPTY_HEADERS) + metadata: Final = create_batch_data.get("metadata") + name: Final = metadata.get("name") if metadata else None + return XAICreateBatchRequest(name=name or DEFAULT_BATCH_NAME, input_file_id=input_file_id) + + +class OpenAIBatchOutputError(TypedDict): + code: ReadOnly[str] + message: ReadOnly[str] + + +class OpenAIBatchOutputResponse(TypedDict): + status_code: ReadOnly[int] + request_id: ReadOnly[object] + body: ReadOnly[Mapping[str, object]] + + +class OpenAIBatchOutputLine(TypedDict): + id: ReadOnly[str] + custom_id: ReadOnly[str] + response: ReadOnly[OpenAIBatchOutputResponse | None] + error: ReadOnly[OpenAIBatchOutputError | None] + + +def _result_to_openai_line(result: XAIBatchResult) -> OpenAIBatchOutputLine: + """One output JSONL line. xAI wraps the body in a one-key map named after the endpoint + (``chat_get_completion``, ``responses``, ``image_generation``, ...); the value is the OpenAI body.""" + error: Final = result.batch_result.error + response: Final = result.batch_result.response + body: Final = next(iter(response.values()), None) if response else None + if body is None: + message: Final = error.message if error is not None else "xAI returned no response for this request" + code: Final = str(error.code) if error is not None and error.code is not None else "request_failed" + return OpenAIBatchOutputLine( + id=f"batch_req_{result.batch_request_id}", + custom_id=result.batch_request_id, + response=None, + error=OpenAIBatchOutputError(code=code, message=message), + ) + return OpenAIBatchOutputLine( + id=f"batch_req_{result.batch_request_id}", + custom_id=result.batch_request_id, + response=OpenAIBatchOutputResponse(status_code=200, request_id=body.get("id"), body=body), + error=None, + ) + + +def results_to_openai_jsonl(results: Sequence[XAIBatchResult]) -> bytes: + return "".join(f"{json.dumps(_result_to_openai_line(r), ensure_ascii=False)}\n" for r in results).encode() diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index 33ee727dfab..e686d49e689 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -296,7 +296,7 @@ class XAIChatConfig(OpenAIGPTConfig): except Exception as e: verbose_logger.debug("Error extracting X.AI web search usage: %s", e) - self._fold_reasoning_tokens_into_completion(response) + self.fold_reasoning_tokens_into_completion(response) self._normalize_openai_compatible_usage_totals(getattr(response, "usage", None)) restated_usage: Final = _usage_restated_from_xai_ticks(getattr(response, "usage", None)) if restated_usage is not None: @@ -304,7 +304,7 @@ class XAIChatConfig(OpenAIGPTConfig): return response @staticmethod - def _fold_reasoning_tokens_into_completion( + def fold_reasoning_tokens_into_completion( target: ModelResponse | Usage | dict[str, Any] | None, ) -> None: """Reconcile xAI Usage to the OpenAI invariant. @@ -426,7 +426,7 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler): chunk["choices"] = [{"index": 0, "delta": {}, "finish_reason": None}] if "usage" in chunk and chunk["usage"] is not None: - XAIChatConfig._fold_reasoning_tokens_into_completion(chunk["usage"]) + XAIChatConfig.fold_reasoning_tokens_into_completion(chunk["usage"]) XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"]) parsed_chunk: Final = super().chunk_parser(chunk) diff --git a/tests/test_litellm/litellm_core_utils/llm_response_utils/__init__.py b/litellm/llms/xai/files/__init__.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_response_utils/__init__.py rename to litellm/llms/xai/files/__init__.py diff --git a/litellm/llms/xai/files/transformation.py b/litellm/llms/xai/files/transformation.py new file mode 100644 index 00000000000..dbccca47b25 --- /dev/null +++ b/litellm/llms/xai/files/transformation.py @@ -0,0 +1,247 @@ +""" +xAI Files API reference: https://docs.x.ai/developers/rest-api-reference/inference/files + +xAI stores ``purpose`` as an empty string; LiteLLM reports uploads as ``batch``, the only purpose xAI files serve. +""" + +import time +from collections.abc import Mapping, Sequence +from typing import Final + +import httpx +from openai.types.file_deleted import FileDeleted +from pydantic import BaseModel, ConfigDict +from typing_extensions import ReadOnly, TypedDict + +from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data +from litellm.litellm_core_utils.url_utils import encode_url_path_segment +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.files.transformation import BaseFilesConfig, LiteLLMLoggingObj +from litellm.types.llms.openai import ( + CreateFileRequest, + FileContentRequest, + HttpxBinaryResponseContent, + OpenAICreateFileRequestOptionalParams, + OpenAIFileObject, + OpenAIFilesPurpose, +) +from litellm.types.utils import LlmProviders + +from ..batches.transformation import ( + get_xai_api_base, + get_xai_auth_headers, + raise_for_xai_status, + xai_batches_error, +) + +_NO_QUERY_PARAMS: Final[dict[str, str]] = {} # mutable-ok: BaseFilesConfig request transforms return tuple[str, dict] +_DEFAULT_PURPOSE: Final[OpenAIFilesPurpose] = "batch" + + +class XAIMultipartUpload(TypedDict): + file: ReadOnly[tuple[str, object, str]] + purpose: ReadOnly[tuple[None, str]] + + +class XAIFile(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + id: str + bytes: int = 0 + created_at: int | None = None + filename: str = "" + purpose: str = "" + expires_at: int | None = None + + +class XAIFileList(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + data: tuple[XAIFile, ...] = () + pagination_token: str | None = None + + +class XAIFileDeleted(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + id: str + deleted: bool = True + + +def _to_openai_file_object(file: XAIFile) -> OpenAIFileObject: + return OpenAIFileObject( + id=file.id, + bytes=file.bytes, + created_at=file.created_at if file.created_at is not None else int(time.time()), + filename=file.filename, + object="file", + purpose=_DEFAULT_PURPOSE, + status="uploaded", + expires_at=file.expires_at, + ) + + +def _api_base_from(litellm_params: Mapping[str, object]) -> str: + api_base: Final = litellm_params.get("api_base") + return get_xai_api_base(api_base if isinstance(api_base, str) else None) + + +class XAIFilesConfig(BaseFilesConfig): + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.XAI + + def get_complete_url( + self, + api_base: str | None, + api_key: str | None, + model: str, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + stream: bool | None = None, + ) -> str: + return f"{get_xai_api_base(api_base)}/v1/files" + + def _file_url(self, file_id: str, litellm_params: Mapping[str, object], suffix: str = "") -> str: + encoded_file_id: Final = encode_url_path_segment(file_id, field_name="file_id") + return f"{_api_base_from(litellm_params)}/v1/files/{encoded_file_id}{suffix}" + + def get_error_class( + self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers + ) -> BaseLLMException: + return xai_batches_error(error_message, status_code, headers) + + def validate_environment( + self, + headers: Mapping[str, str], + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + api_key: str | None = None, + api_base: str | None = None, + ) -> dict[str, str]: # mutable-ok: BaseFilesConfig signature + return get_xai_auth_headers(headers, api_key) + + def get_supported_openai_params( + self, model: str + ) -> list[OpenAICreateFileRequestOptionalParams]: # mutable-ok: BaseFilesConfig signature + return ["purpose"] # mutable-ok: BaseFilesConfig signature + + def map_openai_params( + self, + non_default_params: Mapping[str, object], + optional_params: dict[str, object], # mutable-ok: BaseConfig signature, returned as-is + model: str, + drop_params: bool, + ) -> dict[str, object]: # mutable-ok: BaseConfig signature + return optional_params + + def transform_create_file_request( + self, + model: str, + create_file_data: CreateFileRequest, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + ) -> dict[str, object]: # mutable-ok: BaseFilesConfig signature + if "file" not in create_file_data: + raise ValueError("File data is required") + extracted: Final = extract_file_data(create_file_data["file"]) + filename: Final = extracted["filename"] or f"file_{int(time.time())}.jsonl" + content_type: Final = extracted.get("content_type") or "application/octet-stream" + upload: Final = XAIMultipartUpload( + file=(filename, extracted["content"], content_type), + purpose=(None, create_file_data.get("purpose") or _DEFAULT_PURPOSE), + ) + return dict(upload) # mutable-ok: BaseFilesConfig signature + + def transform_create_file_response( + self, + model: str | None, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + litellm_params: Mapping[str, object], + ) -> OpenAIFileObject: + return _to_openai_file_object(XAIFile.model_validate(raise_for_xai_status(raw_response).json())) + + def transform_retrieve_file_request( + self, + file_id: str, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + ) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature + return self._file_url(file_id, litellm_params), _NO_QUERY_PARAMS + + def transform_retrieve_file_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + litellm_params: Mapping[str, object], + ) -> OpenAIFileObject: + return _to_openai_file_object(XAIFile.model_validate(raise_for_xai_status(raw_response).json())) + + def transform_delete_file_request( + self, + file_id: str, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + ) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature + return self._file_url(file_id, litellm_params), _NO_QUERY_PARAMS + + def transform_delete_file_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + litellm_params: Mapping[str, object], + ) -> FileDeleted: + deleted: Final = XAIFileDeleted.model_validate(raise_for_xai_status(raw_response).json()) + return FileDeleted(id=deleted.id, deleted=deleted.deleted, object="file") + + def transform_list_files_request( + self, + purpose: str | None, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + ) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature + return f"{_api_base_from(litellm_params)}/v1/files", _NO_QUERY_PARAMS + + def transform_list_files_next_request( + self, + raw_response: httpx.Response, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + ) -> tuple[str, dict[str, str]] | None: # mutable-ok: BaseFilesConfig signature + page: Final = XAIFileList.model_validate(raw_response.json()) + if not page.pagination_token or not page.data: + return None + return f"{_api_base_from(litellm_params)}/v1/files", {"pagination_token": page.pagination_token} + + def transform_list_files_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + litellm_params: Mapping[str, object], + ) -> list[OpenAIFileObject]: # mutable-ok: BaseFilesConfig signature + return [ # mutable-ok: BaseFilesConfig signature + _to_openai_file_object(f) + for f in XAIFileList.model_validate(raise_for_xai_status(raw_response).json()).data + ] + + def transform_file_content_request( + self, + file_content_request: FileContentRequest, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object], + ) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature + file_id: Final = file_content_request.get("file_id") + if file_id is None: + raise ValueError("file_id is required to download file content") + return self._file_url(file_id, litellm_params, suffix="/content"), _NO_QUERY_PARAMS + + def transform_file_content_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + litellm_params: Mapping[str, object], + ) -> HttpxBinaryResponseContent: + return HttpxBinaryResponseContent(response=raw_response) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index c252fe68a7b..c9171872734 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -41488,65 +41488,63 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 8.44944e-07, + "cache_read_input_token_cost": 3.828e-08, + "input_cost_per_token": 4.5936e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.689888e-06, + "output_cost_per_token": 9.1872e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.0412e-08, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 1.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":1.5e-7,"output_cost_per_token":6e-7,"cache_read_input_token_cost":3e-9}, + "output_cost_per_token": 4.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_vision": true, "supports_pdf_input": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 4.62e-07, + "cache_read_input_token_cost": 8.8e-09, + "input_cost_per_token": 2.64e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.386e-06, + "output_cost_per_token": 7.92e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 1.54e-08, - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":0.00000132,"output_cost_per_token":0.00000396,"cache_read_input_token_cost":4.4e-8}, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, @@ -42766,14 +42764,14 @@ "openrouter/qwen/qwen3-coder-plus": { "cache_creation_input_token_cost": 8.125e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, - "input_cost_per_token_above_32k_tokens": 1.17e-06, "cache_creation_input_token_cost_above_32k_tokens": 1.4625e-06, - "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, - "output_cost_per_token_above_32k_tokens": 5.85e-06, "cache_read_input_token_cost": 1.3e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, + "deprecation_date": "2026-10-09", "input_cost_per_token": 6.5e-07, "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.17e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 65536, @@ -42781,6 +42779,7 @@ "mode": "chat", "output_cost_per_token": 3.25e-06, "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 5.85e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42813,6 +42812,7 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.3e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -43110,25 +43110,25 @@ "supports_web_search": false }, "openrouter/z-ai/glm-4.7": { - "input_cost_per_token": 4e-07, - "output_cost_per_token": 1.75e-06, "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 8e-08, + "cache_read_input_token_cost": 1.1e-07, + "input_cost_per_token": 6e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", + "output_cost_per_token": 2.2e-06, "source": "https://openrouter.ai/api/v1/models", - "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_vision": false, - "supports_prompt_caching": true, "supports_assistant_prefill": true, "supports_audio_input": false, + "supports_function_calling": true, "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/z-ai/glm-4.7-flash": { @@ -43173,15 +43173,15 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.1": { - "input_cost_per_token": 9.66e-07, - "output_cost_per_token": 3.036e-06, - "cache_read_input_token_cost": 1.794e-07, "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 1.7914e-07, + "input_cost_per_token": 9.646e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", + "output_cost_per_token": 3.0316e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -51447,13 +51447,16 @@ }, "xai/grok-4.20-0309-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -51461,8 +51464,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true @@ -51491,9 +51497,13 @@ "xai/grok-4.3": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_image_token": 1.25e-06, "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, @@ -51501,6 +51511,8 @@ "mode": "chat", "output_cost_per_token": 2.5e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -51513,9 +51525,13 @@ "xai/grok-4.3-latest": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_image_token": 1.25e-06, "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, @@ -51523,6 +51539,8 @@ "mode": "chat", "output_cost_per_token": 2.5e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -59494,13 +59512,16 @@ }, "xai/grok-4.20-0309-non-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -59508,20 +59529,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-multi-agent-0309": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "responses", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": false, "supports_prompt_caching": true, @@ -59530,8 +59557,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true, "supported_endpoints": [ @@ -60224,18 +60254,18 @@ "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/deepseek-v4p1-flash": { - "cache_read_input_token_cost": 7e-09, - "cache_read_input_token_cost_priority": 8.75e-09, - "input_cost_per_token": 2.2e-07, - "input_cost_per_token_priority": 2.75e-07, + "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost_priority": 7.5e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 3.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 6.6e-07, - "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.5e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60245,13 +60275,16 @@ }, "fireworks_ai/accounts/fireworks/routers/deepseek-v4p1-flash-us": { "cache_read_input_token_cost": 9e-09, + "cache_read_input_token_cost_priority": 1.125e-08, "input_cost_per_token": 4.5e-07, + "input_cost_per_token_priority": 5.625e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 1.8e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, @@ -60262,13 +60295,16 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60320,18 +60356,18 @@ "supports_vision": false }, "fireworks_ai/deepseek-v4p1-flash": { - "cache_read_input_token_cost": 7e-09, - "cache_read_input_token_cost_priority": 8.75e-09, - "input_cost_per_token": 2.2e-07, - "input_cost_per_token_priority": 2.75e-07, + "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost_priority": 7.5e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 3.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 6.6e-07, - "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.5e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60341,13 +60377,16 @@ }, "fireworks_ai/deepseek-v4p1-flash-us": { "cache_read_input_token_cost": 9e-09, + "cache_read_input_token_cost_priority": 1.125e-08, "input_cost_per_token": 4.5e-07, + "input_cost_per_token_priority": 5.625e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 1.8e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, @@ -60358,13 +60397,16 @@ }, "fireworks_ai/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60493,14 +60535,17 @@ }, "fireworks_ai/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -60542,14 +60587,17 @@ }, "fireworks_ai/accounts/fireworks/models/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -62798,13 +62846,16 @@ }, "xai/grok-4.20": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -62812,21 +62863,27 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true }, "xai/grok-4.20-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -62834,21 +62891,27 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true }, "xai/grok-4.20-reasoning-latest": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -62856,8 +62919,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true @@ -63078,13 +63144,16 @@ }, "xai/grok-4.20-non-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -63092,20 +63161,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-non-reasoning-latest": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -63113,20 +63188,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-multi-agent": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "responses", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supported_endpoints": [ "/v1/responses" @@ -63138,20 +63219,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-multi-agent-latest": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "responses", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supported_endpoints": [ "/v1/responses" @@ -63163,8 +63250,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, @@ -66105,23 +66195,23 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3-flash": { + "cache_read_input_token_cost": 1e-08, "input_cost_per_token": 4.5e-08, - "output_cost_per_token": 6e-07, - "cache_read_input_token_cost": 2.85e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 1.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { @@ -66282,24 +66372,24 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "input_cost_per_token": 3e-08, - "output_cost_per_token": 3.2e-07, "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 2.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 3.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, "supports_parallel_function_calling": true, "supports_pdf_input": false, - "supports_vision": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.7-flash": { @@ -66772,46 +66862,47 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.6-max-preview": { - "input_cost_per_token": 1.027e-06, - "output_cost_per_token": 6.162e-06, "cache_creation_input_token_cost": 1.28375e-06, - "input_cost_per_token_above_128k_tokens": 1.58e-06, - "output_cost_per_token_above_128k_tokens": 9.48e-06, "cache_creation_input_token_cost_above_128k_tokens": 1.975e-06, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.027e-06, + "input_cost_per_token_above_128k_tokens": 1.58e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 6.162e-06, + "output_cost_per_token_above_128k_tokens": 9.48e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.6-27b": { - "input_cost_per_token": 3.2e-07, - "output_cost_per_token": 2.7e-06, "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 3.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 262140, "max_tokens": 262140, "mode": "chat", + "output_cost_per_token": 3.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/openai/gpt-5.5-pro": { @@ -66856,23 +66947,23 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 4.9e-08, - "output_cost_per_token": 9.8e-08, - "cache_read_input_token_cost": 9.8e-09, + "cache_read_input_token_cost": 9.408e-09, + "input_cost_per_token": 4.704e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", + "output_cost_per_token": 9.408e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { @@ -66897,22 +66988,22 @@ "supports_web_search": false }, "openrouter/google/gemma-4-26b-a4b-it": { - "cache_read_input_token_cost": 5e-08, - "input_cost_per_token": 9e-08, - "output_cost_per_token": 3e-07, + "cache_read_input_token_cost": 3.75e-08, + "input_cost_per_token": 6.75e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", + "output_cost_per_token": 2.25e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": true, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67197,25 +67288,26 @@ "supports_video_input": true }, "openrouter/qwen/qwen3-max-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67467,59 +67559,62 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-vl-32b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.04e-07, - "output_cost_per_token": 4.16e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.16e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.8e-07, - "output_cost_per_token": 2.1e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.1e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67549,40 +67644,41 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-30b-a3b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-30b-a3b-instruct": { - "input_cost_per_token": 1.3e-07, - "output_cost_per_token": 5.2e-07, + "input_cost_per_token": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67606,21 +67702,22 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-235b-a22b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4e-07, - "output_cost_per_token": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67645,32 +67742,33 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-max": { - "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, - "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, - "cache_read_input_token_cost": 1.56e-07, "cache_creation_input_token_cost": 9.75e-07, - "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, + "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, + "cache_read_input_token_cost": 1.56e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 7.8e-07, + "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, - "supports_response_schema": true, - "supports_vision": false, "supports_prompt_caching": true, "supports_reasoning": false, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.1-terminus": { @@ -67765,23 +67863,24 @@ "openrouter/qwen/qwen-plus-2025-07-28": { "cache_creation_input_token_cost": 3.25e-07, "cache_read_input_token_cost": 5.2e-08, + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.6e-07, - "output_cost_per_token": 7.8e-07, "input_cost_per_token_above_256k_tokens": 7.8e-07, - "output_cost_per_token_above_256k_tokens": 2.34e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 7.8e-07, + "output_cost_per_token_above_256k_tokens": 2.34e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67805,21 +67904,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 81920, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68128,21 +68228,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-8b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68185,21 +68286,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4.55e-07, - "output_cost_per_token": 1.82e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 1.82e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -71268,6 +71370,23 @@ "output_cost_per_token": 0.0, "source": "https://openrouter.ai/typesafe/jev-1.13" }, + "openrouter/typesafe/jev-router": { + "input_cost_per_token": 0, + "output_cost_per_token": 0, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "source": "https://openrouter.ai/typesafe/jev-router", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_audio_input": true, + "supports_video_input": true + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", @@ -72952,14 +73071,14 @@ "supports_web_search": false }, "openrouter/inclusionai/ling-3.0-flash-vl": { - "cache_read_input_token_cost": 1.2e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.8e-07, + "output_cost_per_token": 6.16e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -75470,13 +75589,16 @@ }, "xai/grok-4.20-0309": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -75484,8 +75606,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true @@ -77089,5 +77214,24 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_tool_choice": false + }, + "openrouter/perceptron/perceptron-mk1.5": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 36864, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": false, + "supports_pdf_input": false, + "supports_prompt_caching": false, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": false, + "supports_vision": true, + "supports_web_search": false } } diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index fc67aa8c553..f600c7817f2 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -5,6 +5,8 @@ from typing import Final from pydantic import TypeAdapter +from litellm.litellm_core_utils.request_timeout_resolver import get_configured_request_timeout + DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0 _SECONDS: Final = TypeAdapter(float) @@ -48,8 +50,8 @@ def resolve_llm_passthrough_timeout( Anthropic /v1/messages). Non-streaming precedence: kwargs timeout/request_timeout -> litellm_params - timeout/request_timeout -> router_timeout -> general_settings.pass_through_request_timeout - -> 600s default. + timeout/request_timeout -> router_timeout -> litellm.request_timeout (litellm_settings.request_timeout, + when explicitly set) -> general_settings.pass_through_request_timeout -> 600s default. Streaming (``kwargs["stream"]`` truthy) resolves ``stream_timeout`` at every level before any generic timeout, matching ``Router._get_stream_timeout`` on the completion route: @@ -73,6 +75,7 @@ def resolve_llm_passthrough_timeout( deployment.get("timeout"), deployment.get("request_timeout"), router_timeout, + get_configured_request_timeout(), ) winner: Final = next((val for val in candidates if val is not None), None) return resolve_pass_through_request_timeout() if winner is None else _SECONDS.validate_python(winner) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 89fa0058644..4597872d84e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4019,6 +4019,18 @@ class AllCallbacks(LiteLLMPydanticObjectBase): ], ) + zerobus: CallbackOnUI = CallbackOnUI( + litellm_callback_name="zerobus", + ui_callback_name="Databricks Zerobus", + litellm_callback_params=[ # mutable-ok: the registry field is typed list + "ZEROBUS_WORKSPACE_URL", + "ZEROBUS_SERVER_ENDPOINT", + "ZEROBUS_CLIENT_ID", + "ZEROBUS_CLIENT_SECRET", + "ZEROBUS_TABLE_NAME", + ], + ) + class HTTPExceptionErrorDetail(TypedDict): """The `{"error": }` shape most proxy endpoints raise as `HTTPException.detail`.""" @@ -5355,11 +5367,12 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): default=False, description=( "When True, users whose JWT contains no team claims are authenticated " - "using their database team memberships instead of receiving HTTP 403. " - "Usage is attributed to the user's first resolvable DB team, or to the " - "team specified via the x-litellm-team-id request header (validated " - "against DB membership). Requires user_id_upsert=True so that user " - "records exist before the fallback runs." + "using their database team memberships instead of receiving HTTP 403, " + "with usage attributed to the user's first resolvable DB team. Whether or " + "not the JWT carries team claims, the x-litellm-team-id request header may " + "select any team the user is a member of in the database (validated against " + "DB membership); without the header the JWT team stays the default. Requires " + "user_id_upsert=True so that user records exist before the fallback runs." ), ) issuers: list[JWTIssuerConfig] | None = Field( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index e07b20fd5d5..4f41b283a33 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1930,12 +1930,12 @@ class JWTAuthManager: ) -> HeaderTeam | None: """ The team named by x-litellm-team-id, which may carry a team id or a team - alias. A value that is already an allowed team id (or, under the DB - fallback, an existing team id) never costs an alias lookup; an alias is - accepted only when the team it names would have been accepted by id. - Under the DB fallback only a team row that is provably absent falls - through to the alias lookup; a read that failed for any other reason - keeps the membership denial the id path already gives. + alias. A value that is already an allowed team id never costs a lookup; + under the DB fallback any other value is accepted provisionally, by id + or alias, for the membership check auth_builder runs later. Under the + DB fallback only a team row that is provably absent falls through to + the alias lookup; a read that failed for any other reason keeps the + membership denial the id path already gives. Raises: HTTPException: 403 when neither the value nor the team it aliases is @@ -1948,7 +1948,11 @@ class JWTAuthManager: if not header_value: return None - if fallback_to_db_teams and not allowed_team_ids: + if header_value in allowed_team_ids: + verbose_proxy_logger.debug("Using team_id from x-litellm-team-id header: %s", header_value) + return HeaderTeam(header_value=header_value, team_id=header_value) + + if fallback_to_db_teams: try: await get_team_object( team_id=header_value, @@ -1969,10 +1973,6 @@ class JWTAuthManager: JWTAuthManager._raise_header_team_membership_denial(header_value) return HeaderTeam(header_value=header_value, team_id=header_value) - if header_value in allowed_team_ids: - verbose_proxy_logger.debug("Using team_id from x-litellm-team-id header: %s", header_value) - return HeaderTeam(header_value=header_value, team_id=header_value) - team_id_by_alias: Final = await JWTAuthManager._team_id_by_alias( header_value, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj ) @@ -2353,9 +2353,9 @@ class JWTAuthManager: header_value: str, ) -> None: """ - A provisional team_id from the x-litellm-team-id header (accepted without - JWT-team validation when the JWT carries no team claims) must exist in the - user's DB team memberships before it becomes request context. The denial + A provisional team_id from the x-litellm-team-id header (accepted under + fallback_to_db_teams because it is outside the JWT's teams) must exist in + the user's DB team memberships before it becomes request context. The denial names `header_value`, the id or alias the caller sent, not `team_id`. """ user_team_ids: Final = user_object.teams if user_object else [] @@ -2587,22 +2587,30 @@ class JWTAuthManager: if specific_team_id and not db_team_fallback: all_team_ids.add(specific_team_id) + header_db_fallback: Final = handler.litellm_jwtauth.fallback_to_db_teams and team_id is None + header_team: Final = await JWTAuthManager.resolve_team_from_header( request_headers=request_headers, allowed_team_ids=all_team_ids, - fallback_to_db_teams=db_team_fallback, + fallback_to_db_teams=header_db_fallback, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) + provisional_header_team: Final = ( + header_team + if header_team is not None and header_db_fallback and header_team.team_id not in all_team_ids + else None + ) if header_team: team_id = header_team.team_id - # A provisional header team (accepted only because the JWT carries no - # team claims) is validated against DB membership further down; never - # upsert it here or an attacker-supplied x-litellm-team-id would create - # an orphaned team row before that check runs. A genuine membership team - # already exists, so suppressing the upsert in that case costs nothing. + # A provisional header team (accepted because it is outside the + # JWT's teams under fallback_to_db_teams) is validated against DB + # membership further down; never upsert it here or an + # attacker-supplied x-litellm-team-id would create an orphaned team + # row before that check runs. A genuine membership team already + # exists, so suppressing the upsert in that case costs nothing. try: team_object = await get_team_object( team_id=team_id, @@ -2610,10 +2618,10 @@ class JWTAuthManager: user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, - team_id_upsert=(team_id_upsert and not db_team_fallback), + team_id_upsert=(team_id_upsert and provisional_header_team is None), ) except HTTPException: - if not db_team_fallback: + if provisional_header_team is None: raise JWTAuthManager._raise_header_team_membership_denial(header_team.header_value) elif not team_id and not db_team_fallback: @@ -2756,11 +2764,11 @@ class JWTAuthManager: proxy_logging_obj=proxy_logging_obj, team_id_upsert=team_id_upsert, ) - elif db_team_fallback and header_team is not None and team_id == header_team.team_id: + elif provisional_header_team is not None and team_id == provisional_header_team.team_id: JWTAuthManager._validate_header_team_in_db_membership( team_id=team_id, user_object=user_object, - header_value=header_team.header_value, + header_value=provisional_header_team.header_value, ) if not JWTAuthManager._is_team_route_allowed( route=route, @@ -2770,7 +2778,7 @@ class JWTAuthManager: raise HTTPException( status_code=403, detail=( - f"Team '{header_team.header_value}' (from x-litellm-team-id header) " + f"Team '{provisional_header_team.header_value}' (from x-litellm-team-id header) " f"is not allowed to access route '{route}'." ), ) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 4f9b6b3a96f..64b0c6c1967 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -32,6 +32,7 @@ from starlette.types import Receive, Scope, Send import litellm from litellm._logging import redact_internal_details_from_client_message, verbose_proxy_logger from litellm._uuid import uuid +from litellm.anthropic_interface.exceptions import AnthropicErrorSseFrame, anthropic_error_sse_frame from litellm.constants import ( DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE, DEFAULT_MAX_RECURSE_DEPTH, @@ -102,8 +103,11 @@ from litellm.proxy.common_utils.openai_error_payload import ( ) from litellm.proxy.common_utils.sse_keepalive import ( SSE_COMMENT_PING_BYTES, + SSE_STREAM_START_TAIL, + advance_sse_tail, coerce_keepalive_interval, resolve_ttft_keepalive_interval, + seal_open_sse_frame, wrap_sse_stream_with_keepalive_pings, ) from litellm.proxy.dd_span_tagger import DDSpanTagger @@ -999,6 +1003,17 @@ async def create_response( first_chunk_value = await _buffer_first_chunk_honoring_disconnect(generator, request) resolved_headers: Final = await _resolve_stream_headers(headers, refresh_headers) + if isinstance(first_chunk_value, AnthropicErrorSseFrame): + with contextlib.suppress(Exception): + await generator.aclose() + return JSONResponse( + status_code=first_chunk_value.status_code, + content=first_chunk_value.json_body( + error_body_call_id(general_settings, resolved_headers.get(LITELLM_CALL_ID_HEADER)) + ), + headers=resolved_headers, + ) + if first_chunk_value is not None: try: error_code_from_chunk: Final = await _parse_event_data_for_error(first_chunk_value) @@ -3852,6 +3867,7 @@ class ProxyBaseLLMRequestProcessing: serialize_error: StreamErrorSerializer, request: Request | None = None, flush_tail: Callable[[], bytes] | None = None, + seal_open_frame: Callable[[bytes], str] | None = None, ) -> AsyncGenerator[str, None]: """ Shared streaming data generator: runs proxy iterator hook, per-chunk hook, @@ -3861,6 +3877,12 @@ class ProxyBaseLLMRequestProcessing: ``flush_tail`` runs once after the upstream iterator completes cleanly and its non-empty result is yielded, so a serializer that buffers bytes across chunks can emit anything still held at end of stream. + + ``seal_open_frame`` is given the tail of what has been yielded when the + error frame goes out, and what it returns is written first. A passthrough + relays raw upstream bytes, so an upstream that hangs up mid-frame leaves the + client inside an open frame, where an error frame would be swallowed or + misparsed instead of raised. """ verbose_proxy_logger.debug("inside generator") # Resolve per-stream (not per-chunk) whether the heavy per-chunk path @@ -3877,6 +3899,7 @@ class ProxyBaseLLMRequestProcessing: stream_completed = False client_disconnected = False delivered_chunk = False + recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes try: str_so_far = "" async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( @@ -3922,7 +3945,9 @@ class ProxyBaseLLMRequestProcessing: # False and refunds. A keepalive ping carries no provider output, # so it must not suppress that refund. delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES - yield serialize_chunk(chunk) + serialized = serialize_chunk(chunk) + recent_tail = advance_sse_tail(recent_tail, serialized) + yield serialized held_tail: Final = flush_tail() if flush_tail is not None else b"" if held_tail: yield serialize_chunk(held_tail) @@ -3970,7 +3995,9 @@ class ProxyBaseLLMRequestProcessing: code=stream_error_status, ) stream_completed = True - yield serialize_error(proxy_exception) + error_frame: Final = serialize_error(proxy_exception) + seal: Final = "" if seal_open_frame is None else seal_open_frame(recent_tail) + yield seal + error_frame if seal else error_frame finally: await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( request=request, @@ -3992,7 +4019,7 @@ class ProxyBaseLLMRequestProcessing: restamp_model: str | None = None, ) -> AsyncGenerator[str, None]: """ - Anthropic /messages and Google /generateContent streaming data generator require SSE events. + Anthropic /messages streaming data generator, which requires SSE events. Returns the underlying ``async_streaming_data_generator`` configured with SSE serializers directly (rather than re-wrapping it in another @@ -4010,11 +4037,13 @@ class ProxyBaseLLMRequestProcessing: request_data=request_data, proxy_logging_obj=proxy_logging_obj, serialize_chunk=ProxyBaseLLMRequestProcessing._sse_chunk_serializer(restamper), - serialize_error=lambda proxy_exc: ( - f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n" + serialize_error=lambda proxy_exc: anthropic_error_sse_frame( + status_code=error_status_code(proxy_exc, status.HTTP_500_INTERNAL_SERVER_ERROR), + raw_message=proxy_exc.message, ), request=request, flush_tail=None if restamper is None else restamper.flush, + seal_open_frame=seal_open_sse_frame, ) @overload diff --git a/litellm/proxy/common_utils/sse_keepalive.py b/litellm/proxy/common_utils/sse_keepalive.py index cf98a7e9224..d9685971f52 100644 --- a/litellm/proxy/common_utils/sse_keepalive.py +++ b/litellm/proxy/common_utils/sse_keepalive.py @@ -15,7 +15,7 @@ SSE_COMMENT_PING_BYTES: Final = SSE_COMMENT_PING.encode() # terminates a line with CRLF, LF or CR, so a blank line is any of these three. _SSE_FRAME_DELIMITERS: Final = (b"\r\n\r\n", b"\n\n", b"\r\r") _SSE_DELIMITER_LOOKBACK: Final = max(len(delimiter) for delimiter in _SSE_FRAME_DELIMITERS) -_STREAM_START_TAIL: Final = b"\n\n" +SSE_STREAM_START_TAIL: Final = b"\n\n" _SSE_MEDIA_TYPE: Final = "text/event-stream" @@ -128,7 +128,7 @@ async def _keepalive_ping_byte_stream( # Seeded as a delimiter because a stream starts at a frame boundary, and kept # across chunks because a delimiter can be split between two transport reads, # which testing only the latest chunk would miss for the rest of the stream. - recent_tail = _STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes + recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes try: while True: await asyncio.wait((pending,), timeout=ping_interval_seconds) @@ -155,6 +155,28 @@ async def _keepalive_ping_byte_stream( await stream.aclose() +def advance_sse_tail(recent_tail: bytes, chunk: object) -> bytes: + written: Final = _sse_tail_bytes(chunk) + if not written: + return recent_tail + return (recent_tail + written)[-_SSE_DELIMITER_LOOKBACK:] + + +def _sse_tail_bytes(chunk: object) -> bytes: + if isinstance(chunk, bytes): + return chunk[-_SSE_DELIMITER_LOOKBACK:] + if isinstance(chunk, str): + return chunk[-_SSE_DELIMITER_LOOKBACK:].encode() + return b"" + + +def seal_open_sse_frame(recent_tail: bytes) -> str: + if recent_tail.endswith(_SSE_FRAME_DELIMITERS): + return "" + line_break: Final = "" if recent_tail.endswith((b"\n", b"\r")) else "\n" + return f"{line_break}{ANTHROPIC_PING_SSE_CHUNK}" + + def resolve_ttft_keepalive_interval( deployments: Iterable[Mapping[str, object]], global_interval: float | str | None, diff --git a/litellm/proxy/common_utils/validation_error_body.py b/litellm/proxy/common_utils/validation_error_body.py new file mode 100644 index 00000000000..b21f33a2434 --- /dev/null +++ b/litellm/proxy/common_utils/validation_error_body.py @@ -0,0 +1,13 @@ +from collections.abc import Sequence + +from typing_extensions import ReadOnly, TypedDict + + +class ValidationErrorDetail(TypedDict): + type: ReadOnly[str] + loc: ReadOnly[tuple[int | str, ...]] + msg: ReadOnly[str] + + +def public_validation_errors(errors: Sequence[ValidationErrorDetail]) -> tuple[ValidationErrorDetail, ...]: + return tuple(ValidationErrorDetail(type=error["type"], loc=error["loc"], msg=error["msg"]) for error in errors) diff --git a/litellm/proxy/list_api/common.py b/litellm/proxy/list_api/common.py index daa6414fd94..efa8a271459 100644 --- a/litellm/proxy/list_api/common.py +++ b/litellm/proxy/list_api/common.py @@ -8,8 +8,8 @@ from fastapi import Request from fastapi.dependencies.utils import get_flat_params from fastapi.params import ParamTypes from fastapi.responses import JSONResponse -from typing_extensions import ReadOnly, TypedDict +from litellm.proxy.common_utils.validation_error_body import ValidationErrorDetail from litellm.types.proxy.management_endpoints.management_v1 import ( ListLinks, PageLinks, @@ -58,14 +58,6 @@ def escape_like(value: str) -> str: return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") -class ValidationErrorDetail(TypedDict): - """The keys of a pydantic/FastAPI validation error a problem document needs.""" - - type: ReadOnly[str] - loc: ReadOnly[tuple[int | str, ...]] - msg: ReadOnly[str] - - def _is_length_error_of_rejected_items(error: ValidationErrorDetail, errors: Sequence[ValidationErrorDetail]) -> bool: """pydantic counts only items that validated, so a bad item also trips the parent's min_length.""" return error["type"] == "too_short" and any( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 61304f0d919..646eca071d1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -476,6 +476,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( project_spend_counter_key, tag_cache_key, ) +from litellm.proxy.common_utils.validation_error_body import public_validation_errors from litellm.proxy.config_resolvers import ( FieldSource, SettingsStore, @@ -551,7 +552,6 @@ from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_sp from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.list_api.common import ( ManagementProblem, - ValidationErrorDetail, problem_response, request_validation_problem, ) @@ -1983,16 +1983,14 @@ class _ExceptionRow(TypedDict, total=False): @app.exception_handler(RequestValidationError) async def otel_request_validation_exception_handler(request: Request, exc: RequestValidationError): + public_errors: Final = public_validation_errors(exc.errors()) + public_exc: Final = RequestValidationError(public_errors).with_traceback(exc.__traceback__) if request.url.path.startswith(MANAGEMENT_V1_PREFIX): - validation_errors: Final[Sequence[ValidationErrorDetail]] = exc.errors() - problem: Final = request_validation_problem(validation_errors) - _close_dangling_otel_server_span(request, problem.status, exc=exc) + problem: Final = request_validation_problem(public_errors) + _close_dangling_otel_server_span(request, problem.status, exc=public_exc) return problem_response(problem) - _close_dangling_otel_server_span(request, 422, exc=exc) - return JSONResponse( - status_code=422, - content={"detail": jsonable_encoder(exc.errors())}, - ) + _close_dangling_otel_server_span(request, 422, exc=public_exc) + return JSONResponse(status_code=422, content={"detail": public_errors}) @app.exception_handler(Exception) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index c91b1afd64a..227f0e7f795 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -25,6 +25,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys from litellm.proxy._experimental.mcp_server.tool_search import MCP_TOOL_SEARCH_SETTINGS_KEY from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.validation_error_body import public_validation_errors from litellm.proxy.config_resolvers import FieldSource, SettingsStore, source_for from litellm.proxy.config_resolvers.settings_store import ConfigOwnedKeyError from litellm.proxy.config_resolvers.sso import ( @@ -1875,7 +1876,7 @@ async def update_ui_settings( try: settings: Final = effective_cls.model_validate(settings_body) except ValidationError as e: - raise HTTPException(status_code=422, detail=e.errors()) + raise HTTPException(status_code=422, detail=public_validation_errors(e.errors())) unsupported_team_fields: Final = sorted( frozenset(settings.team_admin_editable_team_fields) - SUPPORTED_TEAM_ADMIN_PERMISSIONS diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index fdc702af005..1ef39775bd3 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -10,7 +10,7 @@ from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload, runtime_checkable +from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Protocol, overload, runtime_checkable import httpx from openai._streaming import SSEDecoder @@ -265,6 +265,9 @@ def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool: return not isinstance(status_code, int) or status_code >= 500 or status_code == 429 +_PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"}) + + class BaseResponsesAPIStreamingIterator: """ Base class for streaming iterators that process responses from the Responses API. @@ -292,6 +295,7 @@ class BaseResponsesAPIStreamingIterator: self.start_time = getattr(logging_obj, "start_time", datetime.now()) self._failure_handled = False # Track if failure handler has been called self._yielded_first_chunk = False + self._output_started = False self._generated_content = "" self._generated_tool_arguments = "" self._completed_response_cached = False @@ -879,6 +883,46 @@ class BaseResponsesAPIStreamingIterator: except Exception: pass + def _note_yielded_event(self, event: ResponsesAPIStreamingResponse) -> None: + self._yielded_first_chunk = True + if event.type not in _PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: + self._output_started = True + + def _fallback_error(self, original: Exception) -> MidStreamFallbackError: + return MidStreamFallbackError( + message=str(original), + model=self.model or "", + llm_provider=self.custom_llm_provider or "", + original_exception=original, + generated_content="", + is_pre_first_chunk=not self._yielded_first_chunk, + ) + + def _stream_ended_early_error(self) -> litellm.APIConnectionError: + return litellm.APIConnectionError( + message=( + f"{self.custom_llm_provider or 'provider'} closed the responses stream before any terminal event " + "(response.completed, response.incomplete or response.failed)" + ), + llm_provider=self.custom_llm_provider or "", + model=self.model or "", + ) + + def _raise_if_ended_without_terminal_event(self) -> None: + if self.completed_response is not None: + return + error: Final = self._stream_ended_early_error() + self._handle_failure(error) + if self._output_started: + raise error + raise self._fallback_error(error) from error + + def _raise_for_transport_error(self, error: httpx.ReadError | httpx.RemoteProtocolError) -> NoReturn: + self._handle_failure(error) + if self._output_started: + raise error + raise self._fallback_error(error) from error + async def call_post_streaming_hooks_for_testing( iterator: object, chunk: ResponsesAPIStreamingResponse @@ -934,12 +978,14 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): sse = await self.stream_iterator.__anext__() except StopAsyncIteration: self.finished = True + self._raise_if_ended_without_terminal_event() raise StopAsyncIteration self._check_max_streaming_duration() result = self._process_chunk(sse.data) if self.finished: + self._raise_if_ended_without_terminal_event() raise StopAsyncIteration elif result is not None: self._maybe_raise_for_error_event(result) @@ -948,7 +994,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): result = await self._call_post_streaming_deployment_hook( chunk=result, ) - self._yielded_first_chunk = True + self._note_yielded_event(result) return result # If result is None, continue the loop to get the next chunk @@ -957,10 +1003,9 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise except (httpx.ReadError, httpx.RemoteProtocolError) as e: self.finished = True - if self.completed_response is None: - self._handle_failure(e) - raise - raise StopAsyncIteration from e + if self.completed_response is not None: + raise StopAsyncIteration from e + self._raise_for_transport_error(e) except httpx.HTTPError as e: # Handle HTTP errors self.finished = True @@ -1016,12 +1061,14 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): sse = next(self.stream_iterator) except StopIteration: self.finished = True + self._raise_if_ended_without_terminal_event() raise StopIteration self._check_max_streaming_duration() result = self._process_chunk(sse.data) if self.finished: + self._raise_if_ended_without_terminal_event() raise StopIteration elif result is not None: self._maybe_raise_for_error_event(result) @@ -1030,7 +1077,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): async_function=self._call_post_streaming_deployment_hook, chunk=result, ) - self._yielded_first_chunk = True + self._note_yielded_event(result) return result # If result is None, continue the loop to get the next chunk @@ -1039,10 +1086,9 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise except (httpx.ReadError, httpx.RemoteProtocolError) as e: self.finished = True - if self.completed_response is None: - self._handle_failure(e) - raise - raise StopIteration from e + if self.completed_response is not None: + raise StopIteration from e + self._raise_for_transport_error(e) except httpx.HTTPError as e: # Handle HTTP errors self.finished = True diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 4e84bded9de..64252cbbfb3 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -502,12 +502,12 @@ class RouterBudgetLimiting(CustomLogger): response_cost: Final[float] = standard_logging_payload.get("response_cost", 0) model_id: Final[str] = str(standard_logging_payload.get("model_id", "")) - custom_llm_provider: Final[str] = kwargs.get("litellm_params", {}).get("custom_llm_provider", None) - if custom_llm_provider is None: - raise ValueError("custom_llm_provider is required") + custom_llm_provider: Final[str | None] = standard_logging_payload.get("custom_llm_provider") - budget_config: Final = self._get_budget_config_for_provider(custom_llm_provider) - if budget_config: + budget_config: Final = ( + self._get_budget_config_for_provider(custom_llm_provider) if custom_llm_provider is not None else None + ) + if custom_llm_provider is not None and budget_config is not None: # increment spend for provider spend_key: Final = f"provider_spend:{custom_llm_provider}:{budget_config.budget_duration}" start_time_key: Final = f"provider_budget_start_time:{custom_llm_provider}" diff --git a/litellm/router_strategy/complexity_router/capability_classifier.py b/litellm/router_strategy/complexity_router/capability_classifier.py index 21046ff3421..93077af9e47 100644 --- a/litellm/router_strategy/complexity_router/capability_classifier.py +++ b/litellm/router_strategy/complexity_router/capability_classifier.py @@ -202,15 +202,26 @@ def capability_classifier_system_prompt(mode: Literal["json_schema", "json_objec ) -def unwrap_classifier_json(content: str) -> str: - """Remove the optional Markdown fence without repairing or weakening verdict JSON.""" - text: Final = content.strip() - if not text.startswith("```"): - return text - unfenced: Final = text.removeprefix("```").removeprefix("json").lstrip("\n\r") - return unfenced.removesuffix("```").strip() +_JSON_DECODER: Final = json.JSONDecoder() + + +def _complete_json_object_at(content: str, start: int) -> str | None: + try: + _, end = _JSON_DECODER.raw_decode(content, start) + except (ValueError, RecursionError): + return None + return content[start:end] + + +def extract_classifier_json(content: str) -> str: + """Return the first complete JSON object in the reply, whatever prose or fence surrounds it. + + A reply with no complete object comes back stripped so the caller's validation names the defect.""" + object_starts: Final = (index for index, char in enumerate(content) if char == "{") + candidates: Final = (_complete_json_object_at(content, start) for start in object_starts) + return next((candidate for candidate in candidates if candidate is not None), content.strip()) def parse_capability_classifier_verdict(content: str) -> CapabilityClassifierVerdict: - """Parse raw JSON or the fenced JSON shape tolerated by Switchyard.""" - return CapabilityClassifierVerdict.model_validate_json(unwrap_classifier_json(content)) + """Parse the verdict object out of a bare, fenced, or prose-wrapped reply.""" + return CapabilityClassifierVerdict.model_validate_json(extract_classifier_json(content)) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 0f252952a9d..9df6306436b 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -29,6 +29,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast from pydantic import BaseModel, TypeAdapter, ValidationError, create_model +from pydantic_core import ErrorDetails from litellm._logging import verbose_router_logger from litellm.caching.affinity_cache import claim_affinity_pin @@ -85,8 +86,8 @@ from .capability_classifier import ( CapabilityClassifierForecast, capability_classifier_response_format, capability_classifier_system_prompt, + extract_classifier_json, parse_capability_classifier_verdict, - unwrap_classifier_json, ) from .classification_rubrics import BUSINESS_TIER_CRITERIA, calibration_examples_section from .config import ( @@ -427,6 +428,41 @@ def _effective_turn_off_message_logging(request_kwargs: Mapping[str, object] | N ) +def _classifier_reply_is_private(request_kwargs: Mapping[str, object] | None) -> bool: + from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( + initialize_standard_callback_dynamic_params, + ) + from litellm.litellm_core_utils.redact_messages import should_redact_message_logging + + kwargs: Final = dict(request_kwargs) if request_kwargs else {} + try: + return should_redact_message_logging( + { + "litellm_params": kwargs, + "standard_callback_dynamic_params": initialize_standard_callback_dynamic_params(kwargs), + } + ) + except AttributeError: + return True + + +def _validation_problem(detail: ErrorDetails) -> str: + location: Final = ".".join(str(part) for part in detail["loc"]) + return f"{location}: {detail['msg']}" if location else detail["msg"] + + +def _log_rejected_classifier_verdict( + error: ValidationError, content: str, request_kwargs: Mapping[str, object] | None +) -> None: + problems: Final = "; ".join(_validation_problem(detail) for detail in error.errors()) + reply: Final = ( + "raw reply withheld (message logging is off)" + if _classifier_reply_is_private(request_kwargs) + else f"raw reply: {content!r}" + ) + verbose_router_logger.warning("ComplexityRouter: classifier verdict rejected (%s); %s", problems, reply) + + _REMINDER_OPEN: Final = "" _REMINDER_CLOSE: Final = "" _DEFAULT_REMINDER_MARKERS: Final = ((_REMINDER_OPEN, _REMINDER_CLOSE),) @@ -2040,7 +2076,7 @@ class ComplexityRouter(CustomLogger): except Exception as e: # noqa: BLE001 -- every unavailable or invalid judge verdict must fail closed if breaker is not None and permit is not None: breaker.record_failure(permit, is_timeout=_is_classifier_timeout(e)) - return self._capability_classifier_failure_outcome(f"capability classifier failed ({e})") + return self._capability_classifier_failure_outcome(f"capability classifier failed ({type(e).__name__})") def _capability_classifier_failure_outcome(self, reason: str, signal: str | None = None) -> ClassificationOutcome: """Fail closed to the configured capable tier without consulting another taxonomy.""" @@ -2449,7 +2485,11 @@ class ComplexityRouter(CustomLogger): content, classifier_cost = await self._call_classifier_model( messages_for_call, request_kwargs, encrypted_task=encrypted_task ) - raw_tier: Final = _LabeledTierClassification.model_validate_json(content).tier + try: + raw_tier: Final = _LabeledTierClassification.model_validate_json(extract_classifier_json(content)).tier + except ValidationError as error: + _log_rejected_classifier_verdict(error, content, request_kwargs) + raise tier: Final = self.config.resolve_classified_tier(raw_tier) if tier is None: raise ValueError(f"LLM classifier returned an unrecognized tier: {raw_tier!r}") @@ -2508,7 +2548,11 @@ class ComplexityRouter(CustomLogger): max_output_tokens=capability.max_output_tokens, encrypted_task=encrypted_task, ) - verdict: Final = parse_capability_classifier_verdict(content) + try: + verdict: Final = parse_capability_classifier_verdict(content) + except ValidationError as error: + _log_rejected_classifier_verdict(error, content, request_kwargs) + raise threshold: Final = verdict.routing_threshold(capability.base_threshold, capability.threshold_step) calibration: Final = capability.calibration forecast: Final = CapabilityClassifierForecast( @@ -2563,8 +2607,9 @@ class ComplexityRouter(CustomLogger): messages_for_call, request_kwargs, encrypted_task=encrypted, max_output_tokens=v2.max_output_tokens ) try: - verdict: Final = LLMV2Verdict.model_validate_json(unwrap_classifier_json(content)) - except ValidationError: + verdict: Final = LLMV2Verdict.model_validate_json(extract_classifier_json(content)) + except ValidationError as error: + _log_rejected_classifier_verdict(error, content, request_kwargs) return self._classifier_failure_outcome("Invalid LLM V2 forecast", prompt, system_prompt)._replace( classifier_cost=classifier_cost ) diff --git a/litellm/router_strategy/complexity_router/llm_v2.py b/litellm/router_strategy/complexity_router/llm_v2.py index 18351237e65..8ef2f554ab2 100644 --- a/litellm/router_strategy/complexity_router/llm_v2.py +++ b/litellm/router_strategy/complexity_router/llm_v2.py @@ -16,6 +16,7 @@ from litellm.llms.base_llm.base_utils import ( from litellm.router_strategy.complexity_router.fuse_presets import ProfileText, resolve_fuse_profile ShortText: TypeAlias = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1, max_length=512)] +VerdictText: TypeAlias = Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)] class _SolverProfile(TypedDict): @@ -90,7 +91,7 @@ class LLMV2Demands(BaseModel): class LLMV2SolverForecast(BaseModel): model_config = ConfigDict(extra="forbid", frozen=True) - likely_failure: ShortText + likely_failure: VerdictText p_solve: StrictFloat = Field(ge=0.0, le=1.0) @@ -104,7 +105,7 @@ class LLMV2SolverForecasts(BaseModel): class LLMV2Verdict(BaseModel): model_config = ConfigDict(extra="forbid", frozen=True) - crux: ShortText + crux: VerdictText demands: LLMV2Demands verification: Literal["relevant", "partial", "unavailable", "unknown"] forecasts: LLMV2SolverForecasts diff --git a/litellm/types/integrations/zerobus.py b/litellm/types/integrations/zerobus.py new file mode 100644 index 00000000000..217002dbc65 --- /dev/null +++ b/litellm/types/integrations/zerobus.py @@ -0,0 +1,53 @@ +from dataclasses import dataclass, field +from typing import Final + +from pydantic import Field + +from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams + +RETRYABLE_INGEST_STATUS_CODES: Final = frozenset({408, 429, 500, 502, 503, 504}) + +TOKEN_REFRESH_LEEWAY_SECONDS: Final = 60 + + +class ZerobusInitParams(StandardCustomLoggerInitParams): + """ + Params for initializing a Databricks Zerobus logger on litellm. + + Every connection field falls back to its ``ZEROBUS_*`` environment variable, which is + what the proxy UI writes. ``table_name`` is the fully qualified ``catalog.schema.table``. + """ + + workspace_url: str | None = None + server_endpoint: str | None = None + client_id: str | None = None + client_secret: str | None = None + table_name: str | None = None + batch_size: int = Field(default=100, gt=0) + flush_interval: int = Field(default=10, gt=0) + + +@dataclass(frozen=True, slots=True) +class ZerobusConnection: + """Everything needed to mint a token for one table and post rows to it.""" + + workspace_url: str + workspace_id: str + server_endpoint: str + client_id: str + client_secret: str = field(repr=False) + table_name: str + + +@dataclass(frozen=True, slots=True) +class ZerobusAccessToken: + value: str = field(repr=False) + expires_at: float + + +@dataclass(frozen=True, slots=True) +class ZerobusIngestFailure: + """Why a batch could not be written, and whether a later attempt could still succeed.""" + + detail: str + retryable: bool diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 6e7e9da3498..99ab5920c4f 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -522,7 +522,19 @@ class CreateBatchRequest(TypedDict, total=False): """ completion_window: Literal["24h"] - endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"] + endpoint: Literal[ + "/v1/chat/completions", + "/v1/embeddings", + "/v1/completions", + "/v1/responses", + "/v1/ocr", + "/v1/images/generations", + "/v1/images/edits", + "/v1/videos/generations", + "/v1/videos", + "/v1/videos/edits", + "/v1/videos/extensions", + ] input_file_id: str metadata: dict[str, str] | None output_expires_after: FileExpiresAfter diff --git a/litellm/types/router.py b/litellm/types/router.py index 0d89c9b5081..6362cea4eec 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -345,6 +345,7 @@ class CredentialLiteLLMParams(BaseModel): ## OBJECT STORAGE (files / batches) ## gcs_bucket_name: str | None = None + bucket_name: str | None = None ## AWS BEDROCK / SAGEMAKER ## aws_access_key_id: str | None = None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index bb2899fb0a6..857e161599f 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -299,6 +299,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): cache_read_input_token_cost_above_272k_tokens_flex: float | None cache_read_input_token_cost_above_512k_tokens: float | None cache_read_input_token_cost_batches: ReadOnly[float | None] + cache_read_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None] cache_read_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None] cache_creation_input_token_cost_batches: ReadOnly[float | None] cache_creation_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None] @@ -327,8 +328,10 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): input_cost_per_second: float | None # for OpenAI Speech models input_cost_per_token_batches: float | None input_cost_per_video_token_batches: ReadOnly[float | None] + input_cost_per_token_above_200k_tokens_batches: ReadOnly[float | None] input_cost_per_token_above_272k_tokens_batches: ReadOnly[float | None] output_cost_per_token_batches: float | None + output_cost_per_token_above_200k_tokens_batches: ReadOnly[float | None] output_cost_per_token_above_272k_tokens_batches: ReadOnly[float | None] output_cost_per_token: Required[float | None] output_cost_per_token_flex: float | None # OpenAI flex service tier pricing @@ -3731,6 +3734,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): cache_read_input_token_cost_above_272k_tokens_priority: float | None = None cache_read_input_token_cost_above_272k_tokens_flex: float | None = None cache_read_input_token_cost_batches: float | None = None + cache_read_input_token_cost_above_200k_tokens_batches: float | None = None cache_read_input_token_cost_above_272k_tokens_batches: float | None = None cache_creation_input_token_cost_batches: float | None = None cache_creation_input_token_cost_above_272k_tokens_batches: float | None = None @@ -3744,6 +3748,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): input_cost_per_token_above_200k_tokens_priority: float | None = None input_cost_per_token_above_272k_tokens_priority: float | None = None input_cost_per_token_above_272k_tokens_flex: float | None = None + input_cost_per_token_above_200k_tokens_batches: float | None = None input_cost_per_token_above_272k_tokens_batches: float | None = None input_cost_per_query: float | None = None input_cost_per_image: float | None = None @@ -3768,6 +3773,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): output_cost_per_token_above_200k_tokens_priority: float | None = None output_cost_per_token_above_272k_tokens_priority: float | None = None output_cost_per_token_above_272k_tokens_flex: float | None = None + output_cost_per_token_above_200k_tokens_batches: float | None = None output_cost_per_token_above_272k_tokens_batches: float | None = None output_cost_per_character_above_128k_tokens: float | None = None output_cost_per_image: float | None = None @@ -4142,7 +4148,7 @@ FILE_CONTENT_STREAMING_PROVIDERS: Final[frozenset[str]] = frozenset( LITELLM_EXECUTED_BATCH_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.HOSTED_VLLM.value}) -ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai"] +ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai", "xai"] LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider)) diff --git a/litellm/utils.py b/litellm/utils.py index 4551ffcedc6..424d4f367f5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6160,6 +6160,9 @@ def _get_model_info_helper( cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None), cache_read_input_token_cost_ultrafast=_model_info.get("cache_read_input_token_cost_ultrafast", None), cache_read_input_token_cost_batches=_model_info.get("cache_read_input_token_cost_batches"), + cache_read_input_token_cost_above_200k_tokens_batches=_model_info.get( + "cache_read_input_token_cost_above_200k_tokens_batches" + ), cache_read_input_token_cost_above_272k_tokens_batches=_model_info.get( "cache_read_input_token_cost_above_272k_tokens_batches" ), @@ -6197,10 +6200,16 @@ def _get_model_info_helper( input_cost_per_video_per_second=_model_info.get("input_cost_per_video_per_second", None), input_cost_per_token_batches=_model_info.get("input_cost_per_token_batches"), input_cost_per_video_token_batches=_model_info.get("input_cost_per_video_token_batches", None), + input_cost_per_token_above_200k_tokens_batches=_model_info.get( + "input_cost_per_token_above_200k_tokens_batches" + ), input_cost_per_token_above_272k_tokens_batches=_model_info.get( "input_cost_per_token_above_272k_tokens_batches" ), output_cost_per_token_batches=_model_info.get("output_cost_per_token_batches"), + output_cost_per_token_above_200k_tokens_batches=_model_info.get( + "output_cost_per_token_above_200k_tokens_batches" + ), output_cost_per_token_above_272k_tokens_batches=_model_info.get( "output_cost_per_token_above_272k_tokens_batches" ), @@ -9358,6 +9367,10 @@ class ProviderConfigManager: from litellm.llms.mistral.files.transformation import MistralFilesConfig return MistralFilesConfig() + elif LlmProviders.XAI == provider: + from litellm.llms.xai.files.transformation import XAIFilesConfig + + return XAIFilesConfig() return None @staticmethod diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c252fe68a7b..c9171872734 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -41488,65 +41488,63 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 8.44944e-07, + "cache_read_input_token_cost": 3.828e-08, + "input_cost_per_token": 4.5936e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.689888e-06, + "output_cost_per_token": 9.1872e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.0412e-08, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "input_cost_per_token": 3e-07, - "output_cost_per_token": 1.2e-06, - "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 1.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":1.5e-7,"output_cost_per_token":6e-7,"cache_read_input_token_cost":3e-9}, + "output_cost_per_token": 4.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_vision": true, "supports_pdf_input": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 4.62e-07, + "cache_read_input_token_cost": 8.8e-09, + "input_cost_per_token": 2.64e-07, "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.386e-06, + "output_cost_per_token": 7.92e-07, "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, "supports_function_calling": true, + "supports_pdf_input": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 1.54e-08, - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"00:00-01:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"04:00-06:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"hours_utc":"10:00-00:00"}],"input_cost_per_token":0.00000132,"output_cost_per_token":0.00000396,"cache_read_input_token_cost":4.4e-8}, - "supports_audio_input": false, - "supports_pdf_input": false, "supports_vision": false, "supports_web_search": false }, @@ -42766,14 +42764,14 @@ "openrouter/qwen/qwen3-coder-plus": { "cache_creation_input_token_cost": 8.125e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, - "input_cost_per_token_above_32k_tokens": 1.17e-06, "cache_creation_input_token_cost_above_32k_tokens": 1.4625e-06, - "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, - "output_cost_per_token_above_32k_tokens": 5.85e-06, "cache_read_input_token_cost": 1.3e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 2.34e-07, + "deprecation_date": "2026-10-09", "input_cost_per_token": 6.5e-07, "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.17e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 65536, @@ -42781,6 +42779,7 @@ "mode": "chat", "output_cost_per_token": 3.25e-06, "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 5.85e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42813,6 +42812,7 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.3e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, @@ -43110,25 +43110,25 @@ "supports_web_search": false }, "openrouter/z-ai/glm-4.7": { - "input_cost_per_token": 4e-07, - "output_cost_per_token": 1.75e-06, "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 8e-08, + "cache_read_input_token_cost": 1.1e-07, + "input_cost_per_token": 6e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", + "output_cost_per_token": 2.2e-06, "source": "https://openrouter.ai/api/v1/models", - "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_vision": false, - "supports_prompt_caching": true, "supports_assistant_prefill": true, "supports_audio_input": false, + "supports_function_calling": true, "supports_pdf_input": false, + "supports_prompt_caching": true, + "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/z-ai/glm-4.7-flash": { @@ -43173,15 +43173,15 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.1": { - "input_cost_per_token": 9.66e-07, - "output_cost_per_token": 3.036e-06, - "cache_read_input_token_cost": 1.794e-07, "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 1.7914e-07, + "input_cost_per_token": 9.646e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", + "output_cost_per_token": 3.0316e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -51447,13 +51447,16 @@ }, "xai/grok-4.20-0309-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -51461,8 +51464,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true @@ -51491,9 +51497,13 @@ "xai/grok-4.3": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_image_token": 1.25e-06, "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, @@ -51501,6 +51511,8 @@ "mode": "chat", "output_cost_per_token": 2.5e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -51513,9 +51525,13 @@ "xai/grok-4.3-latest": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_image_token": 1.25e-06, "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, @@ -51523,6 +51539,8 @@ "mode": "chat", "output_cost_per_token": 2.5e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -59494,13 +59512,16 @@ }, "xai/grok-4.20-0309-non-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -59508,20 +59529,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-multi-agent-0309": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "responses", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": false, "supports_prompt_caching": true, @@ -59530,8 +59557,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true, "supported_endpoints": [ @@ -60224,18 +60254,18 @@ "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/deepseek-v4p1-flash": { - "cache_read_input_token_cost": 7e-09, - "cache_read_input_token_cost_priority": 8.75e-09, - "input_cost_per_token": 2.2e-07, - "input_cost_per_token_priority": 2.75e-07, + "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost_priority": 7.5e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 3.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 6.6e-07, - "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.5e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60245,13 +60275,16 @@ }, "fireworks_ai/accounts/fireworks/routers/deepseek-v4p1-flash-us": { "cache_read_input_token_cost": 9e-09, + "cache_read_input_token_cost_priority": 1.125e-08, "input_cost_per_token": 4.5e-07, + "input_cost_per_token_priority": 5.625e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 1.8e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, @@ -60262,13 +60295,16 @@ }, "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60320,18 +60356,18 @@ "supports_vision": false }, "fireworks_ai/deepseek-v4p1-flash": { - "cache_read_input_token_cost": 7e-09, - "cache_read_input_token_cost_priority": 8.75e-09, - "input_cost_per_token": 2.2e-07, - "input_cost_per_token_priority": 2.75e-07, + "cache_read_input_token_cost": 6e-09, + "cache_read_input_token_cost_priority": 7.5e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_priority": 3.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 6.6e-07, - "output_cost_per_token_priority": 8.25e-07, - "source": "https://api.fireworks.ai/v1/serverless/models", + "output_cost_per_token": 1.2e-06, + "output_cost_per_token_priority": 1.5e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -60341,13 +60377,16 @@ }, "fireworks_ai/deepseek-v4p1-flash-us": { "cache_read_input_token_cost": 9e-09, + "cache_read_input_token_cost_priority": 1.125e-08, "input_cost_per_token": 4.5e-07, + "input_cost_per_token_priority": 5.625e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", "output_cost_per_token": 1.8e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_prompt_caching": true, @@ -60358,13 +60397,16 @@ }, "fireworks_ai/deepseek-v4-flash-vision-exp": { "cache_read_input_token_cost": 7e-09, + "cache_read_input_token_cost_priority": 8.75e-09, "deprecation_date": "2026-09-25", "input_cost_per_token": 2.2e-07, + "input_cost_per_token_priority": 2.75e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 1048576, "max_tokens": 1048576, "mode": "chat", "output_cost_per_token": 6.6e-07, + "output_cost_per_token_priority": 8.25e-07, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_tool_choice": true, @@ -60493,14 +60535,17 @@ }, "fireworks_ai/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -60542,14 +60587,17 @@ }, "fireworks_ai/accounts/fireworks/models/muse-glimmer-30b": { "cache_read_input_token_cost": 4e-08, + "cache_read_input_token_cost_priority": 6e-08, "deprecation_date": "2026-09-25", "input_cost_per_token": 3.5e-07, + "input_cost_per_token_priority": 5.25e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", "output_cost_per_token": 1.5e-06, + "output_cost_per_token_priority": 2.25e-06, "source": "https://api.fireworks.ai/v1/serverless/models", "supports_function_calling": true, "supports_reasoning": true, @@ -62798,13 +62846,16 @@ }, "xai/grok-4.20": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -62812,21 +62863,27 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true }, "xai/grok-4.20-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -62834,21 +62891,27 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true }, "xai/grok-4.20-reasoning-latest": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -62856,8 +62919,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true @@ -63078,13 +63144,16 @@ }, "xai/grok-4.20-non-reasoning": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -63092,20 +63161,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-non-reasoning-latest": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_prompt_caching": true, @@ -63113,20 +63188,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-multi-agent": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "responses", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supported_endpoints": [ "/v1/responses" @@ -63138,20 +63219,26 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, "xai/grok-4.20-multi-agent-latest": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "responses", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supported_endpoints": [ "/v1/responses" @@ -63163,8 +63250,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_response_schema": true }, @@ -66105,23 +66195,23 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3-flash": { + "cache_read_input_token_cost": 1e-08, "input_cost_per_token": 4.5e-08, - "output_cost_per_token": 6e-07, - "cache_read_input_token_cost": 2.85e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 1.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { @@ -66282,24 +66372,24 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "input_cost_per_token": 3e-08, - "output_cost_per_token": 3.2e-07, "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 2.2e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", + "output_cost_per_token": 3.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, - "supports_tool_choice": true, - "supports_reasoning": true, - "supports_response_schema": true, "supports_parallel_function_calling": true, "supports_pdf_input": false, - "supports_vision": false, "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.7-flash": { @@ -66772,46 +66862,47 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.6-max-preview": { - "input_cost_per_token": 1.027e-06, - "output_cost_per_token": 6.162e-06, "cache_creation_input_token_cost": 1.28375e-06, - "input_cost_per_token_above_128k_tokens": 1.58e-06, - "output_cost_per_token_above_128k_tokens": 9.48e-06, "cache_creation_input_token_cost_above_128k_tokens": 1.975e-06, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 1.027e-06, + "input_cost_per_token_above_128k_tokens": 1.58e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 6.162e-06, + "output_cost_per_token_above_128k_tokens": 9.48e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, "openrouter/qwen/qwen3.6-27b": { - "input_cost_per_token": 3.2e-07, - "output_cost_per_token": 2.7e-06, "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 3.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 262140, "max_tokens": 262140, "mode": "chat", + "output_cost_per_token": 3.2e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/openai/gpt-5.5-pro": { @@ -66856,23 +66947,23 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 4.9e-08, - "output_cost_per_token": 9.8e-08, - "cache_read_input_token_cost": 9.8e-09, + "cache_read_input_token_cost": 9.408e-09, + "input_cost_per_token": 4.704e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", + "output_cost_per_token": 9.408e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, - "supports_prompt_caching": true, "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { @@ -66897,22 +66988,22 @@ "supports_web_search": false }, "openrouter/google/gemma-4-26b-a4b-it": { - "cache_read_input_token_cost": 5e-08, - "input_cost_per_token": 9e-08, - "output_cost_per_token": 3e-07, + "cache_read_input_token_cost": 3.75e-08, + "input_cost_per_token": 6.75e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", + "output_cost_per_token": 2.25e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": true, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67197,25 +67288,26 @@ "supports_video_input": true }, "openrouter/qwen/qwen3-max-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67467,59 +67559,62 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-vl-32b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.04e-07, - "output_cost_per_token": 4.16e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.16e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.8e-07, - "output_cost_per_token": 2.1e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.1e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-8b-instruct": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67549,40 +67644,41 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-30b-a3b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, "openrouter/qwen/qwen3-vl-30b-a3b-instruct": { - "input_cost_per_token": 1.3e-07, - "output_cost_per_token": 5.2e-07, + "input_cost_per_token": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67606,21 +67702,22 @@ "supports_web_search": true }, "openrouter/qwen/qwen3-vl-235b-a22b-thinking": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4e-07, - "output_cost_per_token": 4e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": true, "supports_web_search": false }, @@ -67645,32 +67742,33 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-max": { - "input_cost_per_token": 7.8e-07, - "input_cost_per_token_above_32k_tokens": 1.56e-06, - "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, - "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, - "output_cost_per_token_above_32k_tokens": 7.8e-06, - "output_cost_per_token": 3.9e-06, - "cache_read_input_token_cost": 1.56e-07, "cache_creation_input_token_cost": 9.75e-07, - "input_cost_per_token_above_128k_tokens": 1.95e-06, - "output_cost_per_token_above_128k_tokens": 9.75e-06, - "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, "cache_creation_input_token_cost_above_128k_tokens": 2.4375e-06, + "cache_creation_input_token_cost_above_32k_tokens": 1.95e-06, + "cache_read_input_token_cost": 1.56e-07, + "cache_read_input_token_cost_above_128k_tokens": 3.9e-07, + "cache_read_input_token_cost_above_32k_tokens": 3.12e-07, + "deprecation_date": "2026-10-09", + "input_cost_per_token": 7.8e-07, + "input_cost_per_token_above_128k_tokens": 1.95e-06, + "input_cost_per_token_above_32k_tokens": 1.56e-06, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", + "output_cost_per_token": 3.9e-06, + "output_cost_per_token_above_128k_tokens": 9.75e-06, + "output_cost_per_token_above_32k_tokens": 7.8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, - "supports_tool_choice": true, - "supports_response_schema": true, - "supports_vision": false, "supports_prompt_caching": true, "supports_reasoning": false, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false, "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.1-terminus": { @@ -67765,23 +67863,24 @@ "openrouter/qwen/qwen-plus-2025-07-28": { "cache_creation_input_token_cost": 3.25e-07, "cache_read_input_token_cost": 5.2e-08, + "deprecation_date": "2026-10-09", "input_cost_per_token": 2.6e-07, - "output_cost_per_token": 7.8e-07, "input_cost_per_token_above_256k_tokens": 7.8e-07, - "output_cost_per_token_above_256k_tokens": 2.34e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 7.8e-07, + "output_cost_per_token_above_256k_tokens": 2.34e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, "supports_reasoning": false, - "supports_tool_choice": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -67805,21 +67904,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-30b-a3b-thinking-2507": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.4e-06, "litellm_provider": "openrouter", "max_input_tokens": 81920, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68128,21 +68228,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-8b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 1.17e-07, - "output_cost_per_token": 4.55e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 4.55e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -68185,21 +68286,22 @@ "supports_web_search": false }, "openrouter/qwen/qwen3-235b-a22b": { + "deprecation_date": "2026-10-09", "input_cost_per_token": 4.55e-07, - "output_cost_per_token": 1.82e-06, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", + "output_cost_per_token": 1.82e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, "supports_pdf_input": false, "supports_prompt_caching": false, - "supports_tool_choice": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false }, @@ -71268,6 +71370,23 @@ "output_cost_per_token": 0.0, "source": "https://openrouter.ai/typesafe/jev-1.13" }, + "openrouter/typesafe/jev-router": { + "input_cost_per_token": 0, + "output_cost_per_token": 0, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "source": "https://openrouter.ai/typesafe/jev-router", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_audio_input": true, + "supports_video_input": true + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", @@ -72952,14 +73071,14 @@ "supports_web_search": false }, "openrouter/inclusionai/ling-3.0-flash-vl": { - "cache_read_input_token_cost": 1.2e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 2.1e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.8e-07, + "output_cost_per_token": 6.16e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -75470,13 +75589,16 @@ }, "xai/grok-4.20-0309": { "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1.6e-07, "input_cost_per_token": 1.25e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "xai", "max_input_tokens": 1000000, "max_output_tokens": 1000000, "max_tokens": 1000000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 2e-06, "source": "https://api.x.ai/v1/language-models", "supports_function_calling": true, "supports_reasoning": true, @@ -75484,8 +75606,11 @@ "supports_vision": true, "supports_web_search": true, "input_cost_per_token_above_200k_tokens": 2.5e-06, + "input_cost_per_token_above_200k_tokens_batches": 2e-06, "output_cost_per_token_above_200k_tokens": 5e-06, + "output_cost_per_token_above_200k_tokens_batches": 4e-06, "cache_read_input_token_cost_above_200k_tokens": 4e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07, "input_cost_per_image_token": 1.25e-06, "supports_prompt_caching": true, "supports_response_schema": true @@ -77089,5 +77214,24 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_function_calling": true, "supports_tool_choice": false + }, + "openrouter/perceptron/perceptron-mk1.5": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 36864, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": false, + "supports_function_calling": false, + "supports_pdf_input": false, + "supports_prompt_caching": false, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": false, + "supports_vision": true, + "supports_web_search": false } } diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index a05a6e56514..0fbf990b671 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -170,6 +170,11 @@ "minimum": 0, "description": "Rate applied once the prompt exceeds the token threshold in the field name." }, + "cache_read_input_token_cost_above_200k_tokens_batches": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_read_input_token_cost_above_200k_tokens_priority": { "type": "number", "minimum": 0, @@ -355,6 +360,11 @@ "minimum": 0, "description": "Rate applied once the prompt exceeds the token threshold in the field name." }, + "input_cost_per_token_above_200k_tokens_batches": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "input_cost_per_token_above_200k_tokens_priority": { "type": "number", "minimum": 0, @@ -712,6 +722,11 @@ "minimum": 0, "description": "Rate applied once the prompt exceeds the token threshold in the field name." }, + "output_cost_per_token_above_200k_tokens_batches": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "output_cost_per_token_above_200k_tokens_priority": { "type": "number", "minimum": 0, diff --git a/pyproject.toml b/pyproject.toml index f2364b5e77b..28b00379cc7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -249,6 +249,7 @@ proxy-dev = [ "prisma==0.11.0", "hypercorn==0.17.3", "prometheus-client==0.20.0", + "sentry-sdk==2.21.0", "opentelemetry-api==1.33.1", "opentelemetry-sdk==1.33.1", "opentelemetry-exporter-otlp==1.33.1", diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index dc1f8592612..659dc438f2d 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -72,6 +72,7 @@ IGNORE_FUNCTIONS = [ "_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible). "_replace_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible). "_sort_processed_sets", # bounded by the nesting depth of the log-record extra it walks (a finite JSON tree, no cycles possible). + "scrub_json_strings", # max depth set (MAX_SCRUB_DEPTH); fails closed by returning "[Filtered]" for anything nested past the cap. ] diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 20c87dbbd74..9cfe6e33ed6 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -103,6 +103,8 @@ - {id: llm.chat_completions.anthropic.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic over /chat/completions: cost header and spend row agree"} - {id: llm.chat_completions.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /chat/completions"} - {id: llm.messages.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /v1/messages"} +- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_event, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_event], source: "customer report", rationale: "An upstream that hangs up mid-stream must reach Anthropic clients as an event: error frame, not an OpenAI-shaped data-only error they silently drop"} +- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_status, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_status], source: "customer report", rationale: "An upstream that hangs up before its first byte must answer as a JSON error carrying its status, so Anthropic clients raise the status-specific error and retry on it instead of reading a 200 stream that only carries an error event"} - {id: llm.messages.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI models served on the Anthropic Messages contract"} - {id: llm.messages.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages streams the Anthropic event grammar"} - {id: llm.messages.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages: cost header and spend row agree"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index 5fd19212ab7..fec1934059c 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -86,6 +86,7 @@ LlmCapability = Literal[ "tool_search", "tool_search_history", "tool_use", + "upstream_stream_failure", "vision", "web_search", "web_search_server_tool", diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py index d048d1343eb..871fd2f9aef 100644 --- a/tests/e2e/llm_translation/test_messages_e2e.py +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -11,8 +11,11 @@ litellm-regression-tests/tests/test_inference_endpoints.py. from __future__ import annotations import time +from collections.abc import Callable +from types import MappingProxyType from typing import Final +import anthropic import pytest from anthropic import Anthropic from anthropic.types import ( @@ -30,12 +33,21 @@ from anthropic.types import ( ToolParam, ToolUseBlock, ) -from e2e_config import STREAM_MIN_LEAD_SECONDS, provider_edge_base, provider_paces_stream, unique_marker +from e2e_config import ( + PROVIDER_EDGE_ADVERTISE_HOST, + PROVIDER_EDGE_BIND_HOST, + STREAM_MIN_LEAD_SECONDS, + provider_edge_base, + provider_paces_stream, + unique_marker, +) from e2e_http import assert_client_error from lifecycle import ResourceManager -from models import ChatMessage, LiteLLMParamsBody, SpendLogRow +from models import AnthropicErrorEvent, AnthropicMessagesBody, ChatMessage, LiteLLMParamsBody, SpendLogRow +from provider_edge import EDGE_MOUNTS, LiveEdge, RunningEdge, StreamCut, start_provider_edge +from provider_edge_bedrock import bedrock_signer from proxy_client import ProxyClient -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError from sdk_clients import NO_PROXY_CACHE, SdkClients, response_header pytestmark = [pytest.mark.e2e, pytest.mark.replayable] @@ -385,3 +397,245 @@ class TestOpenAIMessagesToolContinuation: ) assert _text(continuation).strip() == receipt, "continuation did not consume the correlated tool result" assert all(not isinstance(block, ToolUseBlock) for block in continuation.content) + + +BEDROCK_BACKEND: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" +BEDROCK_EDGE_REGION: Final = "us-east-1" +_STREAM_FAILURE_PROMPT: Final = "Count from 1 to 100, one number per line." +_FRAME_PAYLOAD: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +_AT_FRAME_BOUNDARY: Final = StreamCut(after_content=True) +_MID_FRAME: Final = StreamCut(after_content=True, mid_chunk=True) +_BEFORE_FIRST_BYTE: Final = StreamCut(after_content=False) + +type _CutRegistration = Callable[[ProxyClient, ResourceManager, StreamCut], tuple[str, str]] + + +def _cut_edge(backend: LiveEdge, mount: str) -> RunningEdge: + return start_provider_edge( + backend, + mounts=MappingProxyType({mount: EDGE_MOUNTS[mount]}), + bind_host=PROVIDER_EDGE_BIND_HOST, + advertise_host=PROVIDER_EDGE_ADVERTISE_HOST, + ) + + +def _register_cut_bedrock(proxy: ProxyClient, resources: ResourceManager, cut: StreamCut) -> tuple[str, str]: + mount: Final = f"bedrock/{BEDROCK_EDGE_REGION}" + edge: Final = _cut_edge(LiveEdge(cut=cut, sign=bedrock_signer(BEDROCK_EDGE_REGION)), mount) + resources.defer(edge.shutdown) + return _register( + proxy, + resources, + LiteLLMParamsBody( + model=BEDROCK_BACKEND, + api_base=edge.edge.api_base(mount), + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name=BEDROCK_EDGE_REGION, + ), + prefix="e2e-messages-cut", + ) + + +def _register_cut_anthropic(proxy: ProxyClient, resources: ResourceManager, cut: StreamCut) -> tuple[str, str]: + edge: Final = _cut_edge(LiveEdge(cut=cut), "anthropic") + resources.defer(edge.shutdown) + return _register( + proxy, + resources, + LiteLLMParamsBody( + model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY", api_base=edge.edge.api_base("anthropic") + ), + prefix="e2e-messages-cut", + ) + + +_DROPPED_UPSTREAMS: Final[tuple[tuple[str, _CutRegistration, StreamCut], ...]] = ( + ("bedrock_at_a_frame_boundary", _register_cut_bedrock, _AT_FRAME_BOUNDARY), + ("anthropic_at_a_frame_boundary", _register_cut_anthropic, _AT_FRAME_BOUNDARY), + ("anthropic_mid_frame", _register_cut_anthropic, _MID_FRAME), +) +_DROPPED_BEFORE_FIRST_BYTE: Final[tuple[tuple[str, _CutRegistration, StreamCut], ...]] = ( + ("bedrock_before_the_first_byte", _register_cut_bedrock, _BEFORE_FIRST_BYTE), + ("anthropic_before_the_first_byte", _register_cut_anthropic, _BEFORE_FIRST_BYTE), +) + + +def _payload(frame: str) -> JsonValue | None: + try: + return _FRAME_PAYLOAD.validate_json(frame) + except ValidationError: + return None + + +def _bare_error_frame(frame: str) -> bool: + payload: Final = _payload(frame) + return isinstance(payload, dict) and "error" in payload and payload.get("type") != "error" + + +@pytest.mark.provider_edge_host +@pytest.mark.provider_live +class TestMessagesUpstreamStreamFailure: + @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event") + @pytest.mark.parametrize( + ("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS] + ) + def test_interrupted_upstream_stream_raises_in_the_anthropic_sdk( + self, + proxy: ProxyClient, + resources: ResourceManager, + sdk: SdkClients, + register: _CutRegistration, + cut: StreamCut, + ) -> None: + model, key = register(proxy, resources, cut) + client: Final = sdk.anthropic(key) + + stream: Final = client.messages.create( + model=model, + max_tokens=300, + stream=True, + messages=[_user_turn(_STREAM_FAILURE_PROMPT)], + extra_body=NO_PROXY_CACHE, + ) + first: Final = next(stream) + assert first.type == "message_start", ( + f"the stream produced a first event that is not message_start, so this run proves a " + f"startup failure, not an interrupted stream: {first!r}" + ) + with pytest.raises(anthropic.APIStatusError) as raised: + for _ in stream: + pass + try: + AnthropicErrorEvent.model_validate(raised.value.body) + except ValidationError: + pytest.fail( + f"the SDK raised on the interrupted stream but without the Anthropic error envelope a " + f"client reads the failure from: body={raised.value.body!r} message={raised.value}" + ) + + @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event") + @pytest.mark.parametrize( + ("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS] + ) + def test_interrupted_upstream_stream_is_an_anthropic_error_event( + self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut + ) -> None: + model, key = register(proxy, resources, cut) + + outcome: Final = proxy.messages_stream( + key, + AnthropicMessagesBody( + model=model, + max_tokens=300, + stream=True, + messages=[ChatMessage(role="user", content=_STREAM_FAILURE_PROMPT)], + ), + ) + frames: Final = outcome.stream_events + assert outcome.is_streaming, ( + f"/v1/messages did not answer with an SSE stream: status={outcome.status_code} body={outcome.body}" + ) + assert frames, ( + f"the proxy sent no SSE data frames although the upstream hung up; stream_error={outcome.stream_error!r}" + ) + assert outcome.stream_error == "event: error", ( + f"the interrupted stream was not announced by an 'event: error' line Anthropic clients read; " + f"stream_error={outcome.stream_error!r} frames={frames}" + ) + try: + AnthropicErrorEvent.model_validate_json(frames[-1]) + except ValidationError: + pytest.fail( + f'the last SSE frame was not an Anthropic {{"type": "error", "error": ...}} envelope; frames={frames}' + ) + torn: Final = tuple(index for index, frame in enumerate(frames) if _payload(frame) is None) + expected_torn: Final = 1 if cut.mid_chunk else 0 + assert len(torn) == expected_torn, ( + f"expected {expected_torn} data line(s) that are not JSON, since the edge tears one only when it " + f"cuts mid-frame, but the proxy relayed {[frames[index] for index in torn]}; all frames={frames}" + ) + for index in torn: + assert _payload(frames[index + 1]) == {"type": "ping"}, ( + f"the frame the upstream tore was not closed as a ping event before the error, so an " + f"Anthropic client parses the error inside it: after {frames[index]!r} came " + f"{frames[index + 1]!r}; all frames={frames}" + ) + bare: Final = tuple(frame for frame in frames if _bare_error_frame(frame)) + assert not bare, ( + f"the proxy emitted error frames without the Anthropic envelope, which Anthropic clients drop: " + f"{bare}; all frames={frames}" + ) + + @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status") + @pytest.mark.parametrize( + ("register", "cut"), + [case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE], + ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE], + ) + def test_upstream_that_hangs_up_before_the_first_byte_raises_with_its_status_in_the_anthropic_sdk( + self, + proxy: ProxyClient, + resources: ResourceManager, + sdk: SdkClients, + register: _CutRegistration, + cut: StreamCut, + ) -> None: + model, key = register(proxy, resources, cut) + client: Final = sdk.anthropic(key) + + with pytest.raises(anthropic.APIStatusError) as raised: + client.messages.create( + model=model, + max_tokens=300, + stream=True, + messages=[_user_turn(_STREAM_FAILURE_PROMPT)], + extra_body=NO_PROXY_CACHE, + ) + assert 500 <= raised.value.status_code < 600, ( + f"an upstream that hung up before sending anything must answer with a server error status the SDK " + f"retries on, not {raised.value.status_code}: {raised.value}" + ) + try: + AnthropicErrorEvent.model_validate(raised.value.body) + except ValidationError: + pytest.fail( + f"the SDK raised with the right status but without the Anthropic error envelope a client reads " + f"the failure from: body={raised.value.body!r} message={raised.value}" + ) + + @pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status") + @pytest.mark.parametrize( + ("register", "cut"), + [case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE], + ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE], + ) + def test_upstream_that_hangs_up_before_the_first_byte_is_a_json_error_with_its_status( + self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut + ) -> None: + model, key = register(proxy, resources, cut) + + outcome: Final = proxy.messages_stream( + key, + AnthropicMessagesBody( + model=model, + max_tokens=300, + stream=True, + messages=[ChatMessage(role="user", content=_STREAM_FAILURE_PROMPT)], + ), + ) + assert not outcome.is_streaming, ( + f"nothing had been streamed when the upstream hung up, yet /v1/messages opened a 200 SSE stream " + f"instead of answering with the failure's status: stream_error={outcome.stream_error!r} " + f"frames={outcome.stream_events}" + ) + assert 500 <= outcome.status_code < 600, ( + f"/v1/messages answered {outcome.status_code} for an upstream that hung up before its first byte; " + f"body={outcome.body}" + ) + try: + AnthropicErrorEvent.model_validate_json(outcome.body) + except ValidationError: + pytest.fail( + f'the error body is not an Anthropic {{"type": "error", "error": ...}} envelope; body={outcome.body}' + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index c96f4b0bef1..84399fd6155 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -598,6 +598,16 @@ class CountTokensResponse(BaseModel): input_tokens: int +class AnthropicErrorBody(BaseModel): + type: str + message: str + + +class AnthropicErrorEvent(BaseModel): + type: Literal["error"] + error: AnthropicErrorBody + + # ---------- mcp servers ---------- diff --git a/tests/e2e/provider_edge.py b/tests/e2e/provider_edge.py index fc10dde2a77..3680375b6af 100644 --- a/tests/e2e/provider_edge.py +++ b/tests/e2e/provider_edge.py @@ -45,6 +45,7 @@ import hashlib import os import re import threading +import time from collections import deque from collections.abc import Callable, Generator, Mapping, Sequence from contextlib import closing, contextmanager @@ -56,6 +57,7 @@ from types import MappingProxyType from typing import Final, Literal, assert_never from urllib.parse import parse_qsl, urlsplit +from botocore.eventstream import EventStreamBuffer from e2e_http import ( NetworkError, StreamChunk, @@ -96,16 +98,18 @@ from fixture_mode import ( ) from fixture_profile import IneligibleRequest, MatchProfile, match_profile, strict_identity from provider_cache import ( + JSON_VALUE, SIGNATURE_HEADERS, CacheEdge, MountPolicy, RequestSigner, + invoke_chunk_value, is_bedrock, scoped_edge_base, split_test_segment, ) from provider_cache_routing import LIVE_PROVIDER_REQUIRED -from pydantic import JsonValue, TypeAdapter +from pydantic import JsonValue, TypeAdapter, ValidationError BEDROCK_REGIONS: Final[tuple[str, ...]] = ("us-east-1",) @@ -537,10 +541,33 @@ class ReplayEdge: source: ReplaySource +@dataclass(frozen=True, slots=True) +class StreamCut: + """Where a live edge hangs up on a streamed upstream body: before its first byte, or with + ``after_content`` set, right after the first transfer chunk carrying assistant output (a + ``content_block_delta``). That frame is what commits the proxy's mid-stream fallback + wrapper to the client: it holds the lifecycle frames before it back and drops them when + the transport fails first, so a cut after a fixed number of chunks landed on either side + of that commit depending on how the provider batched its frames. With ``mid_chunk`` set + the hang-up comes part way through the next ``data:`` line the provider sends after that, + so the client is left inside an SSE frame the way a dropped transport leaves it. + + Whatever was relayed sits on the wire for ``_CUT_SETTLE_SECONDS`` before the hang-up, so + the client has read it by then instead of receiving the data and the close in one burst, + where its reader can surface the close before what it buffered.""" + + after_content: bool + mid_chunk: bool = False + + +_CUT_SETTLE_SECONDS: Final = 1.0 + + @dataclass(frozen=True, slots=True) class LiveEdge: observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None sign: RequestSigner | None = None + cut: StreamCut | None = None type EdgeBackend = RecordEdge | ReplayEdge | LiveEdge | CacheEdge @@ -786,11 +813,128 @@ def _handle_record( assert_never(head) +def _data_line_start(data: bytes) -> int: + if data.startswith(b"data:"): + return 0 + at_line_start: Final = data.find(b"\ndata:") + return -1 if at_line_start < 0 else at_line_start + 1 + + +def _torn_prefix(data: bytes) -> bytes: + start: Final = _data_line_start(data) + line_end: Final = data.find(b"\n", start) + end: Final = len(data) if line_end < 0 else line_end + return data[: start + (end - start) // 2] + + +class _DataLineTearer: + __slots__ = ("_unfinished_line",) + + _unfinished_line: bytes + + def __init__(self) -> None: + self._unfinished_line = b"" + + def observe(self, data: bytes) -> None: + self._unfinished_line = (self._unfinished_line + data).rsplit(b"\n", 1)[-1] + + def tear(self, data: bytes) -> bytes | None: + buffered: Final = self._unfinished_line + data + if _data_line_start(buffered) < 0: + self.observe(data) + return None + return _torn_prefix(buffered)[len(self._unfinished_line):] + + +def _is_content_delta(value: JsonValue | None) -> bool: + return isinstance(value, dict) and value.get("type") == "content_block_delta" + + +def _sse_data_carries_content(line: bytes) -> bool: + if not line.startswith(b"data:"): + return False + try: + return _is_content_delta(JSON_VALUE.validate_json(line[len(b"data:"):].strip())) + except ValidationError: + return False + + +class _AnthropicContentDetector: + __slots__ = ("_unfinished_line",) + + _unfinished_line: bytes + + def __init__(self) -> None: + self._unfinished_line = b"" + + def __call__(self, data: bytes) -> bool: + lines: Final = (self._unfinished_line + data).split(b"\n") + self._unfinished_line = lines[-1] + return any(_sse_data_carries_content(line.rstrip(b"\r")) for line in lines[:-1]) + + +def _invoke_frame_carries_content(payload: bytes) -> bool: + try: + return _is_content_delta(invoke_chunk_value(JSON_VALUE.validate_json(payload))) + except ValidationError: + return False + + +def _bedrock_content_detector() -> Callable[[bytes], bool]: + """Bedrock's invoke stream wraps each Anthropic event in an eventstream frame that a + transfer chunk can split, so the frames are reassembled across chunks before being read.""" + frames: Final = EventStreamBuffer() + + def carries_content(data: bytes) -> bool: + frames.add_data(data) + return any(_invoke_frame_carries_content(frame.payload) for frame in frames) + + return carries_content + + +def _content_detector(mount: str) -> Callable[[bytes], bool]: + return _bedrock_content_detector() if is_bedrock(mount) else _AnthropicContentDetector() + + +def _cut_steps( + steps: Generator[StreamStep, None, None], cut: StreamCut, carries_content: Callable[[bytes], bool] +) -> Generator[StreamStep, None, None]: + with closing(steps) as source: + tearer: Final = _DataLineTearer() + if cut.after_content: + for step in source: + yield step + if isinstance(step, StreamTruncation): + return + tearer.observe(step.data) + if carries_content(step.data): + break + else: + return + if cut.mid_chunk: + for step in source: + if isinstance(step, StreamTruncation): + yield step + return + if (torn := tearer.tear(step.data)) is None: + yield step + continue + if torn: + yield StreamChunk(data=torn) + break + else: + return + if cut.after_content or cut.mid_chunk: + time.sleep(_CUT_SETTLE_SECONDS) + yield StreamTruncation(reason=f"edge cut the upstream stream: {cut!r}") + + def _handle_live( method: str, url: str, headers: Mapping[str, str], body: bytes | None, timeout: float, cache: CacheEdge | None = None, mount: str = "", test_key: str | None = None, observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None, sign: RequestSigner | None = None, + cut: StreamCut | None = None, ) -> EdgeOutcome: forwarded: Final = { name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS @@ -805,6 +949,8 @@ def _handle_live( match head: case NetworkError(message=message): return _recorded_outcome(_network_error_response(message)) + case StreamHead() if cut is not None: + return EdgeStream(head.status_code, _filtered_response_headers(head.headers), _cut_steps(head.steps, cut, _content_detector(mount))) case StreamHead() if _is_streamed(head.headers): return EdgeStream(head.status_code, _filtered_response_headers(head.headers), head.steps) case StreamHead(): @@ -875,10 +1021,10 @@ def handle_edge_request( method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout, backend, mount, test_key, ) - case LiveEdge(observe_request=observe_request, sign=sign): + case LiveEdge(observe_request=observe_request, sign=sign, cut=cut): return _handle_live( method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout, - observe_request=observe_request, sign=sign, + mount=mount, observe_request=observe_request, sign=sign, cut=cut, ) case RecordEdge(): return _handle_record( diff --git a/tests/e2e/test_provider_edge.py b/tests/e2e/test_provider_edge.py index d776c338ef7..978c3671a77 100644 --- a/tests/e2e/test_provider_edge.py +++ b/tests/e2e/test_provider_edge.py @@ -54,11 +54,13 @@ from provider_edge import ( EdgeBackend, EdgeReply, EdgeStream, + LiveEdge, ProviderEdge, ProviderRequestObservation, RecordEdge, ReplayEdge, ReplaySource, + StreamCut, edge_request, handle_edge_request, observed_provider_edge, @@ -1000,6 +1002,36 @@ def stream_chunks(response: RecordedStreamedResponse) -> list[bytes]: return [base64.b64decode(chunk) for chunk in response.chunks_b64] +SECOND_DATA_LINE: Final = b'data: {"type":"content_block_delta","delta":{"text":" two"}}' +SPLIT_MARKER_CHUNKS: tuple[bytes, ...] = ( + b'data: {"type":"content_block_delta","delta":{"text":"one"}}\n\nda', + b"ta" + SECOND_DATA_LINE[4:] + b"\n\nda", + b'ta: {"type":"message_delta","usage":{"output_tokens":7}}\n\nda', + b"ta: [DONE]\n\n", +) + + +class TestStreamCut: + def test_a_mid_frame_cut_tears_a_data_line_whose_marker_is_split_across_chunks(self) -> None: + """Every ``data:`` marker after the first content delta straddles a transfer + chunk boundary, so a tearer that inspects each chunk on its own never finds + one and lets the stream finish cleanly instead of cutting it.""" + backend: Final = LiveEdge(cut=StreamCut(after_content=True, mid_chunk=True)) + with chunked_provider(chunks=SPLIT_MARKER_CHUNKS) as provider: + with running_edge(backend, {"openai": provider_url(provider)}) as edge: + head, chunks, ending = raw_stream_post(edge.port, STREAM_PATH, STREAM_BODY) + + assert head.startswith("HTTP/1.1 200 OK") + assert ending == "truncated" + relayed: Final = b"".join(chunks) + whole: Final = b"".join(SPLIT_MARKER_CHUNKS) + assert whole.startswith(relayed) and relayed != whole + assert relayed.startswith(SPLIT_MARKER_CHUNKS[0]) + torn_line: Final = relayed.rsplit(b"\n", 1)[-1] + assert torn_line and SECOND_DATA_LINE.startswith(torn_line) and torn_line != SECOND_DATA_LINE + assert b"[DONE]" not in relayed + + class TestStreamingFidelity: """LIT-5742: a streamed response records and replays as the chunk sequence the provider actually sent, not as one coalesced body. The unit of fidelity is the diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py index 74c0478b08b..fbcf97839b9 100644 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ b/tests/llm_responses_api_testing/base_responses_api.py @@ -742,7 +742,7 @@ class BaseResponsesAPITest(ABC): Passes tools=[{"type": "shell", "environment": {"type": "container_auto"}}]; validates that the request is accepted and returns a valid response. Only runs for OpenAI; offline coverage for the Azure route lives in - tests/test_litellm/responses/test_responses_api_request_body.py. + tests/unit/responses/test_responses_api_request_body.py. """ base_completion_call_args = self.get_base_completion_call_args() model = ( diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index 47b377dc9a4..da37803b64a 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -381,6 +381,36 @@ class TestBaseResponsesAPIStreamingIterator: ) raise + @staticmethod + def _config_completing_after_one_delta() -> Mock: + mock_config = Mock(spec=BaseResponsesAPIConfig) + completed_response = ResponsesAPIResponse( + id="resp_123", + created_at=0, + status="completed", + model="gpt-5.5", + object="response", + output=[], + usage=ResponseAPIUsage(input_tokens=1, output_tokens=1, total_tokens=2), + ) + + def _transform(model, parsed_chunk, logging_obj): + if parsed_chunk.get("type") == "response.completed": + return ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=completed_response, + ) + return OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_123", + output_index=0, + content_index=0, + delta=parsed_chunk["delta"], + ) + + mock_config.transform_streaming_response.side_effect = _transform + return mock_config + @pytest.mark.asyncio async def test_stop_async_iteration_not_logged_as_failure(self): """ @@ -399,6 +429,7 @@ class TestBaseResponsesAPIStreamingIterator: async def mock_aiter_bytes(): yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n' + yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n' mock_response.aiter_bytes = mock_aiter_bytes @@ -408,11 +439,7 @@ class TestBaseResponsesAPIStreamingIterator: mock_logging_obj.async_failure_handler = Mock() mock_logging_obj.failure_handler = Mock() - mock_config = Mock(spec=BaseResponsesAPIConfig) - mock_delta_event = Mock() - mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA - mock_delta_event.delta = "test" - mock_config.transform_streaming_response.return_value = mock_delta_event + mock_config = self._config_completing_after_one_delta() # Create the iterator instance iterator = ResponsesAPIStreamingIterator( @@ -432,8 +459,9 @@ class TestBaseResponsesAPIStreamingIterator: except StopAsyncIteration: pass # This is expected - # Verify we got the chunk - assert len(chunks_received) == 1 + # Verify we got the delta and the terminal event + assert len(chunks_received) == 2 + assert iterator.completed_response is not None # CRITICAL: Verify that failure handlers were NOT called # StopAsyncIteration is a normal end of stream, not a failure @@ -460,6 +488,7 @@ class TestBaseResponsesAPIStreamingIterator: def mock_iter_bytes(): yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n' + yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n' mock_response.iter_bytes = mock_iter_bytes @@ -469,11 +498,7 @@ class TestBaseResponsesAPIStreamingIterator: mock_logging_obj.async_failure_handler = Mock() mock_logging_obj.failure_handler = Mock() - mock_config = Mock(spec=BaseResponsesAPIConfig) - mock_delta_event = Mock() - mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA - mock_delta_event.delta = "test" - mock_config.transform_streaming_response.return_value = mock_delta_event + mock_config = self._config_completing_after_one_delta() # Create the iterator instance iterator = SyncResponsesAPIStreamingIterator( @@ -493,8 +518,9 @@ class TestBaseResponsesAPIStreamingIterator: except StopIteration: pass # This is expected - # Verify we got the chunk - assert len(chunks_received) == 1 + # Verify we got the delta and the terminal event + assert len(chunks_received) == 2 + assert iterator.completed_response is not None # CRITICAL: Verify that failure handlers were NOT called # StopIteration is a normal end of stream, not a failure diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index 0c3eca52dde..1a34e404d7f 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -1800,7 +1800,7 @@ def test_gemini_image_size_limit_exceeded(monkeypatch): that could cause memory issues and pod crashes. The image fetch is mocked (mirroring the LargeImageClient pattern in - tests/test_litellm/litellm_core_utils/test_image_handling.py) so the test + tests/unit/litellm_core_utils/test_image_handling.py) so the test deterministically exercises the size-limit rejection path without any external network dependency. """ diff --git a/tests/test_litellm/caching/test_caching_handler.py b/tests/test_litellm/caching/test_caching_handler.py deleted file mode 100644 index e5a7f1540ca..00000000000 --- a/tests/test_litellm/caching/test_caching_handler.py +++ /dev/null @@ -1,867 +0,0 @@ -import asyncio -import json -import time -from unittest.mock import MagicMock, patch - -import httpx -import pytest -import respx -from fastapi.testclient import TestClient - -from datetime import datetime -from unittest.mock import AsyncMock - -from litellm.caching.caching_handler import _PENDING_CACHE_WRITES, LLMCachingHandler - - -@pytest.mark.asyncio -async def test_process_async_embedding_cached_response(): - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - args = { - "cached_result": [ - { - "embedding": [-0.025122925639152527, -0.019487135112285614], - "index": 0, - "object": "embedding", - } - ] - } - - mock_logging_obj = MagicMock() - mock_logging_obj.async_success_handler = AsyncMock() - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=args["cached_result"], - kwargs={"model": "text-embedding-ada-002", "input": "test"}, - logging_obj=mock_logging_obj, - start_time=datetime.now(), - model="text-embedding-ada-002", - ) - - assert cache_hit - - print(f"response: {response}") - assert len(response.data) == 1 - - -@pytest.mark.asyncio -async def test_embedding_cache_preserves_prompt_tokens_details(): - """Test that prompt_tokens_details (including image_count) survives a full cache hit.""" - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - cached_result = [ - { - "embedding": [-0.025, -0.019], - "index": 0, - "object": "embedding", - "model": "amazon.titan-embed-image-v1", - "prompt_tokens_details": {"image_count": 1}, - } - ] - - mock_logging_obj = MagicMock() - mock_logging_obj.async_success_handler = AsyncMock() - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=cached_result, - kwargs={"model": "amazon.titan-embed-image-v1", "input": "base64imagedata"}, - logging_obj=mock_logging_obj, - start_time=datetime.now(), - model="amazon.titan-embed-image-v1", - ) - - assert cache_hit - assert response.usage is not None - assert response.usage.prompt_tokens_details is not None - assert response.usage.prompt_tokens_details.image_count == 1 - - -@pytest.mark.asyncio -async def test_embedding_cache_backward_compat_no_prompt_tokens_details(): - """Test that old cached items without prompt_tokens_details still work.""" - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - # Old-format cached item — no prompt_tokens_details field - cached_result = [ - { - "embedding": [-0.025, -0.019], - "index": 0, - "object": "embedding", - "model": "text-embedding-ada-002", - } - ] - - mock_logging_obj = MagicMock() - mock_logging_obj.async_success_handler = AsyncMock() - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=cached_result, - kwargs={"model": "text-embedding-ada-002", "input": "test"}, - logging_obj=mock_logging_obj, - start_time=datetime.now(), - model="text-embedding-ada-002", - ) - - assert cache_hit - assert response.usage is not None - assert response.usage.prompt_tokens_details is None - - -@pytest.mark.asyncio -async def test_embedding_cache_aggregates_multiple_image_counts(): - """Test that image_count is summed correctly across multiple cached items.""" - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - cached_result = [ - { - "embedding": [-0.025, -0.019], - "index": 0, - "object": "embedding", - "model": "amazon.titan-embed-image-v1", - "prompt_tokens_details": {"image_count": 1}, - }, - { - "embedding": [0.031, 0.042], - "index": 1, - "object": "embedding", - "model": "amazon.titan-embed-image-v1", - "prompt_tokens_details": {"image_count": 1}, - }, - ] - - mock_logging_obj = MagicMock() - mock_logging_obj.async_success_handler = AsyncMock() - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=cached_result, - kwargs={ - "model": "amazon.titan-embed-image-v1", - "input": ["img1", "img2"], - }, - logging_obj=mock_logging_obj, - start_time=datetime.now(), - model="amazon.titan-embed-image-v1", - ) - - assert cache_hit - assert response.usage.prompt_tokens_details is not None - assert response.usage.prompt_tokens_details.image_count == 2 - - -def test_combine_usage_merges_prompt_tokens_details(): - """Test that combine_usage merges prompt_tokens_details from both Usage objects.""" - from litellm.types.utils import PromptTokensDetailsWrapper, Usage - - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - usage1 = Usage( - prompt_tokens=10, - completion_tokens=0, - total_tokens=10, - prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1), - ) - usage2 = Usage( - prompt_tokens=20, - completion_tokens=0, - total_tokens=20, - prompt_tokens_details=PromptTokensDetailsWrapper(image_count=2), - ) - - combined = llm_caching_handler.combine_usage(usage1, usage2) - - assert combined.prompt_tokens == 30 - assert combined.total_tokens == 30 - assert combined.prompt_tokens_details is not None - assert combined.prompt_tokens_details.image_count == 3 - - -def test_combine_usage_handles_none_details(): - """Test that combine_usage works when one or both sides have null prompt_tokens_details.""" - from litellm.types.utils import PromptTokensDetailsWrapper, Usage - - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - # Both null - usage_a = Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) - usage_b = Usage(prompt_tokens=20, completion_tokens=0, total_tokens=20) - combined = llm_caching_handler.combine_usage(usage_a, usage_b) - assert combined.prompt_tokens_details is None - - # Only first has details - usage_c = Usage( - prompt_tokens=10, - completion_tokens=0, - total_tokens=10, - prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1), - ) - combined = llm_caching_handler.combine_usage(usage_c, usage_b) - assert combined.prompt_tokens_details is not None - assert combined.prompt_tokens_details.image_count == 1 - - # Only second has details - combined = llm_caching_handler.combine_usage(usage_a, usage_c) - assert combined.prompt_tokens_details is not None - assert combined.prompt_tokens_details.image_count == 1 - - -def test_is_chat_completion_cached_dict(): - from litellm.caching.caching_handler import _is_chat_completion_cached_dict - - assert _is_chat_completion_cached_dict( - {"id": "chatcmpl-abc", "object": "chat.completion", "choices": []} - ) - assert _is_chat_completion_cached_dict( - {"id": "other", "object": "chat.completion.chunk", "choices": []} - ) - assert _is_chat_completion_cached_dict( - {"id": "no-object", "choices": [{"index": 0}]} - ) - assert not _is_chat_completion_cached_dict( - {"id": "resp_abc", "object": "response", "output": []} - ) - - -def _build_logging_obj(call_type: str, stream: bool): - import uuid as _uuid - - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging - - return LiteLLMLogging( - litellm_call_id=str(datetime.now()), - call_type=call_type, - model="gpt-5.4", - messages=[], - function_id=str(_uuid.uuid4()), - stream=stream, - start_time=datetime.now(), - ) - - -def test_convert_cached_aresponses_bridge_chat_completion_stream(): - """openai/responses chat-completions bridge: streaming cache hit replays as chat stream.""" - from litellm import aresponses - from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - from litellm.types.utils import CallTypes - - caching_handler = LLMCachingHandler( - original_function=aresponses, request_kwargs={}, start_time=datetime.now() - ) - cached_result = { - "id": "chatcmpl-bridge-cache-test", - "object": "chat.completion", - "created": int(time.time()), - "model": "gpt-5.4", - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": "Hi!"}, - "finish_reason": "stop", - } - ], - "usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}, - } - - result = caching_handler._convert_cached_result_to_model_response( - cached_result=cached_result, - call_type=CallTypes.aresponses.value, - kwargs={ - "model": "gpt-5.4", - "stream": True, - "messages": [{"role": "user", "content": "hi"}], - }, - logging_obj=_build_logging_obj(CallTypes.aresponses.value, stream=True), - model="gpt-5.4", - args=(), - ) - - assert isinstance(result, CustomStreamWrapper) - - -def test_convert_cached_responses_bridge_chat_completion_nonstream(): - """openai/responses chat-completions bridge: non-streaming cache hit replays as ModelResponse.""" - from litellm import responses - from litellm.types.utils import CallTypes, ModelResponse - - caching_handler = LLMCachingHandler( - original_function=responses, request_kwargs={}, start_time=datetime.now() - ) - cached_result = { - "id": "chatcmpl-bridge-nonstream", - "object": "chat.completion", - "created": int(time.time()), - "model": "gpt-5.4", - "choices": [ - { - "index": 0, - "message": {"role": "assistant", "content": "Hi!"}, - "finish_reason": "stop", - } - ], - "usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}, - } - - result = caching_handler._convert_cached_result_to_model_response( - cached_result=cached_result, - call_type=CallTypes.responses.value, - kwargs={ - "model": "gpt-5.4", - "stream": False, - "messages": [{"role": "user", "content": "hi"}], - }, - logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False), - model="gpt-5.4", - args=(), - ) - - assert isinstance(result, ModelResponse) - assert result.choices[0].message.content == "Hi!" - - -def test_convert_cached_responses_legacy_nonstream_path(): - """Genuine ResponsesAPIResponse dict (no chatcmpl/choices) falls through legacy path.""" - from litellm import responses - from litellm.types.llms.openai import ResponsesAPIResponse - from litellm.types.utils import CallTypes - - caching_handler = LLMCachingHandler( - original_function=responses, request_kwargs={}, start_time=datetime.now() - ) - cached_result = { - "id": "resp_legacy_nonstream", - "created_at": int(time.time()), - "status": "completed", - "model": "gpt-4o", - "object": "response", - "output": [ - { - "type": "message", - "id": "msg_legacy", - "status": "completed", - "role": "assistant", - "content": [ - { - "type": "output_text", - "text": "legacy response", - "annotations": [], - } - ], - } - ], - } - - result = caching_handler._convert_cached_result_to_model_response( - cached_result=cached_result, - call_type=CallTypes.responses.value, - kwargs={"model": "gpt-4o", "input": "hi", "stream": False}, - logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False), - model="gpt-4o", - args=(), - ) - - assert isinstance(result, ResponsesAPIResponse) - assert result.id == "resp_legacy_nonstream" - - -def test_convert_cached_responses_legacy_stream_path(): - """Genuine ResponsesAPIResponse dict (no chatcmpl/choices) on stream falls through legacy path.""" - from litellm import responses - from litellm.responses.streaming_iterator import ( - CachedResponsesAPIStreamingIterator, - ) - from litellm.types.utils import CallTypes - - caching_handler = LLMCachingHandler( - original_function=responses, request_kwargs={}, start_time=datetime.now() - ) - cached_result = { - "id": "resp_legacy_stream", - "created_at": int(time.time()), - "status": "completed", - "model": "gpt-4o", - "object": "response", - "output": [ - { - "type": "message", - "id": "msg_legacy_stream", - "status": "completed", - "role": "assistant", - "content": [ - { - "type": "output_text", - "text": "legacy stream", - "annotations": [], - } - ], - } - ], - } - - result = caching_handler._convert_cached_result_to_model_response( - cached_result=cached_result, - call_type=CallTypes.responses.value, - kwargs={"model": "gpt-4o", "input": "hi", "stream": True}, - logging_obj=_build_logging_obj(CallTypes.responses.value, stream=True), - model="gpt-4o", - args=(), - ) - - assert isinstance(result, CachedResponsesAPIStreamingIterator) - - -@pytest.mark.asyncio -async def test_embedding_cache_restores_stored_prompt_tokens_for_image_input(): - """Image-embedding cache hit restores prompt_tokens=0 from the stored value - instead of recomputing a bogus count by tokenizing the base64 input.""" - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - # base64-like blob — token_counter over this would return a large nonzero count - image_input = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk" * 50 - - cached_result = [ - { - "embedding": [-0.025, -0.019], - "index": 0, - "object": "embedding", - "model": "amazon.titan-embed-image-v1", - "prompt_tokens": 0, - "prompt_tokens_details": {"image_count": 1}, - } - ] - - mock_logging_obj = MagicMock() - mock_logging_obj.async_success_handler = AsyncMock() - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=cached_result, - kwargs={"model": "amazon.titan-embed-image-v1", "input": image_input}, - logging_obj=mock_logging_obj, - start_time=datetime.now(), - model="amazon.titan-embed-image-v1", - ) - - assert cache_hit - assert response.usage is not None - assert response.usage.prompt_tokens == 0 - assert response.usage.total_tokens == 0 - assert response.usage.prompt_tokens_details.image_count == 1 - - -@pytest.mark.asyncio -async def test_embedding_cache_sums_stored_prompt_tokens_across_items(): - """A multi-item cache hit sums the stored per-item prompt_tokens back to the total.""" - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - cached_result = [ - { - "embedding": [-0.01], - "index": 0, - "object": "embedding", - "model": "text-embedding-3-small", - "prompt_tokens": 5, - }, - { - "embedding": [-0.02], - "index": 1, - "object": "embedding", - "model": "text-embedding-3-small", - "prompt_tokens": 4, - }, - ] - - mock_logging_obj = MagicMock() - mock_logging_obj.async_success_handler = AsyncMock() - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=cached_result, - kwargs={"model": "text-embedding-3-small", "input": ["hello world", "foo bar"]}, - logging_obj=mock_logging_obj, - start_time=datetime.now(), - model="text-embedding-3-small", - ) - - assert cache_hit - assert response.usage.prompt_tokens == 9 - assert response.usage.total_tokens == 9 - - -@pytest.mark.asyncio -async def test_embedding_cache_falls_back_to_token_counter_for_legacy_entries(): - """Legacy cache entries with no stored prompt_tokens still recompute via token_counter - for str inputs (backward compatibility).""" - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - # No prompt_tokens key — pre-fix entry - cached_result = [ - { - "embedding": [-0.025, -0.019], - "index": 0, - "object": "embedding", - "model": "text-embedding-ada-002", - }, - ] - - mock_logging_obj = MagicMock() - mock_logging_obj.async_success_handler = AsyncMock() - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=cached_result, - kwargs={"model": "text-embedding-ada-002", "input": "hello world"}, - logging_obj=mock_logging_obj, - start_time=datetime.now(), - model="text-embedding-ada-002", - ) - - assert cache_hit - # token_counter over "hello world" yields a nonzero count — fallback path still runs - assert response.usage.prompt_tokens > 0 - - -@pytest.mark.asyncio -async def test_embedding_cache_hit_sets_custom_llm_provider_on_logging_obj(): - """A full embedding cache hit must stamp the resolved provider onto the logging - obj so spend logs record the provider instead of None/unknown.""" - from litellm.types.utils import CallTypes - - llm_caching_handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs={}, - start_time=datetime.now(), - ) - - cached_result = [ - { - "embedding": [-0.025, -0.019], - "index": 0, - "object": "embedding", - "model": "text-embedding-3-small", - "prompt_tokens": 5, - } - ] - - logging_obj = _build_logging_obj(CallTypes.aembedding.value, stream=False) - logging_obj.async_success_handler = AsyncMock() - - response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( - final_embedding_cached_response=None, - cached_result=cached_result, - kwargs={"model": "text-embedding-3-small", "input": "hello world"}, - logging_obj=logging_obj, - start_time=datetime.now(), - model="text-embedding-3-small", - ) - - assert cache_hit - assert logging_obj.model_call_details["custom_llm_provider"] == "openai" - - -def test_sync_stream_responses_cache_hit_sets_custom_llm_provider_on_logging_obj(monkeypatch): - import litellm - from litellm.caching.caching import Cache - from litellm.types.utils import CallTypes - - monkeypatch.setattr(litellm, "cache", Cache(type="local")) - kwargs = {"model": "azure/gpt-5.4-mini", "input": "hello", "stream": True} - cached_response = { - "id": "resp_sync_stream", - "created_at": int(time.time()), - "status": "completed", - "model": "gpt-5.4-mini", - "object": "response", - "output": [ - { - "type": "message", - "id": "msg_sync_stream", - "status": "completed", - "role": "assistant", - "content": [{"type": "output_text", "text": "hi", "annotations": []}], - } - ], - } - litellm.cache.add_cache(json.dumps(cached_response), **kwargs) - handler = LLMCachingHandler(original_function=litellm.responses, request_kwargs=kwargs, start_time=datetime.now()) - logging_obj = _build_logging_obj(CallTypes.responses.value, stream=True) - - hit = handler._sync_get_cache( - model="azure/gpt-5.4-mini", - original_function=litellm.responses, - logging_obj=logging_obj, - start_time=datetime.now(), - call_type=CallTypes.responses.value, - kwargs=kwargs, - args=(), - ) - - assert hit.cached_result is not None - assert logging_obj.model_call_details["custom_llm_provider"] == "azure" - assert logging_obj.model_call_details["litellm_params"]["custom_llm_provider"] == "azure" - - -def test_request_kwargs_does_not_retain_logging_obj(): - """ - The caching handler lives on logging_obj._llm_caching_handler, so keeping - litellm_logging_obj inside request_kwargs closes a reference cycle - (Logging -> LLMCachingHandler -> kwargs -> Logging). That cycle keeps the - full request payload alive until a generational GC pass instead of being - freed by refcount when the request finishes; under bursts of large-token - requests this presents as stepwise RSS growth that never returns to - baseline. Other kwargs (messages included) must be preserved. - """ - logging_obj = MagicMock() - kwargs = { - "model": "gpt-4o", - "messages": [{"role": "user", "content": "hello"}], - "litellm_logging_obj": logging_obj, - } - - handler = LLMCachingHandler( - original_function=MagicMock(), - request_kwargs=kwargs, - start_time=datetime.now(), - ) - - assert "litellm_logging_obj" not in handler.request_kwargs - assert handler.request_kwargs["messages"] == kwargs["messages"] - assert handler.request_kwargs["model"] == "gpt-4o" - - -def test_async_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch): - """ - Regression test for the SDK losing async cache writes in short-lived scripts: - async_set_cache dispatched the write as a bare fire-and-forget task, so - asyncio.run cancelled it at loop close before the write landed (LIT-6184, - deterministic with hiredis installed). The write must survive loop shutdown. - """ - import litellm - - writes = [] - - class _SlowWriteCache: - supported_call_types = ["acompletion"] - cache = None - - async def async_add_cache(self, result, dynamic_cache_object=None, **kwargs): - await asyncio.sleep(0.2) - writes.append(result) - - async def acompletion(**kwargs): - return None - - handler = LLMCachingHandler( - original_function=acompletion, - request_kwargs={}, - start_time=datetime.now(), - ) - monkeypatch.setattr(litellm, "cache", _SlowWriteCache()) - - async def _short_lived_script(): - await handler.async_set_cache( - result=litellm.ModelResponse(), - original_function=acompletion, - kwargs={}, - ) - - asyncio.run(_short_lived_script()) - - assert len(writes) == 1 - - -@pytest.mark.asyncio -async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monkeypatch): - """The spend log for a cache hit must reuse the key the lookup already computed instead of hashing again.""" - import litellm - from litellm.caching.caching import Cache - from litellm.types.utils import CallTypes - - async def acompletion(**kwargs): - return None - - monkeypatch.setattr(litellm, "cache", Cache(type="local")) - kwargs = {"model": "gpt-5.4", "messages": [{"role": "user", "content": "hello"}], "caching": True} - await litellm.cache.async_add_cache( - litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "hi"}}]), **kwargs - ) - handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now()) - logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False) - logging_obj.async_success_handler = AsyncMock() - - hit = await handler._async_get_cache( - model="gpt-5.4", - original_function=acompletion, - logging_obj=logging_obj, - start_time=datetime.now(), - call_type=CallTypes.acompletion.value, - kwargs=kwargs, - args=(), - ) - - assert hit is not None and hit.cached_result is not None - assert handler.preset_cache_key is not None - assert logging_obj.litellm_params["preset_cache_key"] == handler.preset_cache_key - assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key - - -@pytest.mark.asyncio -async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_time(monkeypatch): - import litellm - from litellm.caching.caching import Cache - from litellm.types.utils import CallTypes - - async def aanthropic_messages(**kwargs): - return None - - monkeypatch.setattr(litellm, "cache", Cache(type="local")) - kwargs = { - "model": "claude-sonnet-5", - "messages": [{"role": "user", "content": "hello"}], - "max_tokens": 16, - "caching": True, - "stream": False, - "_websearch_interception_converted_stream": True, - } - cached_message = { - "id": "msg_1", - "type": "message", - "role": "assistant", - "content": [{"type": "text", "text": "hi"}], - } - await litellm.cache.async_add_cache(cached_message, **kwargs) - handler = LLMCachingHandler(original_function=aanthropic_messages, request_kwargs=kwargs, start_time=datetime.now()) - logging_obj = _build_logging_obj(CallTypes.aanthropic_messages.value, stream=False) - logging_obj.async_success_handler = AsyncMock() - logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() - - hit = await handler._async_get_cache( - model="claude-sonnet-5", - original_function=aanthropic_messages, - logging_obj=logging_obj, - start_time=datetime.now(), - call_type=CallTypes.aanthropic_messages.value, - kwargs=kwargs, - args=(), - ) - - assert hit is not None and hit.cached_result == cached_message - logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once() - assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True - - -@pytest.mark.asyncio -async def test_agentic_loop_followup_cache_hit_with_converted_stream_marker_replays_as_plain_object(monkeypatch): - import litellm - from litellm.caching.caching import Cache - from litellm.types.utils import CallTypes - - async def acompletion(**kwargs): - return None - - monkeypatch.setattr(litellm, "cache", Cache(type="local")) - kwargs = { - "model": "gpt-5.6", - "messages": [{"role": "user", "content": "run the code"}], - "caching": True, - "stream": False, - "_code_interpreter_interception_converted_stream": True, - "_agentic_loop_depth": 1, - } - await litellm.cache.async_add_cache( - litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "done"}}]), **kwargs - ) - handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now()) - logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False) - logging_obj.async_success_handler = AsyncMock() - logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() - - hit = await handler._async_get_cache( - model="gpt-5.6", - original_function=acompletion, - logging_obj=logging_obj, - start_time=datetime.now(), - call_type=CallTypes.acompletion.value, - kwargs=kwargs, - args=(), - ) - - assert hit is not None and isinstance(hit.cached_result, litellm.ModelResponse) - assert hit.cached_result.choices[0].message.content == "done" - logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once() - assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True - - -@pytest.mark.asyncio -async def test_partial_embedding_cache_hit_sends_only_misses_and_keeps_input_order(monkeypatch): - import litellm - from litellm import CustomLLM - from litellm.caching.caching import Cache - from litellm.types.utils import Embedding, EmbeddingResponse - - class RecordingEmbedder(CustomLLM): - provider_inputs: tuple[tuple[str, ...], ...] = () - - async def aembedding(self, model, input, model_response, **kwargs) -> EmbeddingResponse: - self.provider_inputs = (*self.provider_inputs, tuple(input)) - return EmbeddingResponse( - model=model, - data=[ - Embedding(embedding=[float(len(text))], index=idx, object="embedding") - for idx, text in enumerate(input) - ], - ) - - embedder = RecordingEmbedder() - monkeypatch.setattr(litellm, "custom_provider_map", [{"provider": "recording-embedder", "custom_handler": embedder}]) - monkeypatch.setattr(litellm, "provider_list", [*litellm.provider_list, "recording-embedder"]) - monkeypatch.setattr(litellm, "_custom_providers", [*litellm._custom_providers, "recording-embedder"]) - monkeypatch.setattr(litellm, "cache", Cache(type="local")) - - await litellm.aembedding(model="recording-embedder/m", input=["aa", "bbbb"]) - await asyncio.gather(*_PENDING_CACHE_WRITES) - mixed_input = ["c", "aa", "ddd", "bbbb", "eeeee"] - response = await litellm.aembedding(model="recording-embedder/m", input=mixed_input) - await asyncio.gather(*_PENDING_CACHE_WRITES) - - assert embedder.provider_inputs == (("aa", "bbbb"), ("c", "ddd", "eeeee")), embedder.provider_inputs - assert [item["index"] for item in response.data] == [0, 1, 2, 3, 4] - assert [item["embedding"] for item in response.data] == [[float(len(text))] for text in mixed_input] - assert response._hidden_params["cache_hit"] is True, "a partial hit must still be reported as a cache hit" - - repeat = await litellm.aembedding(model="recording-embedder/m", input=mixed_input) - - assert len(embedder.provider_inputs) == 2, embedder.provider_inputs - assert [item["embedding"] for item in repeat.data] == [[float(len(text))] for text in mixed_input] diff --git a/tests/test_litellm/conftest.py b/tests/test_litellm/conftest.py index f8c7d5273d1..f83c1e76b3a 100644 --- a/tests/test_litellm/conftest.py +++ b/tests/test_litellm/conftest.py @@ -22,19 +22,6 @@ import litellm from litellm import router as litellm_router_module from litellm import utils as litellm_utils_module from litellm._logging import ALL_LOGGERS -from litellm.litellm_core_utils.cli_keyring import ( - KeyringDiscardsWrites, - KeyringUnreachable, - KeyringUnusable, - SecretErase, - SecretErased, - SecretFound, - SecretMissing, - SecretRead, - SecretStored, - SecretStranded, - SecretWrite, -) from litellm.litellm_core_utils.prompt_templates import ( image_handling as image_handling_module, ) @@ -42,6 +29,7 @@ from litellm.llms.custom_httpx.async_client_cleanup import ( close_litellm_async_clients, ) from litellm.proxy.db import tool_registry_writer as tool_registry_writer_module +from tests.unit.litellm_core_utils.fake_secret_vault import FakeSecretVault def _reset_module_level_aws_auth_caches(): @@ -128,60 +116,6 @@ def isolate_host_os_keychain(monkeypatch): monkeypatch.setenv("LITELLM_CLI_DISABLE_KEYRING", "1") -class FakeSecretVault: - """In-memory stand-in for the OS keychain, injected wherever CLI credential storage is exercised. - - `available=False` models a keychain that is locked or has no backend, `writable=False` one that - refuses to store, `erasable=False` one that will not release what it already holds, and `failure` - picks which unusable state those report. `discards=True` is keyring's null backend, which answers - reads and erases like any other yet keeps nothing it is given, so only writes report it. - """ - - def __init__( - self, - blob: str | None = None, - *, - available: bool = True, - writable: bool = True, - erasable: bool = True, - discards: bool = False, - failure: KeyringUnusable = KeyringUnreachable(), - ) -> None: - self.blob: str | None = blob - self.available: bool = available - self.writable: bool = writable - self.erasable: bool = erasable - self.discards: bool = discards - self.failure: KeyringUnusable = failure - self.reads: int = 0 - self.writes: list[str] = [] - self.erases: int = 0 - - def read(self) -> SecretRead: - self.reads += 1 - if not self.available: - return self.failure - return SecretMissing() if self.blob is None else SecretFound(self.blob) - - def write(self, blob: str) -> SecretWrite: - self.writes.append(blob) - if not (self.available and self.writable): - return self.failure - if self.discards: - return KeyringDiscardsWrites() - self.blob = blob - return SecretStored() - - def erase(self) -> SecretErase: - self.erases += 1 - if not self.available: - return self.failure - if not self.erasable: - return SecretStranded() if self.blob is not None else SecretErased() - self.blob = None - return SecretErased() - - @pytest.fixture def secret_vault_factory(): """Build FakeSecretVault instances; see its docstring for the failure modes it can model.""" diff --git a/tests/test_litellm/litellm_core_utils/__init__.py b/tests/test_litellm/litellm_core_utils/__init__.py index 8c64613a5da..e69de29bb2d 100644 --- a/tests/test_litellm/litellm_core_utils/__init__.py +++ b/tests/test_litellm/litellm_core_utils/__init__.py @@ -1 +0,0 @@ -# This file makes the tests/litellm/litellm_core_utils directory a Python package diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py index cfa5b00e6ca..1e10b7e82b1 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -1,445 +1,6 @@ -#### What this tests #### -# This tests litellm.token_counter.token_counter() function -import asyncio -import base64 -import importlib -import struct -import threading -import time -import traceback -from concurrent.futures import Future, wait -from typing import Final -from unittest.mock import MagicMock - -import anyio.to_thread import pytest -import tiktoken - -from unittest.mock import AsyncMock, patch - -import litellm -from litellm import create_pretrained_tokenizer, decode, encode, get_modified_max_tokens -from litellm import token_counter as token_counter_old -import litellm.constants -from litellm.constants import TOKEN_COUNTER_MAX_CONCURRENT_COUNTS -from litellm.litellm_core_utils.asyncify import asyncify -from litellm.litellm_core_utils.token_counter import ( - _get_exact_count_function, - _get_extrapolating_count_function, - _get_tiktoken_count_function, - calculate_img_tokens, - get_image_dimensions, - high_detail_image_token_upper_bound, - image_dimensions_from_bytes, - offload_token_count, -) -from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new -from tests.large_text import text -from tests.test_litellm.litellm_core_utils.event_loop_lag import ( - assert_loop_stayed_free, - timed_with_loop_lags, - warm_tokenizer, -) -from tests.test_litellm.litellm_core_utils.messages_with_counts import ( - MESSAGES_TEXT, - MESSAGES_WITH_IMAGES, - MESSAGES_WITH_TOOLS, -) - - -def token_counter_both_assert_same(**args): - new = token_counter_new(**args) - old = token_counter_old(**args) - assert new == old, f"New token counter {new} does not match old token counter {old}" - return new - - -## Choose which token_counter the test will use. - -# token_counter = token_counter_new -# token_counter = token_counter_old -token_counter = token_counter_both_assert_same - - -def test_token_counter_basic(): - assert ( - token_counter( - model="claude-2", - messages=[ - { - "role": "user", - "content": "This is a long message that definitely exceeds the token limit.", - } - ], - ) - == 19 - ) - - -def test_token_counter_large_repeated_text_is_fast(): - messages = [{"role": "user", "content": [{"type": "text", "text": "A" * 1024 * 1024}]}] - - start_time = time.perf_counter() - tokens = token_counter_new(model="us.anthropic.claude-sonnet-4-6", messages=messages) - elapsed = time.perf_counter() - start_time - - assert elapsed < 2, f"Token counting took too long: {elapsed:.2f}s" - assert tokens > 0 - - -@pytest.mark.parametrize( - "text", - [ - "Short text", - "This is a normal message with punctuation, numbers, and a few words.", - ], -) -def test_token_counter_short_text_matches_tiktoken(text): - encoding = tiktoken.get_encoding("cl100k_base") - expected = len(encoding.encode(text, disallowed_special=())) - - assert token_counter_new(model="us.anthropic.claude-sonnet-4-6", text=text) == expected - - -def test_token_counter_default_encoding_matches_cl100k(): - encoding: Final = tiktoken.get_encoding("cl100k_base") - expected: Final = len(encoding.encode("hello world", disallowed_special=())) - - assert token_counter_new(model=None, text="hello world") == expected - - -def test_token_counter_text_over_chunk_boundary_stays_close_to_tiktoken(): - text = ("The quick brown fox jumps over the lazy dog. " * 30)[:1025] - encoding = tiktoken.get_encoding("cl100k_base") - expected = len(encoding.encode(text, disallowed_special=())) - - actual = token_counter_new(model="us.anthropic.claude-sonnet-4-6", text=text) - - assert abs(actual - expected) <= 4 - - -@pytest.mark.parametrize( - "configured", - ["0", "-1", "-1024", "not-an-int", "", " ", "999999999", "inf", "1e9"], -) -def test_invalid_chunk_size_config_stays_usable(monkeypatch, configured): - """A misconfigured chunk size must not raise, count zero, or restore the quadratic encode cost.""" - monkeypatch.setenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS", configured) - try: - reloaded = importlib.reload(litellm.constants) - chunk_size = reloaded.TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS - assert 1 <= chunk_size <= reloaded.TIKTOKEN_ENCODE_MAX_CHUNK_SIZE_CHARS - - encoding = tiktoken.get_encoding("cl100k_base") - count_tokens = _get_tiktoken_count_function( - lambda text: len(encoding.encode(text, disallowed_special=())), - chunk_size=chunk_size, - ) - assert count_tokens("The quick brown fox jumps over the lazy dog. " * 40) > 0 - finally: - monkeypatch.delenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS") - importlib.reload(litellm.constants) - - -def test_valid_chunk_size_config_is_honoured(monkeypatch): - monkeypatch.setenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS", "2048") - try: - assert importlib.reload(litellm.constants).TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS == 2048 - finally: - monkeypatch.delenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS") - importlib.reload(litellm.constants) - - -async def test_huggingface_count_in_a_worker_thread_leaves_the_event_loop_free(): - warm_tokenizer("claude-fable-5") - - tokens, took, lags = await timed_with_loop_lags( - lambda: asyncify(token_counter_new)(model="claude-fable-5", text=text * 100) - ) - - assert tokens > 0 - assert_loop_stayed_free(took, lags) - - -@pytest.mark.parametrize("max_exact_chars", [64, 1_000, 2_500]) -def test_count_above_the_cap_samples_the_whole_string_and_scales(max_exact_chars: int): - count_exactly: Final = MagicMock(side_effect=lambda chunk: chunk.count("a") + len(chunk)) - front_heavy: Final = "a" * 1_000 + "b" * 4_000 - exact: Final = 1_000 + len(front_heavy) - - estimate: Final = _get_extrapolating_count_function(count_exactly, max_exact_chars=max_exact_chars)(front_heavy) - - assert abs(estimate - exact) <= exact // 100 - assert sum(len(call.args[0]) for call in count_exactly.call_args_list) <= max_exact_chars - - -def test_count_at_or_below_the_cap_is_exact(): - count_exactly: Final = MagicMock(side_effect=len) - - assert _get_extrapolating_count_function(count_exactly, max_exact_chars=5_000)("a" * 5_000) == 5_000 - assert count_exactly.call_args_list == [(("a" * 5_000,),)] - - -class _SlowEncoder: - def __init__(self) -> None: - self._lock: Final = threading.Lock() - self.in_flight = 0 - self.peak_in_flight = 0 - - def encode_batch_fast(self, texts: list[str]) -> list[list[int]]: - with self._lock: - self.in_flight += 1 - self.peak_in_flight = max(self.peak_in_flight, self.in_flight) - time.sleep(0.1) - with self._lock: - self.in_flight -= 1 - return [[0] * len(text) for text in texts] - - -@pytest.mark.asyncio -async def test_offloaded_counts_do_not_borrow_from_the_shared_thread_pool(): - encoder: Final = _SlowEncoder() - count: Final = _get_exact_count_function(None, {"type": "huggingface_tokenizer", "tokenizer": encoder}) - shared_pool: Final = anyio.to_thread.current_default_thread_limiter() - burst: Final = 2 * TOKEN_COUNTER_MAX_CONCURRENT_COUNTS - - async def shared_pool_borrowed_until_done(counting: asyncio.Future[list[int]]) -> tuple[int, ...]: - if counting.done(): - return () - await asyncio.sleep(0.01) - return (shared_pool.borrowed_tokens, *await shared_pool_borrowed_until_done(counting)) - - counting: Final = asyncio.ensure_future(asyncio.gather(*(offload_token_count(count)("abc") for _ in range(burst)))) - borrowed: Final = await shared_pool_borrowed_until_done(counting) - - assert await counting == [3] * burst - assert len(borrowed) > 1 and max(borrowed) == 0 - assert 1 < encoder.peak_in_flight <= TOKEN_COUNTER_MAX_CONCURRENT_COUNTS - - -def _count_in_a_fresh_event_loop(text: str, result: Future[int]) -> None: - def slow_count(counted: str) -> int: - time.sleep(0.1) - return len(counted) - - result.set_result(asyncio.run(offload_token_count(slow_count)(text))) - - -def test_offloaded_counts_finish_in_every_event_loop_that_shares_the_process(): - loops: Final = 2 * TOKEN_COUNTER_MAX_CONCURRENT_COUNTS - results: Final = tuple(Future[int]() for _ in range(loops)) - threads: Final = tuple( - threading.Thread(target=_count_in_a_fresh_event_loop, args=("a" * size, result), daemon=True) - for size, result in enumerate(results, start=1) - ) - for thread in threads: - thread.start() - - _, pending = wait(results, timeout=5) - - assert not pending - assert tuple(result.result() for result in results) == tuple(range(1, loops + 1)) - - -@pytest.mark.parametrize( - ("configured", "expected"), - [("8", 8), ("0", 4), ("not-an-int", 4)], -) -def test_max_concurrent_counts_config_is_honoured(monkeypatch: pytest.MonkeyPatch, configured: str, expected: int): - monkeypatch.setenv("TOKEN_COUNTER_MAX_CONCURRENT_COUNTS", configured) - try: - assert importlib.reload(litellm.constants).TOKEN_COUNTER_MAX_CONCURRENT_COUNTS == expected - finally: - monkeypatch.delenv("TOKEN_COUNTER_MAX_CONCURRENT_COUNTS") - importlib.reload(litellm.constants) - - -def test_token_counter_applies_the_default_cap(): - max_exact_chars: Final = litellm.constants.TOKEN_COUNTER_MAX_EXACT_CHARS - prose: Final = ("The quick brown fox jumps over the lazy dog. " * (max_exact_chars // 45 + 1))[:max_exact_chars] - over_the_cap: Final = prose + "a" * 200_000 - exact: Final = _get_exact_count_function("gpt-5.6")(over_the_cap) - - estimate: Final = token_counter_new(model="gpt-5.6", text=over_the_cap) - - assert estimate != exact - assert abs(estimate - exact) <= exact // 100 - - -@pytest.mark.parametrize( - ("configured", "expected"), - [("2048", 2048), ("0", 4_000_000), ("not-an-int", 4_000_000)], -) -def test_max_exact_chars_config_is_honoured(monkeypatch: pytest.MonkeyPatch, configured: str, expected: int): - monkeypatch.setenv("TOKEN_COUNTER_MAX_EXACT_CHARS", configured) - try: - assert importlib.reload(litellm.constants).TOKEN_COUNTER_MAX_EXACT_CHARS == expected - finally: - monkeypatch.delenv("TOKEN_COUNTER_MAX_EXACT_CHARS") - importlib.reload(litellm.constants) - - -def test_token_counter_with_prefix(): - messages = [ - {"role": "user", "content": "Who won the world cup in 2022?"}, - {"role": "assistant", "content": "Argentina", "prefix": True}, - ] - tokens = token_counter(model="gpt-3.5-turbo", messages=messages) - assert tokens == 22, f"Expected 22 tokens, got {tokens}" - - -def test_token_counter_normal_plus_function_calling(): - messages = [ - {"role": "system", "content": "System prompt"}, - {"role": "user", "content": "content1"}, - {"role": "assistant", "content": "content2"}, - {"role": "user", "content": "conten3"}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_E0lOb1h6qtmflUyok4L06TgY", - "function": { - "arguments": '{"query":"search query","domain":"google.ca","gl":"ca","hl":"en"}', - "name": "SearchInternet", - }, - "type": "function", - } - ], - }, - { - "tool_call_id": "call_E0lOb1h6qtmflUyok4L06TgY", - "role": "tool", - "name": "SearchInternet", - "content": "tool content", - }, - ] - tokens = token_counter(model="gpt-3.5-turbo", messages=messages) - assert tokens == 80 - - -# test_token_counter_normal_plus_function_calling() - - -def test_token_counter_legacy_function_call_counts_arguments(): - """ - Regression for VERIA-492 (Token-counter function_call bypass). - - The legacy OpenAI assistant `function_call` field carries arbitrary text in - `arguments`. Before the fix, `_count_messages` had no branch for - `function_call` and fell through to the unsupported-key `continue`, so an - assistant turn could smuggle unlimited text past `token_counter` and the - proxy `/utils/token_counter` endpoint (and downstream pre-call budget / - `get_modified_max_tokens` math). After the fix it must be counted the - same as the equivalent `tool_calls` payload. - """ - long_arg = "A" * 4000 - fc_messages = [ - {"role": "user", "content": "hi"}, - { - "role": "assistant", - "content": None, - "function_call": {"name": "search", "arguments": long_arg}, - }, - ] - tc_messages = [ - {"role": "user", "content": "hi"}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_1", - "type": "function", - "function": {"name": "search", "arguments": long_arg}, - } - ], - }, - ] - fc_tokens = token_counter(model="gpt-3.5-turbo", messages=fc_messages) - tc_tokens = token_counter(model="gpt-3.5-turbo", messages=tc_messages) - assert fc_tokens == tc_tokens, ( - f"function_call arguments must count like tool_calls arguments; " - f"got function_call={fc_tokens}, tool_calls={tc_tokens}" - ) - assert fc_tokens > 500, f"4000-char arguments payload must contribute real tokens, got {fc_tokens}" - - -@pytest.mark.parametrize( - "message_count_pair", - MESSAGES_TEXT, -) -def test_token_counter_textonly(message_count_pair): - counted_tokens = token_counter( - model="gpt-35-turbo", messages=[message_count_pair["message"]] - ) - assert counted_tokens == message_count_pair["count"] - - -@pytest.mark.parametrize( - "message_count_pair", - MESSAGES_TEXT, -) -def test_token_counter_count_response_tokens(message_count_pair): - counted_tokens = token_counter( - model="gpt-35-turbo", - messages=[message_count_pair["message"]], - count_response_tokens=True, - ) - # 3 tokens are not added because of count_response_tokens=True - expected = message_count_pair["count"] - 3 - assert counted_tokens == expected - - -@pytest.mark.parametrize( - "message_count_pair", - MESSAGES_WITH_IMAGES, -) -def test_token_counter_with_images(message_count_pair): - counted_tokens = token_counter( - model="gpt-4o", messages=[message_count_pair["message"]] - ) - assert counted_tokens == message_count_pair["count"] - - -@pytest.mark.parametrize( - "message_count_pair", - MESSAGES_WITH_TOOLS, -) -def test_token_counter_with_tools(message_count_pair): - counted_tokens = token_counter( - model="gpt-35-turbo", - messages=[message_count_pair["system_message"]], - tools=message_count_pair["tools"], - tool_choice=message_count_pair["tool_choice"], - ) - expected_tokens = message_count_pair["count"] - actual_diff = counted_tokens - expected_tokens - - if "count-tolerate" in message_count_pair: - if message_count_pair["count-tolerate"] == counted_tokens: - pass # expected - else: - tolerated_diff = message_count_pair["count-tolerate"] - expected_tokens - assert ( - actual_diff <= tolerated_diff - ), f"Expected {expected_tokens} tokens, got {counted_tokens}. Counted tokens is only allowed to be off by {tolerated_diff} in the over-counting direction." - if actual_diff != tolerated_diff: - raise NeedsToleranceUpdateError( - f"SOMETHING BROKEN GOT FIXED! THIS is good! Adjust 'count-tolerate' from {message_count_pair['count-tolerate']} to {counted_tokens}" - ) - - else: - assert ( - expected_tokens == counted_tokens - ), f"Expected {expected_tokens} tokens, got {counted_tokens}." - - -class NeedsToleranceUpdateError(Exception): - """Custom exception to mark tests that have improved""" - - pass +from litellm import create_pretrained_tokenizer +from tests.unit.litellm_core_utils.test_token_counter import token_counter def test_tokenizers(): @@ -452,32 +13,22 @@ def test_tokenizers(): openai_tokens = token_counter(model="gpt-3.5-turbo", text=sample_text) # claude tokenizer - claude_tokens = token_counter( - model="claude-3-5-haiku-20241022", text=sample_text - ) + claude_tokens = token_counter(model="claude-3-5-haiku-20241022", text=sample_text) # cohere tokenizer cohere_tokens = token_counter(model="command-nightly", text=sample_text) # llama2 tokenizer - llama2_tokens = token_counter( - model="meta-llama/Llama-2-7b-chat", text=sample_text - ) + llama2_tokens = token_counter(model="meta-llama/Llama-2-7b-chat", text=sample_text) # llama3 tokenizer (also testing custom tokenizer) - llama3_tokens_1 = token_counter( - model="meta-llama/llama-3-70b-instruct", text=sample_text - ) + llama3_tokens_1 = token_counter(model="meta-llama/llama-3-70b-instruct", text=sample_text) try: llama3_tokenizer = create_pretrained_tokenizer("Xenova/llama-3-tokenizer") except Exception as e: - pytest.skip( - f"custom tokenizer download failed (HF hub unreachable): {e}" - ) - llama3_tokens_2 = token_counter( - custom_tokenizer=llama3_tokenizer, text=sample_text - ) + pytest.skip(f"custom tokenizer download failed (HF hub unreachable): {e}") + llama3_tokens_2 = token_counter(custom_tokenizer=llama3_tokenizer, text=sample_text) print( f"openai tokens: {openai_tokens}; claude tokens: {claude_tokens}; cohere tokens: {cohere_tokens}; llama2 tokens: {llama2_tokens}; llama3 tokens: {llama3_tokens_1}" @@ -488,1238 +39,13 @@ def test_tokenizers(): # model hub is unreachable (e.g. in CI). In that case the count will # equal the openai count and the differentiation assertion is skipped. if openai_tokens == llama2_tokens: - pytest.skip( - "llama2 fell back to tiktoken (HF hub unreachable); skipping differentiation assertion" - ) + pytest.skip("llama2 fell back to tiktoken (HF hub unreachable); skipping differentiation assertion") assert llama2_tokens != llama3_tokens_1, "Token values are not different." - assert ( - llama3_tokens_1 == llama3_tokens_2 - ), "Custom tokenizer is not being used! It has been configured to use the same tokenizer as the built in llama3 tokenizer and the results should be the same." + assert llama3_tokens_1 == llama3_tokens_2, ( + "Custom tokenizer is not being used! It has been configured to use the same tokenizer as the built in llama3 tokenizer and the results should be the same." + ) print("test tokenizer: It worked!") except Exception as e: pytest.fail(f"An exception occured: {e}") - - -# test_tokenizers() - - -def test_encoding_and_decoding(): - try: - sample_text = "Hellö World, this is my input string!" - # openai encoding + decoding - openai_tokens = encode(model="gpt-3.5-turbo", text=sample_text) - openai_text = decode(model="gpt-3.5-turbo", tokens=openai_tokens) - - assert openai_text == sample_text - - # claude encoding + decoding - claude_tokens = encode(model="claude-3-5-haiku-20241022", text=sample_text) - - claude_text = decode(model="claude-3-5-haiku-20241022", tokens=claude_tokens) - - assert claude_text == sample_text - - # cohere encoding + decoding - cohere_tokens = encode(model="command-nightly", text=sample_text) - cohere_text = decode(model="command-nightly", tokens=cohere_tokens) - - assert cohere_text == sample_text - - # llama2 encoding + decoding - llama2_tokens = encode(model="meta-llama/Llama-2-7b-chat", text=sample_text) - llama2_text = decode(model="meta-llama/Llama-2-7b-chat", tokens=llama2_tokens) - - assert llama2_text == sample_text - except Exception as e: - pytest.fail(f"An exception occured: {e}\n{traceback.format_exc()}") - - -# test_encoding_and_decoding() - - -def test_gpt_vision_token_counting(): - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What’s in this image?"}, - { - "type": "image_url", - "image_url": "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png", - }, - ], - } - ] - tokens = token_counter(model="gpt-4-vision-preview", messages=messages) - print(f"tokens: {tokens}") - - -# test_gpt_vision_token_counting() - - -@pytest.mark.parametrize( - "model", - [ - "gpt-4-vision-preview", - "gpt-4o", - "claude-3-opus-20240229", - "command-nightly", - "mistral/mistral-tiny", - ], -) -def test_load_test_token_counter(model): - """ - Token count large prompt 100 times. - - Assert time taken is < 1.5s. - """ - import tiktoken - - messages = [{"role": "user", "content": text}] * 10 - - start_time = time.time() - for _ in range(10): - _ = token_counter(model=model, messages=messages) - # enc.encode("".join(m["content"] for m in messages)) - - end_time = time.time() - - total_time = end_time - start_time - print("model={}, total test time={}".format(model, total_time)) - assert total_time < 10, f"Total encoding time > 10s, {total_time}" - - -def test_openai_token_with_image_and_text(): - model = "gpt-4o" - full_request = { - "model": "gpt-4o", - "tools": [ - { - "type": "function", - "function": { - "name": "json", - "parameters": { - "type": "object", - "required": ["clause"], - "properties": {"clause": {"type": "string"}}, - }, - "description": "Respond with a JSON object.", - }, - } - ], - "logprobs": False, - "messages": [ - { - "role": "user", - "content": [ - { - "text": "\n Just some long text, long long text, and you know it will be longer than 7 tokens definetly.", - "type": "text", - } - ], - } - ], - "tool_choice": {"type": "function", "function": {"name": "json"}}, - "exclude_models": [], - "disable_fallback": False, - "exclude_providers": [], - } - messages = full_request.get("messages", []) - - token_count = token_counter(model=model, messages=messages) - print(token_count) - - -@pytest.mark.parametrize( - "model, base_model, input_tokens, user_max_tokens, expected_value", - [ - ("random-model", "random-model", 1024, 1024, 1024), - ("gpt-3.5-turbo", "gpt-3.5-turbo", 4000, 5000, 4096), # model max output = 4096 - ], -) -def test_get_modified_max_tokens( - model, base_model, input_tokens, user_max_tokens, expected_value -): - """ - - Test when max_output is not known => expect user_max_tokens - - Test when max_output == max_input, - - input > max_output, no max_tokens => expect None - - input + max_tokens > max_output => expect remainder - - input + max_tokens < max_output => expect max_tokens - - Test when max_tokens > max_output => expect max_output - """ - args = locals() - import litellm - - litellm.token_counter = MagicMock() - - def _mock_token_counter(*args, **kwargs): - return input_tokens - - litellm.token_counter.side_effect = _mock_token_counter - print(f"_mock_token_counter: {_mock_token_counter()}") - messages = [{"role": "user", "content": "Hello world!"}] - - calculated_value = get_modified_max_tokens( - model=model, - base_model=base_model, - messages=messages, - user_max_tokens=user_max_tokens, - buffer_perc=0, - buffer_num=0, - ) - - if expected_value is None: - assert calculated_value is None - else: - assert ( - calculated_value == expected_value - ), "Got={}, Expected={}, Params={}".format( - calculated_value, expected_value, args - ) - - -def test_empty_tools(): - messages = [{"role": "user", "content": "hey, how's it going?", "tool_calls": None}] - - result = token_counter( - messages=messages, - ) - - print(result) - - -@pytest.mark.skip( - reason="Skipping this test temporarily because it relies on a function being called that I am removing." -) -def test_gpt_4o_token_counter(): - with patch.object( - litellm.utils, "openai_token_counter", new=MagicMock() - ) as mock_client: - token_counter( - model="gpt-4o-2024-05-13", messages=[{"role": "user", "content": "Hey!"}] - ) - - mock_client.assert_called() - - -@pytest.mark.parametrize( - "img_url", - [ - "https://example.com/test-image.png", - "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAL0AAAC9CAMAAADRCYwCAAAAh1BMVEX///8AAAD8/Pz5+fkEBAT39/cJCQn09PRNTU3y8vIMDAwzMzPe3t7v7+8QEBCOjo7FxcXR0dHn5+elpaWGhoYYGBivr686OjocHBy0tLQtLS1TU1PY2Ni6urpaWlpERER3d3ecnJxoaGiUlJRiYmIlJSU4ODhBQUFycnKAgIDBwcFnZ2chISE7EjuwAAAI/UlEQVR4nO1caXfiOgz1bhJIyAJhX1JoSzv8/9/3LNlpYd4rhX6o4/N8Z2lKM2cURZau5JsQEhERERERERERERERERERERHx/wBjhDPC3OGN8+Cc5JeMuheaETSdO8vZFyCScHtmz2CsktoeMn7rLM1u3h0PMAEhyYX7v/Q9wQvoGdB0hlbzm45lEq/wd6y6G9aezvBk9AXwp1r3LHJIRsh6s2maxaJpmvqgvkC7WFS3loUnaFJtKRVUCEoV/RpCnHRvAsesVQ1hw+vd7Mpo+424tLs72NplkvQgcdrsvXkW/zJWqH/fA0FT84M/xnQJt4to3+ZLuanbM6X5lfXKHosO9COgREqpCR5i86pf2zPS7j9tTj+9nO7bQz3+xGEyGW9zqgQ1tyQ/VsxEDvce/4dcUPNb5OD9yXvR4Z2QisuP0xiGWPnemgugU5q/troHhGEjIF5sTOyW648aC0TssuaaCEsYEIkGzjWXOp3A0vVsf6kgRyqaDk+T7DIVWrb58b2tT5xpUucKwodOD/5LbrZC1ws6YSaBZJ/8xlh+XZSYXaMJ2ezNqjB3IPXuehPcx2U6b4t1dS/xNdFzguUt8ie7arnPeyCZroxLHzGgGdqVcspwafizPWEXBee+9G1OaufGdvNng/9C+gwgZ3PH3r87G6zXTZ5D5De2G2DeFoANXfbACkT+fxBQ22YFsTTJF9hjFVO6VbqxZXko4WJ8s52P4PnuxO5KRzu0/hlix1ySt8iXjgaQ+4IHPA9nVzNkdduM9LFT/Aacj4FtKrHA7iAw602Vnht6R8Vq1IOS+wNMKLYqayAYfRuufQPGeGb7sZogQQoLZrGPgZ6KoYn70Iw30O92BNEDpvwouCFn6wH2uS+EhRb3WF/HObZk3HuxfRQM3Y/Of/VH0n4MKNHZDiZvO9+m/ABALfkOcuar/7nOo7B95ACGVAFaz4jMiJwJhdaHBkySmzlGTu82gr6FSTik2kJvLnY9nOd/D90qcH268m3I/cgI1xg1maE5CuZYaWLH+UHANCIck0yt7Mx5zBm5vVHXHwChsZ35kKqUpmo5Svq5/fzfAI5g2vDtFPYo1HiEA85QrDeGm9g//LG7K0scO3sdpj2CBDgCa+0OFs0bkvVgnnM/QBDwllOMm+cN7vMSHlB7Uu4haHKaTwgGkv8tlK+hP8fzmFuK/RQTpaLPWvbd58yWIo66HHM0OsPoPhVqmtaEVL7N+wYcTLTbb0DLdgp23Eyy2VYJ2N7bkLFAAibtoLPe5sLt6Oa2bvU+zyeMa8wrixO0gRTn9tO9NCSThTLGqcqtsDvphlfmx/cPBZVvw24jg1LE2lPuEo35Mhi58U0I/Ga8n5w+NS8i34MAQLos5B1u0xL1ZvCVYVRw/Fs2q53KLaXJMWwOZZ/4MPYV19bAHmgGDKB6f01xoeJKFbl63q9J34KdaVNPJWztQyRkzA3KNs1AdAEDowMxh10emXTCx75CkurtbY/ZpdNDGdsn2UcHKHsQ8Ai3WZi48IfkvtjOhsLpuIRSKZTX9FA4o+0d6o/zOWqQzVJMynL9NsxhSJOaourq6nBVQBueMSyubsX2xHrmuABZN2Ns9jr5nwLFlLF/2R6atjW/67Yd11YQ1Z+kA9Zk9dPTM/o6dVo6HHVgC0JR8oUfmI93T9u3gvTG94bAH02Y5xeqRcjuwnKCK6Q2+ajl8KXJ3GSh22P3Zfx6S+n008ROhJn+JRIUVu6o7OXl8w1SeyhuqNDwNI7SjbK08QrqPxS95jy4G7nCXVq6G3HNu0LtK5J0e226CfC005WKK9sVvfxI0eUbcnzutfhWe3rpZHM0nZ/ny/N8tanKYlQ6VEW5Xuym8yV1zZX58vwGhZp/5tFfhybZabdbrQYOs8F+xEhmPsb0/nki6kIyVvzZzUASiOrTfF+Sj9bXC7DoJxeiV8tjQL6loSd0yCx7YyB6rPdLx31U2qCG3F/oXIuDuqd6LFO+4DNIJuxFZqSsU0ea88avovFnWKRYFYRQDfCfcGaBCLn4M4A1ntJ5E57vicwqq2enaZEF5nokCYu9TbKqCC5yCDfL+GhLxT4w4xEJs+anqgou8DOY2q8FMryjb2MehC1dRJ9s4g9NXeTwPkWON4RH+FhIe0AWR/S9ekvQ+t70XHeimGF78LzuU7d7PwrswdIG2VpgF8C53qVQsTDtBJc4CdnkQPbnZY9mbPdDFra3PCXBBQ5QBn2aQqtyhvlyYM4Hb2/mdhsxCUen04GZVvIJZw5PAamMOmjzq8Q+dzAKLXDQ3RUZItWsg4t7W2DP+JDrJDymoMH7E5zQtuEpG03GTIjGCW3LQqOYEsXgFc78x76NeRwY6SNM+IfQoh6myJKRBIcLYxZcwscJ/gI2isTBty2Po9IkYzP0/SS4hGlxRjFAG5z1Jt1LckiB57yWvo35EaolbvA+6fBa24xodL2YjsPpTnj3JgJOqhcgOeLVsYYwoK0wjY+m1D3rGc40CukkaHnkEjarlXrF1B9M6ECQ6Ow0V7R7N4G3LfOHAXtymoyXOb4QhaYHJ/gNBJUkxclpSs7DNcgWWDDmM7Ke5MJpGuioe7w5EOvfTunUKRzOh7G2ylL+6ynHrD54oQO3//cN3yVO+5qMVsPZq0CZIOx4TlcJ8+Vz7V5waL+7WekzUpRFMTnnTlSCq3X5usi8qmIleW/rit1+oQZn1WGSU/sKBYEqMNh1mBOc6PhK8yCfKHdUNQk8o/G19ZPTs5MYfai+DLs5vmee37zEyyH48WW3XA6Xw6+Az8lMhci7N/KleToo7PtTKm+RA887Kqc6E9dyqL/QPTugzMHLbLZtJKqKLFfzVWRNJ63c+95uWT/F7R0U5dDVvuS409AJXhJvD0EwWaWdW8UN11u/7+umaYjT8mJtzZwP/MD4r57fihiHlC5fylHfaqnJdro+Dr7DajvO+vi2EwyD70s8nCH71nzIO1l5Zl+v1DMCb5ebvCMkGHvobXy/hPumGLyX0218/3RyD1GRLOuf9u/OGQyDmto32yMiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIv7GP8YjWPR/czH2AAAAAElFTkSuQmCC", - ], -) -def test_img_url_token_counter(img_url, monkeypatch): - """ - Verify get_image_dimensions returns valid (width, height) for both an - HTTPS URL and a base64 data URI. The HTTPS branch is exercised with a - mocked HTTP fetch so the test is hermetic - it can't break when a - third-party image URL goes away. - """ - import base64 - from litellm.litellm_core_utils.token_counter import get_image_dimensions - - # Minimal valid 1x1 PNG, served by the mocked safe_get for the URL case. - _tiny_png = base64.b64decode( - "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII=" - ) - - if img_url.startswith(("http://", "https://")): - - class _FakeResponse: - headers = {"Content-Length": str(len(_tiny_png))} - - def read(self): - return _tiny_png - - monkeypatch.setattr( - "litellm.litellm_core_utils.token_counter.safe_get", - lambda client, url, **kw: _FakeResponse(), - ) - - width, height = get_image_dimensions(data=img_url) - - print(width, height) - - assert width is not None - assert height is not None - - -def test_token_encode_disallowed_special(): - encode(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>") - token_counter(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>") - - -def test_token_counter(): - try: - messages = [{"role": "user", "content": "hi how are you what time is it"}] - tokens = token_counter(model="gpt-3.5-turbo", messages=messages) - print("gpt-35-turbo") - print(tokens) - assert tokens > 0 - - tokens = token_counter(model="claude-2", messages=messages) - print("claude-2") - print(tokens) - assert tokens > 0 - - tokens = token_counter(model="gemini/chat-bison", messages=messages) - print("gemini/chat-bison") - print(tokens) - assert tokens > 0 - - tokens = token_counter(model="ollama/llama2", messages=messages) - print("ollama/llama2") - print(tokens) - assert tokens > 0 - - tokens = token_counter(model="anthropic.claude-instant-v1", messages=messages) - print("anthropic.claude-instant-v1") - print(tokens) - assert tokens > 0 - except Exception as e: - pytest.fail(f"Error occurred: {e}") - - -import unittest - -from litellm.utils import _load_huggingface_tokenizer, _select_tokenizer_helper, claude_json_str, encoding - -# Clear the cache at module load to ensure clean state -_load_huggingface_tokenizer.cache_clear() - - -class TestTokenizerSelection(unittest.TestCase): - def setUp(self): - """Clear the LRU cache before each test method. - - The HuggingFace tokenizers behind _select_tokenizer_helper are cached with - @lru_cache, which can cause cache hits from previous tests when running with - --dist=loadscope (tests from same file run on same worker). - """ - _load_huggingface_tokenizer.cache_clear() - - @patch("litellm.utils.tokenizer_dispatch.from_pretrained") - def test_llama3_tokenizer_api_failure(self, mock_from_pretrained): - # Setup mock to raise an error - mock_from_pretrained.side_effect = Exception("Failed to load tokenizer") - - # Test with llama-3 model - result = _select_tokenizer_helper("llama-3-7b") - - # Verify the attempt to load Llama-3 tokenizer - mock_from_pretrained.assert_called_once_with("Xenova/llama-3-tokenizer") - - # Verify fallback to OpenAI tokenizer - self.assertEqual(result["type"], "openai_tokenizer") - self.assertEqual(result["tokenizer"], encoding) - - @patch("litellm.utils.tokenizer_dispatch.from_pretrained") - def test_cohere_tokenizer_api_failure(self, mock_from_pretrained): - # Setup mock to raise an error - mock_from_pretrained.side_effect = Exception("Failed to load tokenizer") - - # Add Cohere model to the list for testing - litellm.cohere_models = ["command-r-v1"] - - # Test with Cohere model - result = _select_tokenizer_helper("command-r-v1") - - # Verify the attempt to load Cohere tokenizer - mock_from_pretrained.assert_called_once_with( - "Xenova/c4ai-command-r-v01-tokenizer" - ) - - # Verify fallback to OpenAI tokenizer - self.assertEqual(result["type"], "openai_tokenizer") - self.assertEqual(result["tokenizer"], encoding) - - @patch("litellm.utils.tokenizer_dispatch.anthropic") - def test_claude_tokenizer_api_failure(self, mock_anthropic): - # Setup mock to raise an error - mock_anthropic.side_effect = Exception("Failed to load tokenizer") - - # Add Claude model to the list for testing - litellm.anthropic_models = ["claude-2"] - - # Test with Claude model - result = _select_tokenizer_helper("claude-2") - - # Verify the attempt to load Claude tokenizer - mock_anthropic.assert_called_once_with() - - # Verify fallback to OpenAI tokenizer - self.assertEqual(result["type"], "openai_tokenizer") - self.assertEqual(result["tokenizer"], encoding) - - @patch("litellm.utils.tokenizer_dispatch.from_pretrained") - def test_llama2_tokenizer_api_failure(self, mock_from_pretrained): - # Setup mock to raise an error - mock_from_pretrained.side_effect = Exception("Failed to load tokenizer") - - # Test with Llama-2 model - result = _select_tokenizer_helper("llama-2-7b") - - # Verify the attempt to load Llama-2 tokenizer - mock_from_pretrained.assert_called_once_with( - "hf-internal-testing/llama-tokenizer" - ) - - # Verify fallback to OpenAI tokenizer - self.assertEqual(result["type"], "openai_tokenizer") - self.assertEqual(result["tokenizer"], encoding) - - @patch("litellm.utils._return_huggingface_tokenizer") - def test_disable_hf_tokenizer_download(self, mock_return_huggingface_tokenizer): - monkeypatch = pytest.MonkeyPatch() - monkeypatch.setattr(litellm, "disable_hf_tokenizer_download", True) - try: - result = _select_tokenizer_helper("grok-32r22r") - mock_return_huggingface_tokenizer.assert_not_called() - assert result["type"] == "openai_tokenizer" - assert result["tokenizer"] == encoding - finally: - monkeypatch.undo() - - -@pytest.mark.parametrize( - "model", - [ - "gpt-4o", - "claude-3-opus-20240229", - ], -) -@pytest.mark.parametrize( - "messages", - [ - [ - { - "role": "user", - "content": [ - { - "type": "text", - "text": "These are some sample images from a movie. Based on these images, what do you think the tone of the movie is?", - }, - { - "type": "text", - "image_url": { - "url": "https://gratisography.com/wp-content/uploads/2024/11/gratisography-augmented-reality-800x525.jpg", - "detail": "high", - }, - }, - ], - } - ], - ], -) -def test_bad_input_token_counter(model, messages): - """ - Safely handle bad input for token counter. - """ - token_counter( - model=model, - messages=messages, - default_token_count=1000, - ) - - -def test_token_counter_with_anthropic_tool_use(): - """ - Test that _count_anthropic_content() correctly handles tool_use blocks. - - Validates that: - - 'name' field is counted (string) - - 'input' field is counted (dict serialized to string) - - Metadata fields ('type', 'id') are skipped - """ - messages = [ - {"role": "user", "content": "What's the weather in San Francisco?"}, - { - "role": "assistant", - "content": [ - {"type": "text", "text": "I'll check the weather for you."}, - { - "type": "tool_use", - "id": "toolu_01234567890", # Should be skipped - "name": "get_weather", # Should be counted - "input": { # Should be counted (serialized) - "location": "San Francisco, CA", - "unit": "fahrenheit", - }, - }, - ], - }, - ] - - tokens = token_counter(model="gpt-3.5-turbo", messages=messages) - assert tokens > 0, f"Expected positive token count, got {tokens}" - # Should count: user message + "I'll check" text + "get_weather" name + input dict - assert ( - tokens > 15 - ), f"Expected reasonable token count for message with tool_use, got {tokens}" - - -def test_token_counter_with_anthropic_tool_result(): - """ - Test that _count_anthropic_content() correctly handles tool_result blocks. - - Validates that: - - 'content' field (when string) is counted - - Metadata fields ('type', 'tool_use_id') are skipped - - Full conversation with tool_use → tool_result flow works - """ - messages = [ - {"role": "user", "content": "What's the weather in San Francisco?"}, - { - "role": "assistant", - "content": [ - { - "type": "tool_use", - "id": "toolu_01234567890", - "name": "get_weather", - "input": {"location": "San Francisco, CA"}, - } - ], - }, - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "toolu_01234567890", # Should be skipped - "content": "The weather in San Francisco is 65°F and sunny.", # Should be counted - } - ], - }, - ] - - tokens = token_counter(model="gpt-3.5-turbo", messages=messages) - assert tokens > 0, f"Expected positive token count, got {tokens}" - assert ( - tokens > 25 - ), f"Expected reasonable token count for conversation with tool_result, got {tokens}" - - -def test_token_counter_with_nested_tool_result(): - """ - Test that _count_anthropic_content() recursively handles nested content lists. - - Validates that: - - tool_result with 'content' as a list (not string) is handled - - Nested content blocks are recursively counted via _count_content_list() - - TypedDict inference correctly identifies list fields - """ - messages = [ - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "toolu_01234567890", - "content": [ # Nested list - should recursively count - { - "type": "text", - "text": "The weather in San Francisco is 65°F and sunny.", - }, - {"type": "text", "text": "UV index is moderate."}, - ], - } - ], - } - ] - - tokens = token_counter(model="gpt-3.5-turbo", messages=messages) - assert tokens > 0, f"Expected positive token count, got {tokens}" - # Should count both nested text blocks - assert ( - tokens > 15 - ), f"Expected reasonable token count for nested tool_result, got {tokens}" - - -def test_token_counter_tool_use_and_result_combined(): - """ - Test dynamic field inference with multiple tool_use and tool_result blocks. - - Validates that: - - Multiple tool_use blocks in same message are handled - - Multiple tool_result blocks in same message are handled - - skip_fields correctly filters metadata across all blocks - - Full realistic conversation flow works end-to-end - """ - messages = [ - { - "role": "user", - "content": "What's the weather in San Francisco and New York?", - }, - { - "role": "assistant", - "content": [ - { - "type": "text", - "text": "I'll check the weather in both cities for you.", - }, - { - "type": "tool_use", - "id": "toolu_01A", - "name": "get_weather", - "input": {"location": "San Francisco, CA"}, - }, - { - "type": "tool_use", - "id": "toolu_01B", - "name": "get_weather", - "input": {"location": "New York, NY"}, - }, - ], - }, - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "toolu_01A", - "content": "San Francisco: 65°F, sunny", - }, - { - "type": "tool_result", - "tool_use_id": "toolu_01B", - "content": "New York: 45°F, cloudy", - }, - ], - }, - { - "role": "assistant", - "content": "The weather in San Francisco is 65°F and sunny, while New York is cooler at 45°F and cloudy.", - }, - ] - - tokens = token_counter(model="gpt-3.5-turbo", messages=messages) - assert tokens > 0, f"Expected positive token count, got {tokens}" - # Should count all text, tool names, inputs, and results - assert ( - tokens > 60 - ), f"Expected substantial token count for full tool conversation, got {tokens}" - - -def test_token_counter_with_image_url(): - """ - Test that _count_image_tokens() correctly handles image_url content blocks. - - Validates that: - - image_url as dict with 'url' and 'detail' is handled - - image_url as string is handled - - 'detail' field validation works ('low', 'high', 'auto') - - calculate_img_tokens is called with correct parameters - """ - # Test with dict format (detail: low) - messages_dict = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What's in this image?"}, - { - "type": "image_url", - "image_url": { - "url": "https://example.com/image.jpg", - "detail": "low", # Should use low token count (85 base tokens) - }, - }, - ], - } - ] - - tokens_dict = token_counter( - model="gpt-3.5-turbo", - messages=messages_dict, - use_default_image_token_count=True, # Avoid actual HTTP request - ) - assert tokens_dict > 0, f"Expected positive token count, got {tokens_dict}" - assert tokens_dict > 85, f"Expected at least base image tokens, got {tokens_dict}" - - # Test with string format (defaults to auto/low) - messages_str = [ - { - "role": "user", - "content": [ - { - "type": "image_url", - "image_url": "https://example.com/image.jpg", # String format - } - ], - } - ] - - tokens_str = token_counter( - model="gpt-3.5-turbo", messages=messages_str, use_default_image_token_count=True - ) - assert ( - tokens_str > 0 - ), f"Expected positive token count for string image_url, got {tokens_str}" - - # Test invalid detail value raises error - messages_invalid = [ - { - "role": "user", - "content": [ - { - "type": "image_url", - "image_url": { - "url": "https://example.com/image.jpg", - "detail": "invalid", # Should raise ValueError - }, - } - ], - } - ] - - with pytest.raises(ValueError, match="Invalid detail value") as exc_info: - token_counter(model="gpt-3.5-turbo", messages=messages_invalid) - e = exc_info.value - assert "Invalid detail value" in str( - e - ), f"Expected detail validation error, got: {e}" - - -def test_token_counter_with_thinking_content(): - """ - Test that _count_content_list() correctly handles Claude's extended thinking content blocks. - - Validates that: - - 'thinking' content type is recognized and counted - - 'thinking' text field is counted - - 'signature' field is skipped (opaque signature blob) - - Full conversation with thinking blocks works - """ - messages = [ - { - "role": "user", - "content": [ - { - "type": "text", - "text": "Analyze this complex problem: who came first, chicken or egg", - } - ], - }, - { - "role": "assistant", - "content": [ - { - "type": "thinking", - "thinking": "This is actually a fascinating question that touches on philosophy, biology, and semantics. Let me break this down: The egg came first from an evolutionary biology perspective.", - "signature": "EqcLCkYICxgCKkCrqu6lP...", # Should be skipped - }, - { - "type": "text", - "text": "# The Chicken-or-Egg Question: A Multi-Layered Answer\n\n## **The Short Answer: The Egg Came First**", - }, - ], - }, - {"role": "user", "content": [{"type": "text", "text": "Thanks"}]}, - ] - - tokens = token_counter( - model="anthropic/claude-sonnet-4-5-20250929", messages=messages - ) - assert tokens > 0, f"Expected positive token count, got {tokens}" - # Should count: user message + thinking text + response text + "Thanks" - # The thinking text alone is ~30 tokens, plus other content should be > 50 total - assert ( - tokens > 50 - ), f"Expected substantial token count for message with thinking, got {tokens}" - - # Test that thinking block without 'thinking' field doesn't crash (edge case) - messages_no_thinking = [ - { - "role": "assistant", - "content": [ - { - "type": "thinking", - # No 'thinking' field - should count as 0 tokens - "signature": "EqcLCkYICxgCKkCrqu6lP...", - }, - {"type": "text", "text": "Response"}, - ], - } - ] - - tokens_no_thinking = token_counter( - model="anthropic/claude-sonnet-4-5-20250929", messages=messages_no_thinking - ) - assert ( - tokens_no_thinking > 0 - ), f"Expected positive token count even with empty thinking, got {tokens_no_thinking}" - # Should only count "Response" and message overhead - assert ( - tokens_no_thinking < 15 - ), f"Expected minimal token count for empty thinking block, got {tokens_no_thinking}" - - - -def test_token_counter_with_redacted_thinking_content(): - """ - A replayed redacted_thinking block (Anthropic redacted reasoning, or the /v1/messages bridge's stand-in - for a reasoning item with no summary) counts zero tokens for its encrypted payload, like a thinking - block with no text. It used to raise, which made is_prompt_caching_valid_prompt return False and the - prompt_caching pre-call check stop pinning the deployment that held the cached prefix. - """ - model = "anthropic/claude-sonnet-4-5-20250929" - reply = {"type": "text", "text": "Draw from the box labeled Mixed, because that label must be wrong."} - redacted_block = {"type": "redacted_thinking", "data": "EqQBCkYIBRgCKkBjZ2xhc3M" * 30} - user_turn = {"role": "user", "content": [{"type": "text", "text": "Which box do you draw from?"}]} - follow_up = {"role": "user", "content": [{"type": "text", "text": "Restate that in one sentence."}]} - - without_block = [user_turn, {"role": "assistant", "content": [reply]}, follow_up] - with_block = [user_turn, {"role": "assistant", "content": [redacted_block, reply]}, follow_up] - - assert token_counter(model=model, messages=with_block) == token_counter(model=model, messages=without_block) - -def test_token_counter_with_tool_reference_block(): - """ - Regression test: a message containing an Anthropic tool-search - `tool_reference` content block must NOT raise. - - Before the fix, token_counter raised - `Invalid content item type: tool_reference`. On the streaming - anthropic_messages proxy path this nulled response_cost and caused the - SpendLogs row to be dropped, silently undercounting cost. token_counter - must instead count the referenced tool name and return a positive count. - """ - messages = [ - { - "role": "assistant", - "content": [ - {"type": "text", "text": "Let me look up the right tool."}, - {"type": "tool_reference", "tool_name": "search_knowledge_base"}, - ], - } - ] - - # Must not raise, and must produce a positive token count. - tokens = token_counter_new( - model="anthropic/claude-sonnet-4-5-20250929", messages=messages - ) - assert tokens > 0, f"Expected positive token count, got {tokens}" - - # A tool_reference with no/empty tool_name must also be handled gracefully. - messages_empty = [ - { - "role": "assistant", - "content": [{"type": "tool_reference", "tool_name": ""}], - } - ] - tokens_empty = token_counter_new( - model="anthropic/claude-sonnet-4-5-20250929", messages=messages_empty - ) - assert tokens_empty >= 0 - - -def test_count_content_list_rejects_unknown_type(): - """ - An unrecognized content block type must raise, and the error message must - enumerate the supported types (including `tool_reference`). This pins the - catch-all contract so a future block type isn't silently dropped. - """ - from litellm.litellm_core_utils.token_counter import _count_content_list - - with pytest.raises(ValueError, match='Error getting number of tokens from content list: Invalid') as exc_info: - _count_content_list( - count_function=len, - content_list=[{"type": "totally_unknown_block"}], - use_default_image_token_count=False, - default_token_count=None, - ) - - message = str(exc_info.value) - assert "Invalid content item type: totally_unknown_block" in message - assert "tool_reference" in message - - -@pytest.mark.parametrize( - "source", - [ - {"type": "base64", "media_type": "image/png", "data": "iVBORw0KGgo="}, - {"type": "url", "url": "https://example.com/image.png"}, - {"type": "file", "file_id": "file-abc123"}, - ], - ids=["base64", "url", "file"], -) -def test_token_counter_with_anthropic_image_block(source: dict[str, str]): - """Anthropic `image` blocks must count for every source variant, not raise `Invalid content item type` (which the router's context-window pre-call check swallows into an unfiltered dispatch).""" - from litellm.constants import DEFAULT_IMAGE_TOKEN_COUNT - - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this image?"}, - {"type": "image", "source": source}, - ], - } - ] - - tokens = token_counter( - model="anthropic/claude-sonnet-4-5-20250929", - messages=messages, - use_default_image_token_count=True, - ) - assert tokens > DEFAULT_IMAGE_TOKEN_COUNT, ( - f"Expected the image block to contribute tokens, got {tokens}" - ) - - -def test_anthropic_image_block_matches_equivalent_image_url(): - """An Anthropic `image` block prices identically to the OpenAI `image_url` carrying the same bytes.""" - anthropic_messages = [ - { - "role": "user", - "content": [ - { - "type": "image", - "source": { - "type": "base64", - "media_type": "image/png", - "data": "iVBORw0KGgo=", - }, - } - ], - } - ] - openai_messages = [ - { - "role": "user", - "content": [ - { - "type": "image_url", - "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, - } - ], - } - ] - - anthropic_tokens = token_counter( - model="anthropic/claude-sonnet-4-5-20250929", messages=anthropic_messages - ) - openai_tokens = token_counter( - model="anthropic/claude-sonnet-4-5-20250929", messages=openai_messages - ) - assert anthropic_tokens == openai_tokens - - -def test_anthropic_image_block_nested_in_tool_result(): - """An `image` block nested in a `tool_result.content` list is counted through the same recursion.""" - messages = [ - { - "role": "user", - "content": [ - { - "type": "tool_result", - "tool_use_id": "toolu_01", - "content": [ - { - "type": "image", - "source": { - "type": "base64", - "media_type": "image/png", - "data": "iVBORw0KGgo=", - }, - } - ], - } - ], - } - ] - - tokens = token_counter( - model="anthropic/claude-sonnet-4-5-20250929", - messages=messages, - use_default_image_token_count=True, - ) - assert tokens > 0 - - -@pytest.mark.parametrize( - ("source", "expected"), - [ - ({"type": "base64", "media_type": "image/jpeg", "data": "/9j/4AAQ"}, "data:image/jpeg;base64,/9j/4AAQ"), - ({"type": "url", "url": "https://example.com/image.png"}, "https://example.com/image.png"), - ({"type": "file", "file_id": "file-abc123"}, ""), - ], - ids=["base64", "url", "file"], -) -def test_anthropic_image_source_resolves_to_what_the_image_pricer_reads(source: dict[str, str], expected: str): - """base64 sources become a data URI, url sources pass through, file sources resolve to an empty string.""" - from litellm.litellm_core_utils.token_counter import _anthropic_image_source_data - - assert _anthropic_image_source_data(source) == expected - - -def test_anthropic_image_block_with_empty_base64_data(): - """A base64 source with empty `data` prices as an image rather than raising.""" - from litellm.litellm_core_utils.token_counter import _count_content_list - - tokens = _count_content_list( - count_function=len, - content_list=[ - {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}} - ], - use_default_image_token_count=False, - default_token_count=None, - ) - assert tokens > 0 - - -def test_anthropic_image_block_without_source_raises(): - """An `image` block with no `source` raises, matching the OpenAI `image_url`-without-`url` behavior.""" - from litellm.litellm_core_utils.token_counter import _count_content_list - - with pytest.raises(ValueError, match="Error getting number of tokens from content list"): - _count_content_list( - count_function=len, - content_list=[{"type": "image"}], - use_default_image_token_count=False, - default_token_count=None, - ) - - # ... and `default_token_count`, the caller's opt-out from raising, still wins. - assert ( - _count_content_list( - count_function=len, - content_list=[{"type": "image"}], - use_default_image_token_count=False, - default_token_count=7, - ) - == 7 - ) - - -def _count_user_content(content: list[dict]) -> int: - from litellm.litellm_core_utils.token_counter import token_counter - - return token_counter( - model="anthropic/claude-fable-5", - messages=[{"role": "user", "content": content}], - use_default_image_token_count=True, - ) - - -@pytest.mark.parametrize( - "source", - [ - {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"}, - {"type": "url", "url": "https://example.com/report.pdf"}, - {"type": "file", "file_id": "file-abc123"}, - ], - ids=["base64", "url", "file"], -) -def test_anthropic_document_block_with_opaque_source_is_priced_like_an_image(source: dict[str, str]): - """A `document` whose bytes can't be tokenized locally is priced like an `image`, not raised on.""" - prompt = {"type": "text", "text": "Summarize this file."} - - assert _count_user_content([prompt, {"type": "document", "source": source}]) == _count_user_content( - [prompt, {"type": "image", "source": source}] - ) - - -def test_anthropic_document_block_text_sources_count_their_text(): - """`text` and `content` document sources count the text they carry, as inline text blocks would.""" - prompt = {"type": "text", "text": "Summarize this file."} - body = {"type": "text", "text": "Revenue grew eleven percent while churn fell to two percent."} - picture = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "iVBORw0KGgo="}} - - text_source = {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": body["text"]}} - assert _count_user_content([prompt, text_source]) == _count_user_content([prompt, body]) - - string_content = {"type": "document", "source": {"type": "content", "content": body["text"]}} - assert _count_user_content([prompt, string_content]) == _count_user_content([prompt, body]) - - block_content = {"type": "document", "source": {"type": "content", "content": [body, picture]}} - assert _count_user_content([prompt, block_content]) == _count_user_content([prompt, body, picture]) - - -def test_anthropic_document_title_and_context_add_their_tokens(): - prompt = {"type": "text", "text": "Summarize this file."} - source = {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"} - described = {"type": "document", "source": source, "title": "Q3 board packet", "context": "Shared by finance"} - - assert _count_user_content([prompt, described]) == _count_user_content( - [ - prompt, - {"type": "text", "text": "Q3 board packet"}, - {"type": "text", "text": "Shared by finance"}, - {"type": "document", "source": source}, - ] - ) - - -def test_openai_file_block_prices_like_the_equivalent_anthropic_document(): - """An inline `file` is a `document` in the chat-completions dialect, so it must price identically, not raise. - - Before the fix `file` was missing from the content-block match even though `ChatCompletionFileObject` - is in the union this counter accepts, so every local count of a Responses `input_file` raised - `Invalid content item type: file` and surfaced as a 500 on /v1/responses/input_tokens. - """ - prompt = {"type": "text", "text": "Summarize this file."} - inline_file = { - "type": "file", - "file": {"filename": "report.pdf", "file_data": "data:application/pdf;base64,JVBERi0xLjQK"}, - } - document = { - "type": "document", - "title": "report.pdf", - "source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"}, - } - - assert _count_user_content([prompt, inline_file]) == _count_user_content([prompt, document]) - assert _count_user_content([prompt, inline_file]) > _count_user_content([prompt]) - - -def test_openai_file_block_without_inline_bytes_counts_what_it_carries(): - """A `file` block naming an uploaded file has no bytes to price, so it adds only the filename's tokens.""" - prompt = {"type": "text", "text": "Summarize this file."} - - by_id = {"type": "file", "file": {"file_id": "file-abc123"}} - assert _count_user_content([prompt, by_id]) == _count_user_content([prompt]) - - named = {"type": "file", "file": {"file_id": "file-abc123", "filename": "report.pdf"}} - assert _count_user_content([prompt, named]) == _count_user_content( - [prompt, {"type": "text", "text": "report.pdf"}] - ) - - -def _png_data_url(width: int, height: int) -> str: - ihdr = b"\x89PNG\r\n\x1a\n" + (13).to_bytes(4, "big") + b"IHDR" + width.to_bytes(4, "big") + height.to_bytes(4, "big") - return "data:image/png;base64," + base64.b64encode(ihdr + b"\x08\x06\x00\x00\x00").decode() - - -@pytest.mark.parametrize(("width", "height"), [(1, 1), (768, 768), (2000, 768), (768, 2000), (4096, 4096), (8000, 3072)]) -def test_high_detail_image_token_upper_bound_covers_every_image_size(width: int, height: int) -> None: - assert calculate_img_tokens(_png_data_url(width, height), mode="high") <= high_detail_image_token_upper_bound() - - -def test_high_detail_image_token_upper_bound_is_reached_by_the_largest_high_res_image() -> None: - assert calculate_img_tokens(_png_data_url(2000, 768), mode="high") == high_detail_image_token_upper_bound() - assert calculate_img_tokens(_png_data_url(1, 1), mode="high") < high_detail_image_token_upper_bound() - - -def _png_bytes(width: int, height: int) -> bytes: - return ( - b"\x89PNG\r\n\x1a\n" - + (13).to_bytes(4, "big") - + b"IHDR" - + struct.pack(">II", width, height) - + b"\x08\x02\x00\x00\x00" - ) - - -def _gif_bytes(width: int, height: int) -> bytes: - return b"GIF89a" + struct.pack(" bytes: - app: Final = b"".join( - b"\xff\xe0" + struct.pack(">H", 16) + b"JFIF\x00\x01\x01\x00\x00\x01\x00\x01\x00\x00" - for _ in range(app_segments) - ) - sof: Final = ( - b"\xff" + sof_marker + struct.pack(">HBHHB", 17, 8, height, width, 3) + b"\x01\x22\x00\x02\x11\x01\x03\x11\x01" - ) - return b"\xff\xd8" + app + sof - - -def _webp_bytes(chunk: bytes, payload: bytes) -> bytes: - body: Final = chunk + struct.pack(" bytes: - return _webp_bytes(b"VP8 ", b"\x00\x00\x00\x9d\x01\x2a" + struct.pack(" bytes: - return _webp_bytes(b"VP8L", b"\x2f" + struct.pack(" bytes: - return _webp_bytes(b"VP8X", b"\x00" * 4 + (width - 1).to_bytes(3, "little") + (height - 1).to_bytes(3, "little")) - - -@pytest.mark.parametrize( - ("image", "expected"), - [ - pytest.param(_png_bytes(1024, 768), (1024, 768), id="png"), - pytest.param(_gif_bytes(100, 50), (100, 50), id="gif"), - pytest.param(_jpeg_bytes(800, 600, b"\xc0", 1), (800, 600), id="jpeg-baseline"), - pytest.param(_jpeg_bytes(640, 480, b"\xc2", 3), (640, 480), id="jpeg-progressive-after-app-segments"), - pytest.param(_webp_vp8_bytes(640, 480), (640, 480), id="webp-vp8"), - pytest.param(_webp_vp8l_bytes(320, 240), (320, 240), id="webp-vp8l"), - pytest.param(_webp_vp8x_bytes(1920, 1080), (1920, 1080), id="webp-vp8x"), - ], -) -def test_image_dimensions_from_bytes_reads_each_header_format(image: bytes, expected: tuple[int, int]) -> None: - assert image_dimensions_from_bytes(image) == expected - - -@pytest.mark.parametrize( - "image", - [ - pytest.param(b"", id="empty"), - pytest.param(b"BM" + b"\x00" * 30, id="unknown-format"), - pytest.param(_webp_bytes(b"ALPH", b"\x00" * 16), id="webp-without-an-image-chunk"), - pytest.param(b"\x89PNG\r\n\x1a\n\x00\x00", id="png-truncated-before-ihdr"), - pytest.param(b"\xff\xd8\xff\xe0\x00\x10JFIF", id="jpeg-truncated-inside-app0"), - pytest.param(b"\xff\xd8\xff\xe0\x00\x04\x00\x00", id="jpeg-ends-before-sof"), - ], -) -def test_image_dimensions_from_bytes_returns_none_for_unreadable_headers(image: bytes) -> None: - assert image_dimensions_from_bytes(image) is None - - -def _jpeg_sof(width: int, height: int) -> bytes: - return b"\xff\xc0" + struct.pack(">HBHHB", 17, 8, height, width, 3) + b"\x01\x22\x00\x02\x11\x01\x03\x11\x01" - - -@pytest.mark.parametrize( - "image", - [ - pytest.param(b"\xff\xd8" + b"\xff\xe0\x00\x02" * 1025 + _jpeg_sof(800, 600), id="too-many-segments"), - pytest.param(b"\xff\xd8" + b"\xff" * 2000 + _jpeg_sof(800, 600)[1:], id="too-many-fill-bytes"), - pytest.param(b"\xff\xd8\xff\xe0\x00\x00\x02" + _jpeg_sof(800, 600), id="segment-length-below-two"), - ], -) -def test_image_dimensions_from_bytes_gives_up_on_pathological_jpeg_headers(image: bytes) -> None: - assert image_dimensions_from_bytes(image) is None - - -def test_image_dimensions_from_bytes_still_reads_a_jpeg_with_many_real_segments() -> None: - image: Final = b"\xff\xd8" + b"\xff\xe0\x00\x02" * 1000 + b"\xff" * 64 + _jpeg_sof(800, 600)[1:] - - assert image_dimensions_from_bytes(image) == (800, 600) - - -@pytest.mark.parametrize( - "header", - [ - pytest.param(b"\x89PNG\r\n\x1a\n\x00\x00", id="png-truncated"), - pytest.param(b"\xff\xd8\xff\xe0\x00\x10JFIF", id="jpeg-truncated"), - ], -) -def test_get_image_dimensions_still_raises_for_a_truncated_header(header: bytes) -> None: - with pytest.raises((struct.error, TypeError)): - get_image_dimensions(data="data:image/png;base64," + base64.b64encode(header).decode()) - - -@pytest.mark.parametrize( - "image", - [ - pytest.param(b"BM" + b"\x00" * 30, id="unknown-format"), - pytest.param(b"\xff\xd8" + b"\xff\xe0\x00\x02" * 1025 + _jpeg_sof(800, 600), id="pathological-jpeg"), - ], -) -def test_get_image_dimensions_falls_back_to_the_default_size_for_a_header_it_cannot_read(image: bytes) -> None: - assert get_image_dimensions(data="data:image/png;base64," + base64.b64encode(image).decode()) == ( - litellm.constants.DEFAULT_IMAGE_WIDTH, - litellm.constants.DEFAULT_IMAGE_HEIGHT, - ) diff --git a/tests/test_litellm/litellm_core_utils/test_tokenizer.py b/tests/test_litellm/litellm_core_utils/test_tokenizer.py index aa4a0fc6a1c..2171044970c 100644 --- a/tests/test_litellm/litellm_core_utils/test_tokenizer.py +++ b/tests/test_litellm/litellm_core_utils/test_tokenizer.py @@ -1,403 +1,20 @@ -import copy -import os -import pickle -import subprocess -import sys -from pathlib import Path -from typing import Final, Literal - import pytest -import tiktoken -from tokenizers import Tokenizer as ReferenceTokenizer -import litellm -from litellm.caching._embedding_router import truncate_embedding_input -from litellm.litellm_core_utils.tokenizer import HuggingFaceTokenizer, OpenAIEncoding -from litellm.utils import claude_json_str -from tests.test_litellm.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON - - -@pytest.mark.parametrize( - "name", ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "r50k_base", "gpt2", "o200k_harmony") -) -@pytest.mark.parametrize( - "text", ("hello world", "café 漢字 🙂", "", "a\ud800b", "\ud83d\ude42", "🙂\ud83d\ude42\udfff", " " * 64) +from tests.unit.litellm_core_utils.test_tokenizer import ( + UNICODE_TEXTS, + assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface, + assert_openai_encoding_matches_python, ) + +NETWORK_ENCODINGS = ("r50k_base", "gpt2") + + +@pytest.mark.parametrize("name", NETWORK_ENCODINGS) +@pytest.mark.parametrize("text", UNICODE_TEXTS) def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None: - reference: Final = tiktoken.get_encoding(name) - encoding: Final = OpenAIEncoding.from_tiktoken(name) - expected: Final = reference.encode(text) - - assert encoding.encode(text) == expected - assert encoding.count(text) == len(expected) - assert encoding.encode_batch([text], num_threads=2) == reference.encode_batch([text], num_threads=2) - assert encoding.encode_ordinary_batch([text]) == reference.encode_ordinary_batch([text]) - assert encoding.decode_batch([expected]) == reference.decode_batch([expected]) - assert encoding.decode_bytes_batch([expected]) == reference.decode_bytes_batch([expected]) + assert_openai_encoding_matches_python(name, text) -@pytest.mark.parametrize("allowed", (frozenset(), frozenset({"<|endoftext|>"}), "all")) -@pytest.mark.parametrize("disallowed", (frozenset(), frozenset({"<|fim_prefix|>"}), "all")) -def test_openai_special_token_options_match_python( - allowed: frozenset[str] | Literal["all"], disallowed: frozenset[str] | Literal["all"] -) -> None: - reference: Final = tiktoken.get_encoding("cl100k_base") - encoding: Final = OpenAIEncoding.from_tiktoken(reference.name) - text: Final = "hello<|endoftext|><|fim_prefix|>world" - allowed_set: Final = reference.special_tokens_set if allowed == "all" else allowed - disallowed_set: Final = reference.special_tokens_set - allowed_set if disallowed == "all" else disallowed - if any(token in text for token in disallowed_set): - with pytest.raises(ValueError, match="disallowed special token"): - encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed) - return - assert encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed) == reference.encode( - text, allowed_special=allowed, disallowed_special=disallowed - ) - assert encoding.special_tokens_set == reference.special_tokens_set - assert encoding.eot_token == reference.eot_token - - -@pytest.mark.parametrize("errors", ("replace", "ignore", "backslashreplace", "strict")) -def test_openai_partial_token_decoding_preserves_error_policy(errors: str) -> None: - reference: Final = tiktoken.get_encoding("cl100k_base") - encoding: Final = OpenAIEncoding.from_tiktoken(reference.name) - tokens: Final = reference.encode("🙂")[:1] - assert encoding.decode_bytes(tokens) == reference.decode_bytes(tokens) - if errors == "strict": - with pytest.raises(UnicodeDecodeError): - encoding.decode(tokens, errors=errors) - return - assert encoding.decode(tokens, errors=errors) == reference.decode(tokens, errors=errors) - assert encoding.decode_tokens_bytes(tokens) == reference.decode_tokens_bytes(tokens) - - -def test_public_encoding_and_semantic_cache_preserve_truncated_unicode() -> None: - reference: Final = tiktoken.get_encoding(litellm.encoding.name) - text: Final = "🙂" - tokens: Final = reference.encode(text) - - assert litellm.encoding.encode(text, disallowed_special=()) == tokens - assert litellm.encoding.encode_batch([text]) == [tokens] - assert litellm.decode(tokens=tokens[:1]) == reference.decode(tokens[:1]) - assert truncate_embedding_input(text, "", 1) == reference.decode(tokens[:1]) - - -@pytest.mark.parametrize("add_special_tokens", (True, False)) -def test_huggingface_encoding_preserves_result_fields_and_serialization(add_special_tokens: bool) -> None: - reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) - tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON) - expected: Final = reference.encode("Hello World", add_special_tokens=add_special_tokens) - actual: Final = tokenizer.encode("Hello World", add_special_tokens=add_special_tokens) - - assert (actual.ids, actual.tokens, actual.type_ids, actual.offsets, actual.word_ids, actual.sequence_ids) == ( - expected.ids, - expected.tokens, - expected.type_ids, - expected.offsets, - expected.word_ids, - expected.sequence_ids, - ) - assert (actual.attention_mask, actual.special_tokens_mask, actual.n_sequences, len(actual)) == ( - expected.attention_mask, - expected.special_tokens_mask, - expected.n_sequences, - len(expected), - ) - assert copy.deepcopy(actual).ids == expected.ids - assert pickle.loads(pickle.dumps(actual)).offsets == expected.offsets - assert tokenizer.decode(actual.ids, skip_special_tokens=False) == reference.decode( - expected.ids, skip_special_tokens=False - ) - - -def test_huggingface_character_offsets_and_pretokenized_pairs_match_python() -> None: - reference: Final = ReferenceTokenizer.from_str(claude_json_str) - tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str) - text: Final = "café 漢字 🙂" - actual: Final = tokenizer.encode(text) - expected: Final = reference.encode(text) - - assert actual.offsets == expected.offsets - assert actual.ids == expected.ids - assert ( - tokenizer.encode(["hello", "world"], ["again"], is_pretokenized=True).ids - == reference.encode(["hello", "world"], ["again"], is_pretokenized=True).ids - ) - - -def test_huggingface_batches_apply_padding_across_inputs() -> None: - reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) - reference.enable_padding(pad_id=0, pad_token="[UNK]") - tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str()) - inputs: Final = ["Hello", ("Hello World", "World")] - expected: Final = reference.encode_batch(inputs) - actual: Final = tokenizer.encode_batch(inputs) - fast: Final = tokenizer.encode_batch_fast(inputs) - - assert [(item.ids, item.attention_mask, item.offsets) for item in actual] == [ - (item.ids, item.attention_mask, item.offsets) for item in expected - ] - assert [item.ids for item in fast] == [item.ids for item in expected] - assert tokenizer.decode_batch([item.ids for item in actual]) == reference.decode_batch( - [item.ids for item in expected] - ) - - -def test_caller_supplied_huggingface_tokenizer_preserves_public_encode_and_count() -> None: - tokenizer: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) - custom: Final = {"type": "huggingface_tokenizer", "tokenizer": tokenizer} - expected: Final = tokenizer.encode("Hello World").ids - - assert litellm.encode(text="Hello World", custom_tokenizer=custom) == expected - assert litellm.token_counter(text="Hello World", custom_tokenizer=custom) == len(expected) - assert litellm.decode(tokens=expected, custom_tokenizer=custom) == "Hello World" - - -def test_caller_supplied_tiktoken_treats_special_spellings_as_text() -> None: - tokenizer: Final = tiktoken.get_encoding("cl100k_base") - custom: Final = {"type": "openai_tokenizer", "tokenizer": tokenizer} - text: Final = "<|endoftext|>" - - assert litellm.encode(text=text, custom_tokenizer=custom) == tokenizer.encode(text, disallowed_special=()) - - -def test_public_tokenizer_objects_survive_pickle_and_deepcopy(tmp_path: Path) -> None: - custom: Final = litellm.create_tokenizer(TOKENIZER_JSON) - tokenizer: Final = custom["tokenizer"] - path: Final = tmp_path / "tokenizer.json" - tokenizer.save(str(path)) - - assert copy.deepcopy(custom)["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids - assert ( - pickle.loads(pickle.dumps(custom))["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids - ) - assert HuggingFaceTokenizer.from_file(str(path)).encode("Hello World").ids == tokenizer.encode("Hello World").ids - assert copy.deepcopy(litellm.encoding).encode("hello") == litellm.encoding.encode("hello") - assert pickle.loads(pickle.dumps(litellm.encoding)).encode("hello") == litellm.encoding.encode("hello") - - -@pytest.mark.parametrize("offline", ("0", "1")) -def test_hub_loader_preserves_environment_auth_cache_and_offline(tmp_path: Path, offline: str) -> None: - script: Final = """ -import json -import sys -from pathlib import Path -sys.path.insert(0, sys.argv[1]) -import httpx -import huggingface_hub -from huggingface_hub.errors import LocalEntryNotFoundError -import litellm -payload = sys.argv[2].encode() -offline = sys.argv[3] == "1" -observed = [] -def handle(request): - assert not offline, "offline loading issued a request" - if request.url.path.endswith("/tokenizer.json"): - observed.append(request.headers.get("authorization")) - if request.headers.get("authorization") != "Bearer audit-fixture-token": - return httpx.Response(401) - return httpx.Response(200, headers={"content-length": str(len(payload)), "etag": '"fixture"', "x-repo-commit": "a" * 40}, content=payload if request.method == "GET" else b"") -if not offline: - huggingface_hub.set_client_factory(lambda: httpx.Client(transport=httpx.MockTransport(handle))) -try: - tokenizer = litellm.create_pretrained_tokenizer("test-fixture/tokenizer")["tokenizer"] -except LocalEntryNotFoundError: - assert offline - assert observed == [] -else: - assert not offline - assert "Bearer audit-fixture-token" in observed - assert tokenizer.decode(tokenizer.encode("Hello World").ids) == "Hello World" - assert tuple(Path(sys.argv[4]).rglob("tokenizer.json")) -print("compatible") -""" - result: Final = subprocess.run( - [ - sys.executable, - "-I", - "-c", - script, - str(Path(litellm.__file__).parent.parent), - TOKENIZER_JSON, - offline, - str(tmp_path / "cache"), - ], - capture_output=True, - text=True, - timeout=30, - env={ - **os.environ, - "HF_HOME": str(tmp_path / "home"), - "HF_HUB_CACHE": str(tmp_path / "cache"), - "HF_ENDPOINT": "http://127.0.0.1:9", - "HF_TOKEN": "audit-fixture-token", - "HF_HUB_OFFLINE": offline, - "HF_HUB_DISABLE_IMPLICIT_TOKEN": "0", - "LITELLM_LOCAL_MODEL_COST_MAP": "True", - }, - ) - assert result.returncode == 0, result.stdout + result.stderr - assert result.stdout.strip() == "compatible" - - -@pytest.mark.parametrize("rust", (None, "0", "1")) -def test_tokenization_without_native_extension_stays_offline(tmp_path: Path, rust: str | None) -> None: - script: Final = """ -import importlib.abc -import sys -sys.path.insert(0, sys.argv[1]) -def reject_network(event, args): - if event == "socket.connect": - raise AssertionError("tokenizer attempted a network connection") -sys.addaudithook(reject_network) -class Block(importlib.abc.MetaPathFinder): - def find_spec(self, fullname, path=None, target=None): - if fullname == "litellm.rust_bridge._native": - raise ImportError("native extension is unavailable") -sys.meta_path.insert(0, Block()) -import litellm -from litellm.rust_bridge.tokenizer import get_encoding -import tiktoken -from tokenizers import Tokenizer -assert isinstance(litellm.encoding, tiktoken.Encoding) -for name in ("cl100k_base", "o200k_base", "o200k_harmony", "p50k_base", "p50k_edit"): - encoding = get_encoding(name) - text = "offline café 漢字 🙂" + " " * 64 - assert encoding.decode(encoding.encode(text)) == text -ids = litellm.encode(text="hello world") -assert litellm.decode(tokens=ids) == "hello world" -assert litellm.token_counter(model=None, text="hello world") == len(ids) -custom = litellm.create_tokenizer(sys.argv[2]) -assert isinstance(custom["tokenizer"], Tokenizer) -custom["tokenizer"].enable_padding(pad_id=0, pad_token="[UNK]") -assert litellm.decode(tokens=litellm.encode(text="Hello World", custom_tokenizer=custom), custom_tokenizer=custom) == "Hello World" -print("compatible") -""" - result: Final = subprocess.run( - [sys.executable, "-I", "-c", script, str(Path(litellm.__file__).parent.parent), TOKENIZER_JSON], - capture_output=True, - text=True, - timeout=30, - cwd=tmp_path, - env={ - **{key: value for key, value in os.environ.items() if key != "LITELLM_RUST"}, - **({"LITELLM_RUST": rust} if rust is not None else {}), - "LITELLM_LOCAL_MODEL_COST_MAP": "True", - "TIKTOKEN_CACHE_DIR": str(tmp_path / "unused-tokenizer-cache"), - }, - ) - assert result.returncode == 0, result.stdout + result.stderr - assert result.stdout.strip() == "compatible" - assert not (tmp_path / "unused-tokenizer-cache").exists() - - -@pytest.mark.parametrize("is_pretokenized", (False, True)) -def test_huggingface_batch_sequence_containers_match_python(is_pretokenized: bool) -> None: - reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) - tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON) - inputs: Final = [["Hello", "World"], ("Hello", "World")] - actual: Final = tokenizer.encode_batch(inputs, is_pretokenized=is_pretokenized) - expected: Final = reference.encode_batch(inputs, is_pretokenized=is_pretokenized) - assert [(item.ids, item.type_ids, item.sequence_ids) for item in actual] == [ - (item.ids, item.type_ids, item.sequence_ids) for item in expected - ] - - -@pytest.mark.parametrize("name", ("cl100k_base", "o200k_base", "p50k_edit", "gpt2")) +@pytest.mark.parametrize("name", ("gpt2",)) def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: - reference: Final = tiktoken.get_encoding(name) - encoding: Final = OpenAIEncoding.from_tiktoken(name) - text: Final = "hello fanta" - - assert repr(encoding) == repr(reference) == f"" - assert (encoding.name, encoding.n_vocab, encoding.max_token_value) == ( - reference.name, - reference.n_vocab, - reference.max_token_value, - ) - assert encoding.token_byte_values() == reference.token_byte_values() - assert encoding.encode_single_token("hello") == reference.encode_single_token("hello") - assert encoding.encode_single_token(b"<|endoftext|>") == reference.eot_token - assert [encoding.is_special_token(token) for token in (0, reference.eot_token)] == [False, True] - assert encoding.decode_with_offsets(reference.encode(text)) == reference.decode_with_offsets(reference.encode(text)) - assert encoding.encode_to_numpy(text).tolist() == reference.encode_to_numpy(text).tolist() - stable, completions = encoding.encode_with_unstable(text) - expected_stable, expected_completions = reference.encode_with_unstable(text) - assert (stable, sorted(completions)) == (expected_stable, sorted(expected_completions)) - with pytest.raises(KeyError): - encoding.encode_single_token("<|not-a-token|>") - - -def test_huggingface_tokenizer_exposes_the_tokenizers_vocabulary_surface() -> None: - reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) - reference.enable_padding(pad_id=0, pad_token="[UNK]", length=4) - reference.enable_truncation(max_length=3, stride=1, strategy="only_first", direction="left") - tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str()) - - assert tokenizer.token_to_id("Hello") == reference.token_to_id("Hello") == 1 - assert tokenizer.id_to_token(3) == reference.id_to_token(3) == "[BOS]" - assert tokenizer.id_to_token(99) is None - assert tokenizer.get_vocab() == reference.get_vocab() - assert tokenizer.get_vocab(with_added_tokens=False) == reference.get_vocab(with_added_tokens=False) - assert tokenizer.get_vocab_size() == reference.get_vocab_size() == 4 - assert tokenizer.get_vocab_size(with_added_tokens=False) == reference.get_vocab_size(with_added_tokens=False) - added: Final = tokenizer.get_added_tokens_decoder() - expected_added: Final = reference.get_added_tokens_decoder() - assert {token_id: str(token) for token_id, token in added.items()} == { - token_id: str(token) for token_id, token in expected_added.items() - } - assert added[3].special == expected_added[3].special - assert tokenizer.num_special_tokens_to_add(False) == reference.num_special_tokens_to_add(False) == 1 - assert tokenizer.num_special_tokens_to_add(True) == reference.num_special_tokens_to_add(True) == 0 - assert tokenizer.padding == reference.padding - assert tokenizer.truncation == reference.truncation - assert tokenizer.encode_special_tokens == reference.encode_special_tokens is False - assert HuggingFaceTokenizer.from_buffer(TOKENIZER_JSON.encode()).encode("Hello").ids == [3, 1] - assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).padding is None - assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).truncation is None - - -def test_huggingface_encoding_exposes_the_tokenizers_lookup_and_mutation_surface() -> None: - reference: Final = ReferenceTokenizer.from_str(claude_json_str) - tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str) - text: Final = "hello wide world" - actual: Final = tokenizer.encode(text, "again") - expected: Final = reference.encode(text, "again") - - lookups: Final = ( - lambda encoding: [encoding.token_to_chars(index) for index in range(len(encoding))], - lambda encoding: [encoding.token_to_word(index) for index in range(len(encoding))], - lambda encoding: [encoding.token_to_sequence(index) for index in range(len(encoding))], - lambda encoding: [encoding.char_to_token(position) for position in range(len(text))], - lambda encoding: [encoding.char_to_word(position) for position in range(len(text))], - lambda encoding: [encoding.char_to_token(position, 1) for position in range(5)], - lambda encoding: [encoding.word_to_tokens(word) for word in range(3)], - lambda encoding: [encoding.word_to_chars(word) for word in range(3)], - lambda encoding: [encoding.word_to_tokens(0, 1), encoding.word_to_chars(0, 1)], - ) - for lookup in lookups: - assert lookup(actual) == lookup(expected) - assert repr(actual) == repr(expected) - - actual.truncate(4, stride=1, direction="left") - expected.truncate(4, stride=1, direction="left") - assert (actual.ids, [item.ids for item in actual.overflowing]) == ( - expected.ids, - [item.ids for item in expected.overflowing], - ) - actual.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="") - expected.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="") - assert (actual.ids, actual.attention_mask, actual.type_ids, actual.tokens) == ( - expected.ids, - expected.attention_mask, - expected.type_ids, - expected.tokens, - ) - actual.set_sequence_id(3) - expected.set_sequence_id(3) - assert actual.sequence_ids == expected.sequence_ids - merged: Final = type(actual).merge([actual, tokenizer.encode("more")]) - assert merged.ids == type(expected).merge([expected, reference.encode("more")]).ids - assert merged.offsets == type(expected).merge([expected, reference.encode("more")]).offsets - with pytest.raises(ValueError, match="direction"): - actual.pad(8, direction="sideways") + assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name) diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index f8b9043a23f..b1622e0dff0 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -5324,16 +5324,22 @@ def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deploym @pytest.mark.asyncio -async def test_resolve_team_from_header_defers_to_db_membership_only_without_jwt_claims(): +async def test_resolve_team_from_header_accepts_db_teams_provisionally_under_fallback_even_with_jwt_claims(): """With fallback_to_db_teams=True, an x-litellm-team-id header naming an existing - team is accepted provisionally only when the JWT carries no team claims (allowed - set empty). When the JWT does carry team claims, the header must still be validated - against them, and the flag-off behavior must keep rejecting unknown teams.""" + team is accepted provisionally whether or not the JWT carries team claims; the + union of JWT teams and DB memberships is enforced by auth_builder's later + membership check. Unknown values still 403, and the flag-off behavior keeps + rejecting teams outside the JWT's allowed set.""" known_ids = frozenset({"team-from-db"}) deferred, _, _ = await _resolve_header("team-from-db", set(), True, _teams_by_id(known_ids), _team_alias_lookup_404) assert deferred == HeaderTeam(header_value="team-from-db", team_id="team-from-db") + deferred_with_claims, _, _ = await _resolve_header( + "team-from-db", {"team-1"}, True, _teams_by_id(known_ids), _team_alias_lookup_404 + ) + assert deferred_with_claims == HeaderTeam(header_value="team-from-db", team_id="team-from-db") + with pytest.raises(HTTPException) as exc_info: await _resolve_header("team-x", {"team-1", "team-2"}, True, _teams_by_id(known_ids), _team_alias_lookup_404) assert exc_info.value.status_code == 403 @@ -5849,6 +5855,7 @@ async def _run_auth_builder_with_header_team( allowed_team_ids: set, fake_get_team_by_alias=_team_alias_lookup_404, route: str = "/chat/completions", + send_header: bool = True, ): jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = jwt_auth_config @@ -5909,7 +5916,7 @@ async def _run_auth_builder_with_header_team( user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, - request_headers={"x-litellm-team-id": header_team_id}, + request_headers={"x-litellm-team-id": header_team_id} if send_header else {}, ) @@ -7283,6 +7290,130 @@ async def test_auth_builder_header_alias_under_db_fallback_keeps_the_team_allowe assert allowed["team_id"] == "team_member" +@pytest.mark.asyncio +async def test_auth_builder_header_selects_db_membership_team_when_jwt_also_carries_a_team_claim() -> None: + """Under fallback_to_db_teams, x-litellm-team-id may name a DB-membership + team the JWT does not claim (LIT-8656): the allowed set is the JWT teams + union the user's DB memberships, not the JWT teams alone. The flag-off + path keeps rejecting the same header against the JWT's allowed teams.""" + user_object = LiteLLM_UserTable( + user_id="u_mixed", + user_role=LitellmUserRoles.INTERNAL_USER, + teams=["team_member"], + ) + config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_jwt_field="appid") + token = {"sub": "u_mixed", "scope": "", "appid": "team_claimed"} + fake_get_team = _teams_by_id(frozenset({"team_claimed", "team_member"})) + + by_membership = await _run_auth_builder_with_header_team( + config, token, "team_member", user_object, fake_get_team, {"team_claimed"} + ) + assert by_membership["team_id"] == "team_member" + assert by_membership["team_object"].team_id == "team_member" + + by_claim = await _run_auth_builder_with_header_team( + config, token, "team_claimed", user_object, fake_get_team, {"team_claimed"} + ) + assert by_claim["team_id"] == "team_claimed" + + flag_off = LiteLLM_JWTAuth(fallback_to_db_teams=False, team_id_jwt_field="appid") + with pytest.raises(HTTPException) as exc_info: + await _run_auth_builder_with_header_team( + flag_off, token, "team_member", user_object, fake_get_team, {"team_claimed"} + ) + assert exc_info.value.status_code == 403 + assert "JWT's allowed teams" in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_auth_builder_header_non_member_team_is_denied_when_jwt_also_carries_a_team_claim() -> None: + """A header naming a team the user does not belong to stays a membership + denial even when the JWT carries a team claim, and an existing but + non-member team produces the exact same 403 shape as a nonexistent one so + the response is no oracle for which team ids exist.""" + user_object = LiteLLM_UserTable( + user_id="u_mixed", + user_role=LitellmUserRoles.INTERNAL_USER, + teams=["team_member"], + ) + config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_jwt_field="appid") + token = {"sub": "u_mixed", "scope": "", "appid": "team_claimed"} + fake_get_team = _teams_by_id(frozenset({"team_claimed", "team_member", "team_other"})) + + with pytest.raises(HTTPException) as outsider_exc: + await _run_auth_builder_with_header_team( + config, token, "team_other", user_object, fake_get_team, {"team_claimed"} + ) + with pytest.raises(HTTPException) as missing_exc: + await _run_auth_builder_with_header_team( + config, token, "team_ghost", user_object, fake_get_team, {"team_claimed"} + ) + + assert outsider_exc.value.status_code == 403 + assert missing_exc.value.status_code == 403 + assert outsider_exc.value.detail == ( + "x-litellm-team-id 'team_other' does not resolve to a team id or a unique team alias among your " + "team memberships." + ) + assert missing_exc.value.detail.replace("team_ghost", "") == outsider_exc.value.detail.replace( + "team_other", "" + ) + assert "exist" not in missing_exc.value.detail + + +@pytest.mark.asyncio +async def test_auth_builder_no_header_keeps_the_jwt_team_when_fallback_to_db_teams_is_on() -> None: + """With no x-litellm-team-id header, fallback_to_db_teams must not disturb + the claim path: the JWT's own team claim still binds the request.""" + user_object = LiteLLM_UserTable( + user_id="u_mixed", + user_role=LitellmUserRoles.INTERNAL_USER, + teams=["team_member"], + ) + config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_jwt_field="appid") + token = {"sub": "u_mixed", "scope": "", "appid": "team_claimed"} + + result = await _run_auth_builder_with_header_team( + config, + token, + "team_member", + user_object, + _teams_by_id(frozenset({"team_claimed", "team_member"})), + {"team_claimed"}, + send_header=False, + ) + assert result["team_id"] == "team_claimed" + + +@pytest.mark.asyncio +async def test_auth_builder_team_id_default_does_not_widen_the_header_allowed_set() -> None: + """team_id_default fills in a team for claimless tokens but must not widen + the header's allowed set: a header naming the default team is still held + to DB membership under fallback_to_db_teams.""" + user_object = LiteLLM_UserTable( + user_id="u_default", + user_role=LitellmUserRoles.INTERNAL_USER, + teams=["team_member"], + ) + config = LiteLLM_JWTAuth(fallback_to_db_teams=True, team_id_default="team_default") + token = {"sub": "u_default", "scope": ""} + + with pytest.raises(HTTPException) as exc_info: + await _run_auth_builder_with_header_team( + config, + token, + "team_default", + user_object, + _teams_by_id(frozenset({"team_default", "team_member"})), + set(), + ) + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == ( + "x-litellm-team-id 'team_default' does not resolve to a team id or a unique team alias among your " + "team memberships." + ) + + @pytest.mark.asyncio async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_flag(): """Reading the singular team claim during sync is scoped to fallback_to_db_teams. diff --git a/tests/test_litellm/proxy/client/test_chat.py b/tests/test_litellm/proxy/client/test_chat.py index 67b6ee833f2..8fe1bfcbb2f 100644 --- a/tests/test_litellm/proxy/client/test_chat.py +++ b/tests/test_litellm/proxy/client/test_chat.py @@ -13,7 +13,7 @@ from litellm.proxy.client.exceptions import UnauthorizedError def _load_http_mocking_responses(): """Load the third-party `responses` package even if test collection creates - a top-level `responses` namespace package from `tests/test_litellm/responses`. + a top-level `responses` namespace package from `tests/unit/responses`. """ module = importlib.import_module("responses") if hasattr(module, "activate"): diff --git a/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py b/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py index 69b92f5e4d7..228fd5bcae6 100644 --- a/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py +++ b/tests/test_litellm/proxy/common_utils/test_sse_keepalive.py @@ -6,10 +6,13 @@ import pytest from fastapi.responses import StreamingResponse from litellm.proxy.common_request_processing import create_response +from litellm.types.utils import ModelResponse from litellm.proxy.common_utils.sse_keepalive import ( ANTHROPIC_PING_SSE_CHUNK, SSE_COMMENT_PING_BYTES, + advance_sse_tail, resolve_ttft_keepalive_interval, + seal_open_sse_frame, split_complete_sse_frames, wrap_passthrough_sse_bytes_with_keepalive_pings, wrap_sse_stream_with_keepalive_pings, @@ -32,6 +35,12 @@ def test_split_complete_sse_frames_holds_bytes_with_no_complete_frame(): assert split_complete_sse_frames(b"data: unterminated") == (b"", b"data: unterminated") +@pytest.mark.parametrize("chunk", [{"content": "hi"}, ModelResponse()]) +def test_advance_sse_tail_ignores_a_chunk_that_is_not_sse_text(chunk: object): + assert advance_sse_tail(b"\n\n", chunk) == b"\n\n" + assert seal_open_sse_frame(advance_sse_tail(b"data: {", chunk)) == "\n" + ANTHROPIC_PING_SSE_CHUNK + + @pytest.mark.asyncio async def test_pings_fill_mid_stream_silence_and_preserve_chunk_order(): async def gappy_stream() -> AsyncGenerator[str, None]: diff --git a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py index e6795bb22f3..42c1f489bdd 100644 --- a/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py +++ b/tests/test_litellm/proxy/hooks/test_tpm_concurrent.py @@ -3674,7 +3674,7 @@ async def test_post_call_success_hook_contains_header_merge_failures( @pytest.mark.asyncio async def test_the_project_itpm_reservation_counts_the_request_off_the_event_loop(rate_limiter): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index a40741c8fdb..3469df082e0 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -23,6 +23,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import ProxyException, UserAPIKeyAuth @@ -1165,6 +1166,29 @@ def test_resolve_llm_passthrough_timeout_precedence(): assert resolve_llm_passthrough_timeout() == 6.0 +def test_resolve_llm_passthrough_timeout_honors_explicit_global_request_timeout(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr("litellm.request_timeout", 44.0, raising=False) + monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False) + + with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}): + assert resolve_llm_passthrough_timeout() == 44.0 + assert resolve_llm_passthrough_timeout(kwargs={"stream": True}) == 44.0 + assert resolve_llm_passthrough_timeout(router_timeout=120) == 120.0 + assert resolve_llm_passthrough_timeout(kwargs={"stream": True}, router_stream_timeout=900) == 900.0 + assert resolve_llm_passthrough_timeout(litellm_params={"timeout": 90}) == 90.0 + assert resolve_llm_passthrough_timeout(kwargs={"timeout": 45}) == 45.0 + + +def test_resolve_llm_passthrough_timeout_skips_unset_global_request_timeout(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr("litellm.request_timeout", float(DEFAULT_REQUEST_TIMEOUT_SECONDS), raising=False) + monkeypatch.setattr("litellm.request_timeout_explicitly_set", False, raising=False) + + with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}): + assert resolve_llm_passthrough_timeout() == 6.0 + with patch("litellm.proxy.proxy_server.general_settings", {}): + assert resolve_llm_passthrough_timeout() == DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS + + def test_resolve_llm_passthrough_timeout_stream_timeout_precedence(): assert ( resolve_llm_passthrough_timeout( diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py index e91b7ef970c..9d4532df49a 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py @@ -162,7 +162,7 @@ async def test_interrupted_anthropic_stream_recovers_output_tokens_off_the_event from unittest.mock import AsyncMock from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, @@ -201,7 +201,7 @@ async def test_failed_anthropic_stream_records_partial_usage_off_the_event_loop( from unittest.mock import AsyncMock from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py index 089c2d57594..16cb1146ff5 100644 --- a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py +++ b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py @@ -265,7 +265,7 @@ def test_close_dangling_otel_server_span_logger_raises_state_cleared_error(monke @pytest.mark.asyncio async def test_otel_request_validation_exception_handler_returns_422_detail(): - errors = [{"loc": ["body", "model"], "msg": "field required", "type": "missing"}] + errors = [{"loc": ["body", "model"], "msg": "field required", "type": "missing", "input": {"messages": []}}] exc = RequestValidationError(errors) request = _make_request() @@ -273,7 +273,69 @@ async def test_otel_request_validation_exception_handler_returns_422_detail(): body = json.loads(response.body) assert response.status_code == 422 - assert normalize(body) == {"detail": exc.errors()} + assert body == {"detail": [{"type": "missing", "loc": ["body", "model"], "msg": "field required"}]} + + +_SUBMITTED_PASSWORD: Final = "hunter2-Sup3rSecret!" +_PASSWORD_LEAKING_ERRORS: Final = ( + { + "type": "missing", + "loc": ["body", "new_password"], + "msg": "Field required", + "input": {"current_password": _SUBMITTED_PASSWORD}, + }, + { + "type": "value_error", + "loc": ["body", "password"], + "msg": "Value error, password cannot be set via /user/new", + "input": _SUBMITTED_PASSWORD, + "ctx": {"error": ValueError(_SUBMITTED_PASSWORD)}, + }, +) +_PUBLIC_ERRORS: Final = ( + {"type": "missing", "loc": ["body", "new_password"], "msg": "Field required"}, + {"type": "value_error", "loc": ["body", "password"], "msg": "Value error, password cannot be set via /user/new"}, +) + + +@pytest.mark.asyncio +async def test_otel_request_validation_exception_handler_never_echoes_the_submitted_body(): + """A pydantic error carries the offending value as ``input`` (the whole body for a + ``missing`` error) and input-derived values in ``ctx``; a caller who mistyped a + request holding a password must not get that password back.""" + exc = RequestValidationError(list(_PASSWORD_LEAKING_ERRORS)) + + response = await otel_request_validation_exception_handler(request=_make_request(), exc=exc) + + assert response.status_code == 422 + assert json.loads(response.body) == {"detail": list(_PUBLIC_ERRORS)} + assert _SUBMITTED_PASSWORD.encode() not in response.body + + +@pytest.mark.asyncio +async def test_otel_request_validation_exception_handler_hands_the_span_only_the_public_errors(monkeypatch): + """The OTEL SERVER span's error message is ``str(exc)``, which FastAPI builds from + every error dict ``input`` included, so the span gets the same public-only errors + the caller does, and keeps the traceback the original carried.""" + import litellm.proxy.proxy_server as ps + + fake_logger = MagicMock() + monkeypatch.setattr(ps, "open_telemetry_logger", fake_logger, raising=False) + exc = RequestValidationError(list(_PASSWORD_LEAKING_ERRORS)) + try: + raise exc + except RequestValidationError as raised: + original_traceback = raised.__traceback__ + request = _make_request(parent_otel_span=MagicMock()) + + await otel_request_validation_exception_handler(request=request, exc=exc) + + (_span, span_exc, status_code) = fake_logger.record_error_attributes_on_span.call_args.args + assert status_code == 422 + assert isinstance(span_exc, RequestValidationError) + assert list(span_exc.errors()) == list(_PUBLIC_ERRORS) + assert _SUBMITTED_PASSWORD not in str(span_exc) + assert span_exc.__traceback__ is original_traceback @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py b/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py index 778acc1baab..6c1d869d113 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_onboarding.py @@ -317,6 +317,25 @@ def test_claim_onboarding_link_missing_field_422(client, monkeypatch, mock_prism assert any("password" in str(item) for item in body["detail"]) +def test_claim_onboarding_link_422_never_echoes_the_submitted_password(client): + """A body that fails validation is answered with the field path and message only; + pydantic's ``input`` (the whole submitted body for a missing field, password + included) must never come back to the caller or land in whatever logs the response.""" + password = "hunter2-Sup3rSecret!" + + response = client.post( + "/onboarding/claim_token", + json={"invitation_link": "abc", "password": password}, + ) + + assert response.status_code == 422 + assert password.encode() not in response.content + detail = response.json()["detail"] + assert detail[0]["loc"] == ["body", "user_id"] + assert detail[0]["msg"] + assert set(detail[0]) == {"type", "loc", "msg"} + + def test_claim_onboarding_link_bad_onboarding_jwt_401( client, monkeypatch, mock_prisma ): diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index bdf003085ef..c17f41a8b8f 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -7,6 +7,7 @@ from typing import AsyncGenerator, Callable, Final, Iterator, Literal, Optional, from urllib.parse import unquote_plus from unittest.mock import AsyncMock, MagicMock, patch +import anthropic import httpx import pytest from fastapi import HTTPException, Request, Response, status @@ -14,6 +15,7 @@ from fastapi.responses import JSONResponse, StreamingResponse import litellm from litellm._uuid import uuid +from litellm.anthropic_interface.exceptions import AnthropicErrorSseFrame, anthropic_error_sse_frame from litellm.litellm_core_utils.bug_report import ( DISABLE_ENV_VAR, ISSUE_URL_BASE, @@ -55,6 +57,7 @@ from litellm.proxy.common_request_processing import ( sse_error_payload, ) from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header +from litellm.proxy.common_utils.sse_keepalive import ANTHROPIC_PING_SSE_CHUNK from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import ProxyErrorTypes, ProxyException @@ -2543,6 +2546,63 @@ class TestCommonRequestProcessingHelpers: assert response.headers["x-litellm-call-id"] == "call-8302" assert json.loads(response.body) == {"error": {"code": 403, "message": "forbidden"}} + async def test_a_stream_that_fails_before_its_first_byte_answers_as_an_anthropic_json_error(self): + """A /v1/messages stream whose first chunk is already the error frame has nothing + streamed yet, so the failure answers as JSON with the status the upstream gave, + the shape Anthropic clients raise their status-specific errors on""" + + async def stream(): + yield anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable") + yield ANTHROPIC_PING_SSE_CHUNK + + generator: Final = stream() + response = await create_response(generator, "text/event-stream", {"x-litellm-call-id": "call-8609"}) + + assert isinstance(response, JSONResponse) + assert response.status_code == 503 + assert response.headers["content-type"] == "application/json" + assert response.headers["x-litellm-call-id"] == "call-8609" + assert json.loads(response.body) == { + "type": "error", + "error": {"type": "api_error", "message": "upstream unavailable"}, + } + assert generator.ag_frame is None + + async def test_a_stream_that_fails_before_its_first_byte_names_the_call_when_opted_in(self): + async def stream(): + yield anthropic_error_sse_frame(status_code=429, raw_message="slow down") + + response = await create_response( + stream(), + "text/event-stream", + {"x-litellm-call-id": "call-8609"}, + general_settings={"include_call_id_in_error_body": True}, + ) + + assert isinstance(response, JSONResponse) + assert response.status_code == 429 + assert json.loads(response.body) == { + "type": "error", + "error": {"type": "rate_limit_error", "message": "slow down", "litellm_call_id": "call-8609"}, + } + + async def test_an_error_event_after_a_keepalive_ping_still_streams(self): + """Once a keepalive ping went out the headers are committed, so the error frame + streams as an event instead of turning into a JSON answer""" + + async def stream(): + yield ANTHROPIC_PING_SSE_CHUNK + yield anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable") + + response = await create_response(stream(), "text/event-stream", {}) + + assert isinstance(response, StreamingResponse) + assert response.status_code == 200 + assert "".join(await self.consume_stream(response)) == ( + ANTHROPIC_PING_SSE_CHUNK + + 'event: error\ndata: {"type": "error", "error": {"type": "api_error", "message": "upstream unavailable"}}\n\n' + ) + async def test_create_streaming_response_disables_proxy_buffering(self): """Regression for #28384: every StreamingResponse create_response returns must carry the headers that stop nginx/ingress/Envoy from buffering the @@ -9901,6 +9961,209 @@ class TestErrorLogCarriesCallId: assert call_id in record.getMessage() +class TestAnthropicMessagesStreamErrorFrame: + """A ``/v1/messages`` stream that fails after the headers are out has to say so with an + ``event: error`` frame. Anthropic clients pick events by name, so a bare ``data:`` line is + skipped and the request looks like it ended with nothing in it""" + + @staticmethod + def _sse_generator_failing_with(failure: Exception) -> AsyncGenerator[str, None]: + class FailingUpstream: + def __aiter__(self) -> "FailingUpstream": + return self + + async def __anext__(self) -> object: + raise failure + + ProxyLogging._callback_capabilities_cache.clear() + return ProxyBaseLLMRequestProcessing.async_sse_data_generator( + response=FailingUpstream(), + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"), + request_data={"model": "claude-sonnet-4-5"}, + proxy_logging_obj=ProxyLogging(user_api_key_cache=MagicMock()), + ) + + @pytest.mark.parametrize( + "status_code, expected_error_type", + [ + (429, "rate_limit_error"), + (529, "overloaded_error"), + (413, "request_too_large"), + (500, "api_error"), + (502, "api_error"), + (400, "invalid_request_error"), + ], + ) + async def test_mid_stream_failure_arrives_as_an_anthropic_error_event( + self, status_code: int, expected_error_type: str + ) -> None: + class UpstreamFailure(Exception): + def __init__(self) -> None: + super().__init__("upstream stopped sending") + self.status_code: Final = status_code + + frames: Final = [frame async for frame in self._sse_generator_failing_with(UpstreamFailure())] + + assert len(frames) == 1 + event_line, data_line, first_blank, second_blank = frames[0].split("\n") + assert isinstance(frames[0], AnthropicErrorSseFrame) + assert frames[0].status_code == status_code + assert event_line == "event: error" + assert (first_blank, second_blank) == ("", "") + payload: Final = json.loads(data_line.removeprefix("data: ")) + assert payload["type"] == "error" + assert payload["error"]["type"] == expected_error_type + assert "upstream stopped sending" in payload["error"]["message"] + + _CONTENT_DELTA_FRAME: Final = ( + b"event: content_block_delta\n" + b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"1\\n2\\n3"}}\n\n' + ) + _TORN_DATA_LINE: Final = ( + b"event: content_block_delta\n" + b'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"4' + ) + _PING: Final = ANTHROPIC_PING_SSE_CHUNK.encode() + + @staticmethod + def _upstream_failure(status_code: int) -> Exception: + class UpstreamFailure(Exception): + def __init__(self) -> None: + super().__init__("upstream stopped sending") + self.status_code: Final = status_code + + return UpstreamFailure() + + @staticmethod + def _sse_generator_cut_after(relayed: Sequence[bytes], failure: Exception) -> AsyncGenerator[str, None]: + class CutUpstream: + def __init__(self) -> None: + self._remaining: Final = iter(relayed) + + def __aiter__(self) -> "CutUpstream": + return self + + async def __anext__(self) -> object: + chunk: Final = next(self._remaining, None) + if chunk is None: + raise failure + return chunk + + ProxyLogging._callback_capabilities_cache.clear() + return ProxyBaseLLMRequestProcessing.async_sse_data_generator( + response=CutUpstream(), + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"), + request_data={"model": "claude-sonnet-4-5"}, + proxy_logging_obj=ProxyLogging(user_api_key_cache=MagicMock()), + ) + + @staticmethod + def _as_bytes(chunk: object) -> bytes: + if isinstance(chunk, bytes): + return chunk + assert isinstance(chunk, str) + return chunk.encode() + + async def _wire_bytes(self, relayed: Sequence[bytes]) -> bytes: + stream: Final = self._sse_generator_cut_after(relayed, self._upstream_failure(500)) + return b"".join([self._as_bytes(chunk) async for chunk in stream]) + + @staticmethod + def _error_frame_after(wire: bytes, relayed: bytes) -> bytes: + assert wire.startswith(relayed), f"the wire did not open with {relayed!r}: {wire!r}" + return wire.removeprefix(relayed) + + @staticmethod + def _assert_error_frame(frame: bytes) -> None: + event_line, data_line, first_blank, second_blank = frame.split(b"\n") + assert event_line == b"event: error" + assert (first_blank, second_blank) == (b"", b"") + payload: Final = json.loads(data_line.removeprefix(b"data: ")) + assert payload["type"] == "error" + assert "upstream stopped sending" in payload["error"]["message"] + + @pytest.mark.parametrize( + "torn, seal", + [ + (_TORN_DATA_LINE, b"\n" + _PING), + (b"event: content_bl", b"\n" + _PING), + (b"event: content_block_delta\n", _PING), + (b'event: content_block_delta\r\ndata: {"type":"content_block_delta"}\r\n', _PING), + ], + ids=["mid_data_line", "mid_event_line", "after_a_complete_line", "after_a_crlf_line"], + ) + async def test_a_frame_the_upstream_tore_is_closed_as_a_ping_before_the_error_event( + self, torn: bytes, seal: bytes + ) -> None: + wire: Final = await self._wire_bytes((self._CONTENT_DELTA_FRAME, torn)) + + self._assert_error_frame(self._error_frame_after(wire, self._CONTENT_DELTA_FRAME + torn + seal)) + + async def test_a_cut_at_a_frame_boundary_gets_the_error_event_alone(self) -> None: + wire: Final = await self._wire_bytes((self._CONTENT_DELTA_FRAME,)) + + self._assert_error_frame(self._error_frame_after(wire, self._CONTENT_DELTA_FRAME)) + + async def test_a_torn_frame_still_raises_the_error_in_the_anthropic_sdk(self) -> None: + wire: Final = await self._wire_bytes((self._CONTENT_DELTA_FRAME, self._TORN_DATA_LINE)) + + def serve(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=wire) + + client: Final = anthropic.Anthropic( + api_key="sk-test", + base_url="http://proxy.test", + http_client=httpx.Client(transport=httpx.MockTransport(serve)), + max_retries=0, + ) + with pytest.raises(anthropic.APIStatusError) as raised: + for _ in client.messages.create( + model="claude-sonnet-4-5", max_tokens=16, messages=[{"role": "user", "content": "count"}], stream=True + ): + pass + body: Final = raised.value.body + assert isinstance(body, dict) + assert body["type"] == "error" + assert "upstream stopped sending" in body["error"]["message"] + + async def test_a_failure_before_the_first_byte_answers_with_its_status_as_json(self) -> None: + response: Final = await create_response( + self._sse_generator_failing_with(self._upstream_failure(502)), "text/event-stream", {} + ) + + assert isinstance(response, JSONResponse) + assert response.status_code == 502 + body: Final = json.loads(response.body) + assert body["type"] == "error" + assert body["error"]["type"] == "api_error" + assert "upstream stopped sending" in body["error"]["message"] + + async def test_a_failure_before_the_first_byte_raises_with_its_status_in_the_anthropic_sdk(self) -> None: + response: Final = await create_response( + self._sse_generator_failing_with(self._upstream_failure(502)), "text/event-stream", {} + ) + assert isinstance(response, JSONResponse) + + def serve(request: httpx.Request) -> httpx.Response: + return httpx.Response(response.status_code, headers=dict(response.headers), content=response.body) + + client: Final = anthropic.Anthropic( + api_key="sk-test", + base_url="http://proxy.test", + http_client=httpx.Client(transport=httpx.MockTransport(serve)), + max_retries=0, + ) + with pytest.raises(anthropic.APIStatusError) as raised: + client.messages.create( + model="claude-sonnet-4-5", max_tokens=16, messages=[{"role": "user", "content": "count"}], stream=True + ) + assert raised.value.status_code == 502 + body: Final = raised.value.body + assert isinstance(body, dict) + assert body["type"] == "error" + assert "upstream stopped sending" in body["error"]["message"] + + class TestStreamingContainerOwnershipRecordedBeforeDone: """Regression for LIT-8612: the OpenAI SDK closes the connection at ``data: [DONE]`` and starlette cancels the body task, so an ownership row diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 884a9c81500..89156cd19a0 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -14974,7 +14974,7 @@ def test_settings_store_exposes_dashboard_saved_mcp_client_allowlist_to_the_mcp_ async def test_token_counter_keeps_the_event_loop_free_during_a_huggingface_count(monkeypatch): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, @@ -14995,7 +14995,7 @@ async def test_token_counter_loads_a_custom_tokenizer_off_the_event_loop(monkeyp from litellm.rust_bridge._native import Tokenizer from litellm import Router - from tests.test_litellm.litellm_core_utils.event_loop_lag import assert_loop_stayed_free, timed_with_loop_lags + from tests.unit.litellm_core_utils.event_loop_lag import assert_loop_stayed_free, timed_with_loop_lags claude_tokenizer: Final = litellm.utils._select_tokenizer("claude-fable-5")["tokenizer"] diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 0fc7295a717..ea1870d3b73 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -2021,7 +2021,7 @@ async def test_a_dispatched_failure_is_counted_off_the_event_loop(): from unittest.mock import AsyncMock, patch from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/proxy/test_zerobus_dashboard_config.py b/tests/test_litellm/proxy/test_zerobus_dashboard_config.py new file mode 100644 index 00000000000..d2143767480 --- /dev/null +++ b/tests/test_litellm/proxy/test_zerobus_dashboard_config.py @@ -0,0 +1,52 @@ +from pathlib import Path +from typing import Final + +from pydantic import BaseModel, TypeAdapter + +import litellm +from litellm.integrations.custom_logger import CustomLogger + + +class DashboardField(BaseModel): + type: str + required: bool + + +class DashboardCallbackConfig(BaseModel): + id: str + displayName: str + logo: str + supports_key_team_logging: bool + dynamic_params: dict[str, DashboardField] + + +def _zerobus_config() -> DashboardCallbackConfig: + path: Final = Path(litellm.__file__).parent / "integrations" / "callback_configs.json" + configs: Final = TypeAdapter(tuple[DashboardCallbackConfig, ...]).validate_json(path.read_text()) + return next(config for config in configs if config.id == "zerobus") + + +def test_zerobus_appears_in_the_dashboard_callback_dropdown(): + """The dropdown is served from callback_configs.json, so an entry only in the dashboard source is invisible.""" + entry = _zerobus_config() + + assert entry.displayName == "Databricks Zerobus" + assert entry.supports_key_team_logging is False + assert entry.dynamic_params["ZEROBUS_CLIENT_SECRET"].type == "password" + assert all(field.required is True for field in entry.dynamic_params.values()) + + +def test_the_dropdown_logo_asset_exists(): + """A logo the dashboard cannot resolve degrades silently to a letter tile.""" + logo = _zerobus_config().logo + repo_root = Path(litellm.__file__).parent.parent + asset = repo_root / "ui" / "litellm-dashboard" / "public" / "assets" / "logos" / logo + + assert asset.is_file() + + +def test_the_dropdown_fields_are_the_env_vars_the_logger_reads(): + """Naming the fields as stored means the edit form prefills saved values instead of showing blanks.""" + fields = tuple(_zerobus_config().dynamic_params) + + assert fields == tuple(CustomLogger.get_callback_env_vars("zerobus")) diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 08d542df16c..0d7a713a380 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -3817,6 +3817,22 @@ class TestTeamAdminEditableTeamFieldsSetting: assert response.status_code == 422 + def test_patch_422_never_echoes_the_submitted_value(self, monkeypatch): + self._as_proxy_admin(monkeypatch) + submitted = "hunter2-Sup3rSecret!" + + try: + response = client.patch("/update/ui_settings", json={"team_admin_editable_team_fields": submitted}) + finally: + app.dependency_overrides.clear() + + assert response.status_code == 422 + assert submitted.encode() not in response.content + detail = response.json()["detail"] + assert detail[0]["loc"] == ["team_admin_editable_team_fields"] + assert detail[0]["msg"] + assert set(detail[0]) == {"type", "loc", "msg"} + def test_patch_persists_and_syncs_the_list_to_general_settings(self, monkeypatch): mock_prisma = self._as_proxy_admin(monkeypatch) general_settings: dict = {"team_admin_editable_team_fields": []} diff --git a/tests/test_litellm/rust_bridge/messages/test_route_host.py b/tests/test_litellm/rust_bridge/messages/test_route_host.py deleted file mode 100644 index c5a442e0709..00000000000 --- a/tests/test_litellm/rust_bridge/messages/test_route_host.py +++ /dev/null @@ -1,124 +0,0 @@ -from dataclasses import astuple -from typing import Final - -import pytest - -import litellm -from litellm.rust_bridge.messages import route_host - -pytestmark = pytest.mark.usefixtures("local_model_cost_map") - - -def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None: - monkeypatch.setitem( - litellm.model_cost, - name, - { - "litellm_provider": "anthropic", - "mode": "chat", - "input_cost_per_token": 0, - "output_cost_per_token": 0, - **flags, - }, - ) - - -def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None: - _flag_model( - monkeypatch, - "claude-test-adaptive", - supports_reasoning=True, - supports_adaptive_thinking=True, - supports_output_config=True, - supports_xhigh_reasoning_effort=True, - supports_sampling_params=False, - ) - - capabilities: Final = route_host.model_capabilities("anthropic/claude-test-adaptive", None) - - assert capabilities.supports_adaptive_thinking - assert capabilities.supports_output_config - assert not capabilities.supports_legacy_thinking - assert not capabilities.supports_sampling_params - assert capabilities.effort_tiers.xhigh - assert not capabilities.effort_tiers.max - - -def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None: - capabilities: Final = route_host.model_capabilities("anthropic/not-a-real-model", None) - - assert capabilities.supports_sampling_params - assert not capabilities.supports_reasoning - assert not capabilities.supports_adaptive_thinking - assert not any(astuple(capabilities.effort_tiers)) - - -@pytest.mark.parametrize( - ("global_flag", "kwargs", "expected"), - [ - (False, {}, False), - (True, {}, True), - (False, {"drop_params": "true"}, True), - (False, {"drop_params": "nonsense"}, False), - (False, {"drop_params": False}, False), - ], -) -def test_drop_params_merges_the_global_flag_with_the_request( - monkeypatch: pytest.MonkeyPatch, global_flag: bool, kwargs: dict[str, object], expected: bool -) -> None: - monkeypatch.setattr(litellm, "drop_params", global_flag) - - assert route_host.shaping("anthropic/not-a-real-model", None, kwargs)["drop_params"] is expected - - -@pytest.mark.parametrize( - ("configured", "expected"), - [ - (["tools[*].input_examples", 3, "metadata.user_id"], ("tools[*].input_examples", "metadata.user_id")), - ("tools", ()), - (None, ()), - ], -) -def test_additional_drop_params_keep_only_string_paths(configured: object, expected: tuple[str, ...]) -> None: - shaping: Final = route_host.shaping("anthropic/not-a-real-model", None, {"additional_drop_params": configured}) - - assert shaping["additional_drop_params"] == expected - - -def test_native_request_rejections_map_to_the_public_400() -> None: - from types import MappingProxyType - - from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest - - request: Final = LiteLLMMessagesRequest( - model="anthropic/claude-sonnet-5", - messages=(), - max_tokens=8, - stream=None, - api_key=None, - api_base=None, - custom_llm_provider=None, - kwargs=MappingProxyType({}), - ) - rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5") - rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets - - mapped: Final = route_host.map_failure(rejected, request, "anthropic") - - assert isinstance(mapped, litellm.BadRequestError) - assert mapped.status_code == 400 - assert "does not support top_k=5" in mapped.message - assert mapped.model == "claude-sonnet-5" - assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError) - - -def test_stream_hidden_params_projects_upstream_headers_the_way_the_python_handler_does() -> None: - hidden: Final = route_host.stream_hidden_params( - (("request-id", "req_upstream_123"), ("x-ratelimit-remaining-requests", "41")) - ) - - additional: Final = hidden["additional_headers"] - assert isinstance(additional, dict) - assert additional["llm_provider-request-id"] == "req_upstream_123" - assert additional["x-ratelimit-remaining-requests"] == "41" - assert "request-id" not in additional diff --git a/tests/test_litellm_rust/tokenizer/test_fast_count.py b/tests/test_litellm_rust/tokenizer/test_fast_count.py index 2902b79dca8..f91f47e4b86 100644 --- a/tests/test_litellm_rust/tokenizer/test_fast_count.py +++ b/tests/test_litellm_rust/tokenizer/test_fast_count.py @@ -7,7 +7,7 @@ from tokenizers import Tokenizer as ReferenceTokenizer from litellm.rust_bridge import _native from litellm.utils import claude_json_str -from tests.test_litellm.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON +from tests.unit.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON pytestmark = pytest.mark.requires_rust_extension diff --git a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py index abf6a6dda31..2e883e91fda 100644 --- a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py +++ b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py @@ -94,7 +94,7 @@ class _AgentChunk: @pytest.mark.asyncio async def test_stream_completion_counts_tokens_off_the_event_loop(monkeypatch): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/unit/a2a_protocol/test_main.py b/tests/unit/a2a_protocol/test_main.py index c65d171246d..4ba0ef8fa04 100644 --- a/tests/unit/a2a_protocol/test_main.py +++ b/tests/unit/a2a_protocol/test_main.py @@ -469,7 +469,7 @@ class _UsageRecorder(CustomLogger): @pytest.mark.asyncio async def test_asend_message_counts_usage_off_the_event_loop(monkeypatch): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/unit/anthropic_interface/exceptions/test_exception_mapping_utils.py b/tests/unit/anthropic_interface/exceptions/test_exception_mapping_utils.py index ef092b65f28..0d8e7674e7d 100644 --- a/tests/unit/anthropic_interface/exceptions/test_exception_mapping_utils.py +++ b/tests/unit/anthropic_interface/exceptions/test_exception_mapping_utils.py @@ -3,8 +3,15 @@ Tests for AnthropicExceptionMapping class in litellm/anthropic_interface/excepti """ import json +from typing import Final -from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping +import pytest + +from litellm.anthropic_interface.exceptions import ( + AnthropicErrorSseFrame, + AnthropicExceptionMapping, + anthropic_error_sse_frame, +) class TestCreateErrorResponse: @@ -206,3 +213,42 @@ class TestTransformToAnthropicError: ) assert result["type"] == "error" assert result["error"]["message"] == '["error1", "error2"]' + + +class TestAnthropicErrorSseFrame: + @pytest.mark.parametrize( + ("status_code", "expected_error_type"), + [(429, "rate_limit_error"), (503, "api_error"), (400, "invalid_request_error")], + ) + def test_the_frame_is_one_error_event_carrying_the_anthropic_envelope( + self, status_code: int, expected_error_type: str + ) -> None: + frame: Final = anthropic_error_sse_frame(status_code=status_code, raw_message="upstream unavailable") + + event_line, data_line, first_blank, second_blank = frame.split("\n") + assert event_line == "event: error" + assert (first_blank, second_blank) == ("", "") + assert json.loads(data_line.removeprefix("data: ")) == { + "type": "error", + "error": {"type": expected_error_type, "message": "upstream unavailable"}, + } + + def test_the_frame_remembers_the_status_and_body_it_was_built_from(self) -> None: + frame: Final = anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable") + + assert isinstance(frame, AnthropicErrorSseFrame) + assert frame.status_code == 503 + data_line: Final = frame.split("\n")[1] + assert data_line == f"data: {json.dumps(frame.json_body(call_id=None))}" + + def test_the_json_body_names_the_call_only_when_asked(self) -> None: + frame: Final = anthropic_error_sse_frame(status_code=503, raw_message="upstream unavailable") + + assert frame.json_body(call_id="call-1") == { + "type": "error", + "error": {"type": "api_error", "message": "upstream unavailable", "litellm_call_id": "call-1"}, + } + assert frame.json_body(call_id=None) == { + "type": "error", + "error": {"type": "api_error", "message": "upstream unavailable"}, + } diff --git a/tests/unit/batches/test_batch_utils.py b/tests/unit/batches/test_batch_utils.py index dd95addac40..b8b922f72a7 100644 --- a/tests/unit/batches/test_batch_utils.py +++ b/tests/unit/batches/test_batch_utils.py @@ -464,6 +464,40 @@ def test_total_cost_applies_the_long_context_batch_tier_per_line(): assert result.cost == pytest.approx((300_000 * 2e-6) + (10 * 6e-6) + (100 * 1e-6) + (10 * 4e-6)) +def test_xai_output_lines_bill_reasoning_tokens_as_completion_tokens(): + row = _success_row( + model="grok-4.3", + usage={ + "prompt_tokens": 615, + "completion_tokens": 3, + "total_tokens": 993, + "completion_tokens_details": {"reasoning_tokens": 375}, + }, + ) + + result = bu._aggregate_batch_cost_usage_models( + entries=[row], + custom_llm_provider="xai", + model_info=ModelInfo( + key="xai/grok-4.3", + max_tokens=None, + max_input_tokens=None, + max_output_tokens=None, + input_cost_per_token=1.25e-6, + output_cost_per_token=2.5e-6, + litellm_provider="xai", + mode="chat", + supported_openai_params=None, + input_cost_per_token_batches=1e-6, + output_cost_per_token_batches=2e-6, + ), + ) + + assert result.usage.completion_tokens == 378 + assert result.usage.total_tokens == 993 + assert result.cost == pytest.approx((615 * 1e-6) + (378 * 2e-6)) + + def test_total_usage_empty_is_zero(): result = bu._aggregate_batch_cost_usage_models(entries=[], custom_llm_provider="openai") assert result.cost == 0.0 diff --git a/tests/test_litellm/caching/test_azure_blob_cache.py b/tests/unit/caching/test_azure_blob_cache.py similarity index 100% rename from tests/test_litellm/caching/test_azure_blob_cache.py rename to tests/unit/caching/test_azure_blob_cache.py diff --git a/tests/test_litellm/caching/test_caching.py b/tests/unit/caching/test_caching.py similarity index 100% rename from tests/test_litellm/caching/test_caching.py rename to tests/unit/caching/test_caching.py diff --git a/tests/unit/caching/test_caching_handler.py b/tests/unit/caching/test_caching_handler.py index a181ef89fe0..425d657312a 100644 --- a/tests/unit/caching/test_caching_handler.py +++ b/tests/unit/caching/test_caching_handler.py @@ -39,6 +39,11 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm._logging import verbose_logger import logging +import json +import httpx +import respx +from fastapi.testclient import TestClient +from litellm.caching.caching_handler import _PENDING_CACHE_WRITES def setup_cache(): @@ -1062,6 +1067,9 @@ def test_is_chat_completion_cached_dict(): assert _is_chat_completion_cached_dict( {"id": "other", "object": "chat.completion.chunk", "choices": []} ) + assert _is_chat_completion_cached_dict( + {"id": "no-object", "choices": [{"index": 0}]} + ) assert not _is_chat_completion_cached_dict( {"id": "resp_abc", "object": "response", "output": []} ) @@ -1432,3 +1440,799 @@ def test_convert_cached_responses_result_parameterized( assert result is not None assert result.id == cached_result["id"] assert result.status == cached_result["status"] + + +@pytest.mark.asyncio +async def test_process_async_embedding_cached_response(): + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + args = { + "cached_result": [ + { + "embedding": [-0.025122925639152527, -0.019487135112285614], + "index": 0, + "object": "embedding", + } + ] + } + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=args["cached_result"], + kwargs={"model": "text-embedding-ada-002", "input": "test"}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="text-embedding-ada-002", + ) + + assert cache_hit + + print(f"response: {response}") + assert len(response.data) == 1 + + +@pytest.mark.asyncio +async def test_embedding_cache_preserves_prompt_tokens_details(): + """Test that prompt_tokens_details (including image_count) survives a full cache hit.""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "amazon.titan-embed-image-v1", + "prompt_tokens_details": {"image_count": 1}, + } + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "amazon.titan-embed-image-v1", "input": "base64imagedata"}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="amazon.titan-embed-image-v1", + ) + + assert cache_hit + assert response.usage is not None + assert response.usage.prompt_tokens_details is not None + assert response.usage.prompt_tokens_details.image_count == 1 + + +@pytest.mark.asyncio +async def test_embedding_cache_backward_compat_no_prompt_tokens_details(): + """Test that old cached items without prompt_tokens_details still work.""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + # Old-format cached item — no prompt_tokens_details field + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "text-embedding-ada-002", + } + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "text-embedding-ada-002", "input": "test"}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="text-embedding-ada-002", + ) + + assert cache_hit + assert response.usage is not None + assert response.usage.prompt_tokens_details is None + + +@pytest.mark.asyncio +async def test_embedding_cache_aggregates_multiple_image_counts(): + """Test that image_count is summed correctly across multiple cached items.""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "amazon.titan-embed-image-v1", + "prompt_tokens_details": {"image_count": 1}, + }, + { + "embedding": [0.031, 0.042], + "index": 1, + "object": "embedding", + "model": "amazon.titan-embed-image-v1", + "prompt_tokens_details": {"image_count": 1}, + }, + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={ + "model": "amazon.titan-embed-image-v1", + "input": ["img1", "img2"], + }, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="amazon.titan-embed-image-v1", + ) + + assert cache_hit + assert response.usage.prompt_tokens_details is not None + assert response.usage.prompt_tokens_details.image_count == 2 + + +def test_combine_usage_merges_prompt_tokens_details(): + """Test that combine_usage merges prompt_tokens_details from both Usage objects.""" + from litellm.types.utils import PromptTokensDetailsWrapper, Usage + + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + usage1 = Usage( + prompt_tokens=10, + completion_tokens=0, + total_tokens=10, + prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1), + ) + usage2 = Usage( + prompt_tokens=20, + completion_tokens=0, + total_tokens=20, + prompt_tokens_details=PromptTokensDetailsWrapper(image_count=2), + ) + + combined = llm_caching_handler.combine_usage(usage1, usage2) + + assert combined.prompt_tokens == 30 + assert combined.total_tokens == 30 + assert combined.prompt_tokens_details is not None + assert combined.prompt_tokens_details.image_count == 3 + + +def test_combine_usage_handles_none_details(): + """Test that combine_usage works when one or both sides have null prompt_tokens_details.""" + from litellm.types.utils import PromptTokensDetailsWrapper, Usage + + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + # Both null + usage_a = Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) + usage_b = Usage(prompt_tokens=20, completion_tokens=0, total_tokens=20) + combined = llm_caching_handler.combine_usage(usage_a, usage_b) + assert combined.prompt_tokens_details is None + + # Only first has details + usage_c = Usage( + prompt_tokens=10, + completion_tokens=0, + total_tokens=10, + prompt_tokens_details=PromptTokensDetailsWrapper(image_count=1), + ) + combined = llm_caching_handler.combine_usage(usage_c, usage_b) + assert combined.prompt_tokens_details is not None + assert combined.prompt_tokens_details.image_count == 1 + + # Only second has details + combined = llm_caching_handler.combine_usage(usage_a, usage_c) + assert combined.prompt_tokens_details is not None + assert combined.prompt_tokens_details.image_count == 1 + + +def _build_logging_obj(call_type: str, stream: bool): + import uuid as _uuid + + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + return LiteLLMLogging( + litellm_call_id=str(datetime.now()), + call_type=call_type, + model="gpt-5.4", + messages=[], + function_id=str(_uuid.uuid4()), + stream=stream, + start_time=datetime.now(), + ) + + +def test_convert_cached_responses_bridge_chat_completion_nonstream(): + """openai/responses chat-completions bridge: non-streaming cache hit replays as ModelResponse.""" + from litellm import responses + from litellm.types.utils import CallTypes, ModelResponse + + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + cached_result = { + "id": "chatcmpl-bridge-nonstream", + "object": "chat.completion", + "created": int(time.time()), + "model": "gpt-5.4", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Hi!"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}, + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={ + "model": "gpt-5.4", + "stream": False, + "messages": [{"role": "user", "content": "hi"}], + }, + logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False), + model="gpt-5.4", + args=(), + ) + + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "Hi!" + + +def test_convert_cached_responses_legacy_nonstream_path(): + """Genuine ResponsesAPIResponse dict (no chatcmpl/choices) falls through legacy path.""" + from litellm import responses + from litellm.types.llms.openai import ResponsesAPIResponse + from litellm.types.utils import CallTypes + + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + cached_result = { + "id": "resp_legacy_nonstream", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "message", + "id": "msg_legacy", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "legacy response", + "annotations": [], + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "hi", "stream": False}, + logging_obj=_build_logging_obj(CallTypes.responses.value, stream=False), + model="gpt-4o", + args=(), + ) + + assert isinstance(result, ResponsesAPIResponse) + assert result.id == "resp_legacy_nonstream" + + +def test_convert_cached_responses_legacy_stream_path(): + """Genuine ResponsesAPIResponse dict (no chatcmpl/choices) on stream falls through legacy path.""" + from litellm import responses + from litellm.responses.streaming_iterator import ( + CachedResponsesAPIStreamingIterator, + ) + from litellm.types.utils import CallTypes + + caching_handler = LLMCachingHandler( + original_function=responses, request_kwargs={}, start_time=datetime.now() + ) + cached_result = { + "id": "resp_legacy_stream", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-4o", + "object": "response", + "output": [ + { + "type": "message", + "id": "msg_legacy_stream", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "legacy stream", + "annotations": [], + } + ], + } + ], + } + + result = caching_handler._convert_cached_result_to_model_response( + cached_result=cached_result, + call_type=CallTypes.responses.value, + kwargs={"model": "gpt-4o", "input": "hi", "stream": True}, + logging_obj=_build_logging_obj(CallTypes.responses.value, stream=True), + model="gpt-4o", + args=(), + ) + + assert isinstance(result, CachedResponsesAPIStreamingIterator) + + +@pytest.mark.asyncio +async def test_embedding_cache_restores_stored_prompt_tokens_for_image_input(): + """Image-embedding cache hit restores prompt_tokens=0 from the stored value + instead of recomputing a bogus count by tokenizing the base64 input.""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + # base64-like blob — token_counter over this would return a large nonzero count + image_input = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk" * 50 + + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "amazon.titan-embed-image-v1", + "prompt_tokens": 0, + "prompt_tokens_details": {"image_count": 1}, + } + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "amazon.titan-embed-image-v1", "input": image_input}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="amazon.titan-embed-image-v1", + ) + + assert cache_hit + assert response.usage is not None + assert response.usage.prompt_tokens == 0 + assert response.usage.total_tokens == 0 + assert response.usage.prompt_tokens_details.image_count == 1 + + +@pytest.mark.asyncio +async def test_embedding_cache_sums_stored_prompt_tokens_across_items(): + """A multi-item cache hit sums the stored per-item prompt_tokens back to the total.""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + cached_result = [ + { + "embedding": [-0.01], + "index": 0, + "object": "embedding", + "model": "text-embedding-3-small", + "prompt_tokens": 5, + }, + { + "embedding": [-0.02], + "index": 1, + "object": "embedding", + "model": "text-embedding-3-small", + "prompt_tokens": 4, + }, + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "text-embedding-3-small", "input": ["hello world", "foo bar"]}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="text-embedding-3-small", + ) + + assert cache_hit + assert response.usage.prompt_tokens == 9 + assert response.usage.total_tokens == 9 + + +@pytest.mark.asyncio +async def test_embedding_cache_falls_back_to_token_counter_for_legacy_entries(): + """Legacy cache entries with no stored prompt_tokens still recompute via token_counter + for str inputs (backward compatibility).""" + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + # No prompt_tokens key — pre-fix entry + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "text-embedding-ada-002", + }, + ] + + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "text-embedding-ada-002", "input": "hello world"}, + logging_obj=mock_logging_obj, + start_time=datetime.now(), + model="text-embedding-ada-002", + ) + + assert cache_hit + # token_counter over "hello world" yields a nonzero count — fallback path still runs + assert response.usage.prompt_tokens > 0 + + +@pytest.mark.asyncio +async def test_embedding_cache_hit_sets_custom_llm_provider_on_logging_obj(): + """A full embedding cache hit must stamp the resolved provider onto the logging + obj so spend logs record the provider instead of None/unknown.""" + from litellm.types.utils import CallTypes + + llm_caching_handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs={}, + start_time=datetime.now(), + ) + + cached_result = [ + { + "embedding": [-0.025, -0.019], + "index": 0, + "object": "embedding", + "model": "text-embedding-3-small", + "prompt_tokens": 5, + } + ] + + logging_obj = _build_logging_obj(CallTypes.aembedding.value, stream=False) + logging_obj.async_success_handler = AsyncMock() + + response, cache_hit = llm_caching_handler._process_async_embedding_cached_response( + final_embedding_cached_response=None, + cached_result=cached_result, + kwargs={"model": "text-embedding-3-small", "input": "hello world"}, + logging_obj=logging_obj, + start_time=datetime.now(), + model="text-embedding-3-small", + ) + + assert cache_hit + assert logging_obj.model_call_details["custom_llm_provider"] == "openai" + + +def test_sync_stream_responses_cache_hit_sets_custom_llm_provider_on_logging_obj(monkeypatch): + import litellm + from litellm.caching.caching import Cache + from litellm.types.utils import CallTypes + + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + kwargs = {"model": "azure/gpt-5.4-mini", "input": "hello", "stream": True} + cached_response = { + "id": "resp_sync_stream", + "created_at": int(time.time()), + "status": "completed", + "model": "gpt-5.4-mini", + "object": "response", + "output": [ + { + "type": "message", + "id": "msg_sync_stream", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "hi", "annotations": []}], + } + ], + } + litellm.cache.add_cache(json.dumps(cached_response), **kwargs) + handler = LLMCachingHandler(original_function=litellm.responses, request_kwargs=kwargs, start_time=datetime.now()) + logging_obj = _build_logging_obj(CallTypes.responses.value, stream=True) + + hit = handler._sync_get_cache( + model="azure/gpt-5.4-mini", + original_function=litellm.responses, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.responses.value, + kwargs=kwargs, + args=(), + ) + + assert hit.cached_result is not None + assert logging_obj.model_call_details["custom_llm_provider"] == "azure" + assert logging_obj.model_call_details["litellm_params"]["custom_llm_provider"] == "azure" + + +def test_request_kwargs_does_not_retain_logging_obj(): + """ + The caching handler lives on logging_obj._llm_caching_handler, so keeping + litellm_logging_obj inside request_kwargs closes a reference cycle + (Logging -> LLMCachingHandler -> kwargs -> Logging). That cycle keeps the + full request payload alive until a generational GC pass instead of being + freed by refcount when the request finishes; under bursts of large-token + requests this presents as stepwise RSS growth that never returns to + baseline. Other kwargs (messages included) must be preserved. + """ + logging_obj = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hello"}], + "litellm_logging_obj": logging_obj, + } + + handler = LLMCachingHandler( + original_function=MagicMock(), + request_kwargs=kwargs, + start_time=datetime.now(), + ) + + assert "litellm_logging_obj" not in handler.request_kwargs + assert handler.request_kwargs["messages"] == kwargs["messages"] + assert handler.request_kwargs["model"] == "gpt-4o" + + +def test_async_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch): + """ + Regression test for the SDK losing async cache writes in short-lived scripts: + async_set_cache dispatched the write as a bare fire-and-forget task, so + asyncio.run cancelled it at loop close before the write landed (LIT-6184, + deterministic with hiredis installed). The write must survive loop shutdown. + """ + import litellm + + writes = [] + + class _SlowWriteCache: + supported_call_types = ["acompletion"] + cache = None + + async def async_add_cache(self, result, dynamic_cache_object=None, **kwargs): + await asyncio.sleep(0.2) + writes.append(result) + + async def acompletion(**kwargs): + return None + + handler = LLMCachingHandler( + original_function=acompletion, + request_kwargs={}, + start_time=datetime.now(), + ) + monkeypatch.setattr(litellm, "cache", _SlowWriteCache()) + + async def _short_lived_script(): + await handler.async_set_cache( + result=litellm.ModelResponse(), + original_function=acompletion, + kwargs={}, + ) + + asyncio.run(_short_lived_script()) + + assert len(writes) == 1 + + +@pytest.mark.asyncio +async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monkeypatch): + """The spend log for a cache hit must reuse the key the lookup already computed instead of hashing again.""" + import litellm + from litellm.caching.caching import Cache + from litellm.types.utils import CallTypes + + async def acompletion(**kwargs): + return None + + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + kwargs = {"model": "gpt-5.4", "messages": [{"role": "user", "content": "hello"}], "caching": True} + await litellm.cache.async_add_cache( + litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "hi"}}]), **kwargs + ) + handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now()) + logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False) + logging_obj.async_success_handler = AsyncMock() + + hit = await handler._async_get_cache( + model="gpt-5.4", + original_function=acompletion, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.acompletion.value, + kwargs=kwargs, + args=(), + ) + + assert hit is not None and hit.cached_result is not None + assert handler.preset_cache_key is not None + assert logging_obj.litellm_params["preset_cache_key"] == handler.preset_cache_key + assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key + + +@pytest.mark.asyncio +async def test_converted_stream_cache_hit_replayed_as_plain_object_logs_at_hit_time(monkeypatch): + import litellm + from litellm.caching.caching import Cache + from litellm.types.utils import CallTypes + + async def aanthropic_messages(**kwargs): + return None + + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + kwargs = { + "model": "claude-sonnet-5", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 16, + "caching": True, + "stream": False, + "_websearch_interception_converted_stream": True, + } + cached_message = { + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "hi"}], + } + await litellm.cache.async_add_cache(cached_message, **kwargs) + handler = LLMCachingHandler(original_function=aanthropic_messages, request_kwargs=kwargs, start_time=datetime.now()) + logging_obj = _build_logging_obj(CallTypes.aanthropic_messages.value, stream=False) + logging_obj.async_success_handler = AsyncMock() + logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() + + hit = await handler._async_get_cache( + model="claude-sonnet-5", + original_function=aanthropic_messages, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.aanthropic_messages.value, + kwargs=kwargs, + args=(), + ) + + assert hit is not None and hit.cached_result == cached_message + logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once() + assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True + + +@pytest.mark.asyncio +async def test_agentic_loop_followup_cache_hit_with_converted_stream_marker_replays_as_plain_object(monkeypatch): + import litellm + from litellm.caching.caching import Cache + from litellm.types.utils import CallTypes + + async def acompletion(**kwargs): + return None + + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + kwargs = { + "model": "gpt-5.6", + "messages": [{"role": "user", "content": "run the code"}], + "caching": True, + "stream": False, + "_code_interpreter_interception_converted_stream": True, + "_agentic_loop_depth": 1, + } + await litellm.cache.async_add_cache( + litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "done"}}]), **kwargs + ) + handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now()) + logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False) + logging_obj.async_success_handler = AsyncMock() + logging_obj.handle_sync_success_callbacks_for_async_calls = MagicMock() + + hit = await handler._async_get_cache( + model="gpt-5.6", + original_function=acompletion, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.acompletion.value, + kwargs=kwargs, + args=(), + ) + + assert hit is not None and isinstance(hit.cached_result, litellm.ModelResponse) + assert hit.cached_result.choices[0].message.content == "done" + logging_obj.handle_sync_success_callbacks_for_async_calls.assert_called_once() + assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_args.kwargs["cache_hit"] is True + + +@pytest.mark.asyncio +async def test_partial_embedding_cache_hit_sends_only_misses_and_keeps_input_order(monkeypatch): + import litellm + from litellm import CustomLLM + from litellm.caching.caching import Cache + from litellm.types.utils import Embedding, EmbeddingResponse + + class RecordingEmbedder(CustomLLM): + provider_inputs: tuple[tuple[str, ...], ...] = () + + async def aembedding(self, model, input, model_response, **kwargs) -> EmbeddingResponse: + self.provider_inputs = (*self.provider_inputs, tuple(input)) + return EmbeddingResponse( + model=model, + data=[ + Embedding(embedding=[float(len(text))], index=idx, object="embedding") + for idx, text in enumerate(input) + ], + ) + + embedder = RecordingEmbedder() + monkeypatch.setattr(litellm, "custom_provider_map", [{"provider": "recording-embedder", "custom_handler": embedder}]) + monkeypatch.setattr(litellm, "provider_list", [*litellm.provider_list, "recording-embedder"]) + monkeypatch.setattr(litellm, "_custom_providers", [*litellm._custom_providers, "recording-embedder"]) + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + + await litellm.aembedding(model="recording-embedder/m", input=["aa", "bbbb"]) + await asyncio.gather(*_PENDING_CACHE_WRITES) + mixed_input = ["c", "aa", "ddd", "bbbb", "eeeee"] + response = await litellm.aembedding(model="recording-embedder/m", input=mixed_input) + await asyncio.gather(*_PENDING_CACHE_WRITES) + + assert embedder.provider_inputs == (("aa", "bbbb"), ("c", "ddd", "eeeee")), embedder.provider_inputs + assert [item["index"] for item in response.data] == [0, 1, 2, 3, 4] + assert [item["embedding"] for item in response.data] == [[float(len(text))] for text in mixed_input] + assert response._hidden_params["cache_hit"] is True, "a partial hit must still be reported as a cache hit" + + repeat = await litellm.aembedding(model="recording-embedder/m", input=mixed_input) + + assert len(embedder.provider_inputs) == 2, embedder.provider_inputs + assert [item["embedding"] for item in repeat.data] == [[float(len(text))] for text in mixed_input] diff --git a/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py b/tests/unit/caching/test_check_and_fix_namespace_none_guard.py similarity index 100% rename from tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py rename to tests/unit/caching/test_check_and_fix_namespace_none_guard.py diff --git a/tests/test_litellm/caching/test_disk_cache.py b/tests/unit/caching/test_disk_cache.py similarity index 100% rename from tests/test_litellm/caching/test_disk_cache.py rename to tests/unit/caching/test_disk_cache.py diff --git a/tests/test_litellm/caching/test_dual_cache.py b/tests/unit/caching/test_dual_cache.py similarity index 100% rename from tests/test_litellm/caching/test_dual_cache.py rename to tests/unit/caching/test_dual_cache.py diff --git a/tests/test_litellm/caching/test_embedding_router.py b/tests/unit/caching/test_embedding_router.py similarity index 100% rename from tests/test_litellm/caching/test_embedding_router.py rename to tests/unit/caching/test_embedding_router.py diff --git a/tests/test_litellm/caching/test_evicted_client_closer.py b/tests/unit/caching/test_evicted_client_closer.py similarity index 100% rename from tests/test_litellm/caching/test_evicted_client_closer.py rename to tests/unit/caching/test_evicted_client_closer.py diff --git a/tests/test_litellm/caching/test_gcs_cache.py b/tests/unit/caching/test_gcs_cache.py similarity index 100% rename from tests/test_litellm/caching/test_gcs_cache.py rename to tests/unit/caching/test_gcs_cache.py diff --git a/tests/test_litellm/caching/test_in_memory_cache.py b/tests/unit/caching/test_in_memory_cache.py similarity index 100% rename from tests/test_litellm/caching/test_in_memory_cache.py rename to tests/unit/caching/test_in_memory_cache.py diff --git a/tests/test_litellm/caching/test_llm_caching_handler.py b/tests/unit/caching/test_llm_caching_handler.py similarity index 100% rename from tests/test_litellm/caching/test_llm_caching_handler.py rename to tests/unit/caching/test_llm_caching_handler.py diff --git a/tests/test_litellm/caching/test_llm_client_cache_e2e.py b/tests/unit/caching/test_llm_client_cache_e2e.py similarity index 100% rename from tests/test_litellm/caching/test_llm_client_cache_e2e.py rename to tests/unit/caching/test_llm_client_cache_e2e.py diff --git a/tests/test_litellm/caching/test_qdrant_semantic_cache.py b/tests/unit/caching/test_qdrant_semantic_cache.py similarity index 99% rename from tests/test_litellm/caching/test_qdrant_semantic_cache.py rename to tests/unit/caching/test_qdrant_semantic_cache.py index ca7303e4c6d..4f18fb1bca6 100644 --- a/tests/test_litellm/caching/test_qdrant_semantic_cache.py +++ b/tests/unit/caching/test_qdrant_semantic_cache.py @@ -1033,7 +1033,7 @@ def test_qdrant_semantic_cache_defaults_embedding_timeout(): @pytest.mark.asyncio async def test_qdrant_async_embedding_truncates_off_the_event_loop(monkeypatch): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/unit/caching/test_redis_cache.py similarity index 100% rename from tests/test_litellm/caching/test_redis_cache.py rename to tests/unit/caching/test_redis_cache.py diff --git a/tests/test_litellm/caching/test_redis_cluster_cache.py b/tests/unit/caching/test_redis_cluster_cache.py similarity index 100% rename from tests/test_litellm/caching/test_redis_cluster_cache.py rename to tests/unit/caching/test_redis_cluster_cache.py diff --git a/tests/test_litellm/caching/test_redis_cluster_node_isolation.py b/tests/unit/caching/test_redis_cluster_node_isolation.py similarity index 100% rename from tests/test_litellm/caching/test_redis_cluster_node_isolation.py rename to tests/unit/caching/test_redis_cluster_node_isolation.py diff --git a/tests/test_litellm/caching/test_redis_connection_pool.py b/tests/unit/caching/test_redis_connection_pool.py similarity index 100% rename from tests/test_litellm/caching/test_redis_connection_pool.py rename to tests/unit/caching/test_redis_connection_pool.py diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/unit/caching/test_redis_semantic_cache.py similarity index 99% rename from tests/test_litellm/caching/test_redis_semantic_cache.py rename to tests/unit/caching/test_redis_semantic_cache.py index de253b4f10b..461689165bb 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/unit/caching/test_redis_semantic_cache.py @@ -1392,7 +1392,7 @@ def test_redis_semantic_cache_defaults_embedding_timeout(): @pytest.mark.asyncio async def test_redis_async_embedding_truncates_off_the_event_loop(monkeypatch): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/caching/test_s3_cache.py b/tests/unit/caching/test_s3_cache.py similarity index 100% rename from tests/test_litellm/caching/test_s3_cache.py rename to tests/unit/caching/test_s3_cache.py diff --git a/tests/test_litellm/caching/test_valkey_semantic_cache.py b/tests/unit/caching/test_valkey_semantic_cache.py similarity index 100% rename from tests/test_litellm/caching/test_valkey_semantic_cache.py rename to tests/unit/caching/test_valkey_semantic_cache.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/__init__.py b/tests/unit/expected_responses_api_request/__init__.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/__init__.py rename to tests/unit/expected_responses_api_request/__init__.py diff --git a/tests/test_litellm/expected_responses_api_request/azure_shell_tool.json b/tests/unit/expected_responses_api_request/azure_shell_tool.json similarity index 100% rename from tests/test_litellm/expected_responses_api_request/azure_shell_tool.json rename to tests/unit/expected_responses_api_request/azure_shell_tool.json diff --git a/tests/test_litellm/expected_responses_api_request/context_management_and_shell.json b/tests/unit/expected_responses_api_request/context_management_and_shell.json similarity index 100% rename from tests/test_litellm/expected_responses_api_request/context_management_and_shell.json rename to tests/unit/expected_responses_api_request/context_management_and_shell.json diff --git a/tests/unit/integrations/compression_interception/test_compression_interception_handler.py b/tests/unit/integrations/compression_interception/test_compression_interception_handler.py index e66cd654f93..d7a1d6f14e1 100644 --- a/tests/unit/integrations/compression_interception/test_compression_interception_handler.py +++ b/tests/unit/integrations/compression_interception/test_compression_interception_handler.py @@ -528,7 +528,7 @@ async def test_pre_call_hook_no_compression_records_no_savings(monkeypatch): @pytest.mark.asyncio async def test_pre_call_hook_counts_tokens_off_the_event_loop(): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/rust_bridge/__init__.py b/tests/unit/integrations/zerobus/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/__init__.py rename to tests/unit/integrations/zerobus/__init__.py diff --git a/tests/unit/integrations/zerobus/test_zerobus_client.py b/tests/unit/integrations/zerobus/test_zerobus_client.py new file mode 100644 index 00000000000..ae7610536f3 --- /dev/null +++ b/tests/unit/integrations/zerobus/test_zerobus_client.py @@ -0,0 +1,258 @@ +import base64 +import json +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from itertools import chain, repeat + +import httpx +import pytest + +from litellm.integrations.zerobus.client import ZerobusIngestClient +from litellm.types.integrations.zerobus import ZerobusAccessToken, ZerobusConnection, ZerobusIngestFailure + +CONNECTION = ZerobusConnection( + workspace_url="https://dbc-a1b2c3d4-e5f6.cloud.databricks.com/", + workspace_id="1234567890123456", + server_endpoint="https://1234567890123456.zerobus.us-west-2.cloud.databricks.com", + client_id="sp-client-id", + client_secret="sp-client-secret", + table_name="main.litellm.traces", +) +ROWS = ({"id": "a", "model": "gpt-4o"}, {"id": "b", "model": "gpt-4o"}) + + +def _token(value: str = "tok-1", expires_in: float = 3600) -> httpx.Response: + return httpx.Response(200, text=json.dumps({"access_token": value, "expires_in": expires_in})) + + +def _accepted() -> httpx.Response: + return httpx.Response(200, text="{}") + + +@dataclass(frozen=True, slots=True) +class TokenCall: + url: str + data: Mapping[str, str] + headers: Mapping[str, str] + + +@dataclass(frozen=True, slots=True) +class InsertCall: + url: str + content: bytes + headers: Mapping[str, str] + + +def _results(results: Sequence[httpx.Response | Exception]) -> Iterator[httpx.Response | Exception]: + """Results are served in order, and the last one repeats.""" + return chain(results[:-1], repeat(results[-1])) + + +class FakeHTTPClient: + """Stands in for AsyncHTTPHandler, including its habit of raising on error statuses.""" + + def __init__( + self, + token: Sequence[httpx.Response | Exception] = (), + insert: Sequence[httpx.Response | Exception] = (), + ) -> None: + self.token_results = _results(token or (_token(),)) + self.insert_results = _results(insert or (_accepted(),)) + self.token_calls: tuple[TokenCall, ...] = () + self.insert_calls: tuple[InsertCall, ...] = () + + async def post( + self, + url: str, + data: Mapping[str, str] | None = None, + content: bytes | None = None, + headers: Mapping[str, str] | None = None, + ) -> httpx.Response: + if url.endswith("/oidc/v1/token"): + self.token_calls = (*self.token_calls, TokenCall(url, data or {}, headers or {})) + return _raise_like_the_handler(next(self.token_results), url) + self.insert_calls = (*self.insert_calls, InsertCall(url, content or b"", headers or {})) + return _raise_like_the_handler(next(self.insert_results), url) + + +def _raise_like_the_handler(result: httpx.Response | Exception, url: str) -> httpx.Response: + if isinstance(result, Exception): + raise result + if result.status_code >= 300: + raise httpx.HTTPStatusError( + "boom", + request=httpx.Request("POST", url), + response=httpx.Response(result.status_code, text=result.text), + ) + return result + + +class FakeClock: + def __init__(self, now: float = 1_000.0) -> None: + self.now = now + + def __call__(self) -> float: + return self.now + + +def _client(http_client: FakeHTTPClient, clock: FakeClock | None = None) -> ZerobusIngestClient: + return ZerobusIngestClient(connection=CONNECTION, http_client=http_client, clock=clock or FakeClock()) + + +@pytest.mark.asyncio +async def test_rows_are_posted_as_one_json_list_to_the_table_insert_endpoint(): + http_client = FakeHTTPClient() + + outcome = await _client(http_client).insert(ROWS) + + assert outcome is None + (call,) = http_client.insert_calls + # Insert endpoint per the Zerobus Ingest docs, read 2026-09-19: + # https://docs.databricks.com/aws/en/ingestion/lakeflow-connect/zerobus-ingest + assert call.url == ( + "https://1234567890123456.zerobus.us-west-2.cloud.databricks.com/zerobus/v1/tables/main.litellm.traces/insert" + ) + assert json.loads(call.content) == [{"id": "a", "model": "gpt-4o"}, {"id": "b", "model": "gpt-4o"}] + assert call.headers["Content-Type"] == "application/json" + assert call.headers["Authorization"] == "Bearer tok-1" + + +@pytest.mark.asyncio +async def test_the_token_is_minted_for_the_zerobus_resource_with_the_table_privileges(): + """Zerobus refuses a plain workspace token: it must name its own resource and the table's UC privileges.""" + http_client = FakeHTTPClient() + + await _client(http_client).insert(ROWS) + + (call,) = http_client.token_calls + # Token form per the Zerobus Ingest docs (REST API authentication), read 2026-09-19: + # https://docs.databricks.com/aws/en/ingestion/lakeflow-connect/zerobus-ingest + assert call.url == "https://dbc-a1b2c3d4-e5f6.cloud.databricks.com/oidc/v1/token" + assert call.data["grant_type"] == "client_credentials" + assert call.data["scope"] == "all-apis" + assert call.data["resource"] == "api://databricks/workspaces/1234567890123456/zerobusDirectWriteApi" + details = json.loads(call.data["authorization_details"]) + assert [(d["object_type"], d["object_full_path"], d["privileges"]) for d in details] == [ + ("CATALOG", "main", ["USE CATALOG"]), + ("SCHEMA", "main.litellm", ["USE SCHEMA"]), + ("TABLE", "main.litellm.traces", ["SELECT", "MODIFY"]), + ] + assert all(d["type"] == "unity_catalog_privileges" for d in details) + + +@pytest.mark.asyncio +async def test_the_service_principal_authenticates_with_http_basic(): + http_client = FakeHTTPClient() + + await _client(http_client).insert(ROWS) + + scheme, credentials = http_client.token_calls[0].headers["Authorization"].split(" ") + assert scheme == "Basic" + assert base64.b64decode(credentials).decode() == "sp-client-id:sp-client-secret" + + +def test_the_client_secret_and_minted_token_stay_out_of_reprs_and_tracebacks(): + token = ZerobusAccessToken(value="tok-secret", expires_at=1.0) + + assert "sp-client-secret" not in repr(CONNECTION) + assert "sp-client-id" in repr(CONNECTION) + assert "tok-secret" not in repr(token) + assert "expires_at=1.0" in repr(token) + + +@pytest.mark.asyncio +async def test_the_token_is_reused_across_inserts_until_it_nears_expiry(): + clock = FakeClock(now=1_000.0) + http_client = FakeHTTPClient(token=[_token("tok-1", expires_in=600), _token("tok-2")]) + client = _client(http_client, clock) + + await client.insert(ROWS) + clock.now = 1_000.0 + 600 - 61 + await client.insert(ROWS) + clock.now = 1_000.0 + 600 - 59 + await client.insert(ROWS) + + assert len(http_client.token_calls) == 2 + assert [call.headers["Authorization"] for call in http_client.insert_calls] == [ + "Bearer tok-1", + "Bearer tok-1", + "Bearer tok-2", + ] + + +@pytest.mark.asyncio +async def test_a_401_discards_the_token_so_the_next_insert_mints_a_fresh_one(): + http_client = FakeHTTPClient( + token=[_token("tok-1"), _token("tok-2")], + insert=[httpx.Response(401, text="expired"), _accepted()], + ) + client = _client(http_client) + + first = await client.insert(ROWS) + second = await client.insert(ROWS) + + assert first == ZerobusIngestFailure(detail="insert returned 401, token discarded", retryable=True) + assert second is None + assert http_client.insert_calls[1].headers["Authorization"] == "Bearer tok-2" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [429, 500, 503]) +async def test_a_transient_insert_status_is_retryable(status: int): + http_client = FakeHTTPClient(insert=[httpx.Response(status, text="later")]) + + outcome = await _client(http_client).insert(ROWS) + + assert isinstance(outcome, ZerobusIngestFailure) + assert outcome.retryable is True + assert str(status) in outcome.detail + + +@pytest.mark.asyncio +async def test_a_schema_rejection_is_not_retryable_and_says_why(): + http_client = FakeHTTPClient(insert=[httpx.Response(400, text="unknown column foo")]) + + outcome = await _client(http_client).insert(ROWS) + + assert outcome == ZerobusIngestFailure(detail="insert returned 400: unknown column foo", retryable=False) + + +@pytest.mark.asyncio +async def test_a_network_failure_on_insert_is_retryable(): + http_client = FakeHTTPClient(insert=[httpx.ConnectError("connection refused")]) + + outcome = await _client(http_client).insert(ROWS) + + assert isinstance(outcome, ZerobusIngestFailure) + assert outcome.retryable is True + + +@pytest.mark.asyncio +async def test_bad_credentials_fail_the_insert_without_posting_rows(): + http_client = FakeHTTPClient(token=[httpx.Response(401, text="invalid_client")]) + + outcome = await _client(http_client).insert(ROWS) + + assert outcome == ZerobusIngestFailure(detail="token request returned 401: invalid_client", retryable=False) + assert http_client.insert_calls == () + + +@pytest.mark.asyncio +async def test_a_token_endpoint_outage_is_retryable(): + http_client = FakeHTTPClient(token=[httpx.Response(503, text="try later")]) + + outcome = await _client(http_client).insert(ROWS) + + assert isinstance(outcome, ZerobusIngestFailure) + assert outcome.retryable is True + + +@pytest.mark.asyncio +async def test_a_token_response_without_a_token_is_reported_not_raised(): + http_client = FakeHTTPClient(token=[httpx.Response(200, text='{"token_type": "Bearer"}')]) + + outcome = await _client(http_client).insert(ROWS) + + assert isinstance(outcome, ZerobusIngestFailure) + assert outcome.retryable is False + assert "token response" in outcome.detail diff --git a/tests/unit/integrations/zerobus/test_zerobus_logger.py b/tests/unit/integrations/zerobus/test_zerobus_logger.py new file mode 100644 index 00000000000..a85a9e2e6a0 --- /dev/null +++ b/tests/unit/integrations/zerobus/test_zerobus_logger.py @@ -0,0 +1,392 @@ +import asyncio +from collections.abc import Callable, Iterator, Mapping, Sequence +from itertools import chain, repeat + +import pytest + +import litellm +from litellm.integrations.zerobus.client import ZerobusIngestError +from litellm.integrations.zerobus.logger import ZerobusLogger, connection_for +from litellm.types.integrations.zerobus import ZerobusIngestFailure, ZerobusInitParams + +WORKSPACE_URL = "https://dbc-a1b2c3d4-e5f6.cloud.databricks.com" +SERVER_ENDPOINT = "https://1234567890123456.zerobus.us-west-2.cloud.databricks.com" + + +Row = Mapping[str, object] + + +class FakeIngestClient: + """Records the rows each flush would have written; outcomes are served in order and the last one repeats.""" + + def __init__( + self, + outcomes: Sequence[ZerobusIngestFailure | None] = (None,), + on_insert: Callable[[], None] | None = None, + ) -> None: + self.outcomes: Iterator[ZerobusIngestFailure | None] = chain(outcomes[:-1], repeat(outcomes[-1])) + self.on_insert = on_insert + self.batches: tuple[tuple[Row, ...], ...] = () + + async def insert(self, rows: Sequence[Row]) -> ZerobusIngestFailure | None: + if self.on_insert is not None: + self.on_insert() + self.batches = (*self.batches, tuple(rows)) + return next(self.outcomes) + + def ids(self) -> tuple[object, ...]: + return tuple(row["id"] for batch in self.batches for row in batch) + + +def _logger(client: FakeIngestClient, **params: object) -> ZerobusLogger: + return ZerobusLogger(params=ZerobusInitParams.model_validate(params), client=client) + + +def _event(request_id: str, **payload: object) -> dict[str, object]: + return { + "standard_logging_object": { + "id": request_id, + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "response": {"choices": []}, + **payload, + } + } + + +async def _settle(logger: ZerobusLogger) -> None: + for _ in range(200): + await asyncio.sleep(0.001) + task = logger._batch_flush_task + if (task is None or task.done()) and not logger._flushing: + return + + +@pytest.mark.asyncio +async def test_a_full_batch_is_written_as_one_insert_of_table_rows(): + client = FakeIngestClient() + logger = _logger(client, batch_size=3) + + for request_id in ("a", "b", "c"): + await logger.async_log_success_event(_event(request_id), None, None, None) + + await _settle(logger) + assert len(client.batches) == 1 + assert client.ids() == ("a", "b", "c") + assert client.batches[0][0]["model"] == "gpt-4o" + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_rows_are_held_until_the_batch_is_full(): + client = FakeIngestClient() + logger = _logger(client, batch_size=3) + + await logger.async_log_success_event(_event("a"), None, None, None) + + assert client.batches == () + assert len(logger.log_queue) == 1 + + +@pytest.mark.asyncio +async def test_failed_requests_are_written_too(): + client = FakeIngestClient() + logger = _logger(client, batch_size=1) + + await logger.async_log_failure_event(_event("failed", status="failure", error_str="boom"), None, None, None) + + await _settle(logger) + assert client.ids() == ("failed",) + assert client.batches[0][0]["status"] == "failure" + assert client.batches[0][0]["error_str"] == "boom" + + +@pytest.mark.asyncio +async def test_an_event_without_a_standard_payload_is_skipped(): + client = FakeIngestClient() + logger = _logger(client, batch_size=1) + + await logger.async_log_success_event({"kwargs": "but no payload"}, None, None, None) + + assert client.batches == () + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_a_retryable_failure_keeps_the_rows_for_the_next_flush(): + client = FakeIngestClient([ZerobusIngestFailure("zerobus is down", retryable=True)]) + logger = _logger(client, batch_size=2) + + for request_id in ("a", "b"): + await logger.async_log_success_event(_event(request_id), None, None, None) + + await _settle(logger) + assert [row["id"] for row in logger.log_queue] == ["a", "b"] + + +@pytest.mark.asyncio +async def test_a_retryable_failure_surfaces_so_the_base_logger_can_preserve_it(): + client = FakeIngestClient([ZerobusIngestFailure("zerobus is down", retryable=True)]) + logger = _logger(client, batch_size=99) + logger.log_queue.append({"id": "a"}) + + with pytest.raises(ZerobusIngestError, match="zerobus is down"): + await logger.async_send_batch() + + +@pytest.mark.asyncio +async def test_a_rejected_batch_is_dropped_rather_than_blocking_the_queue(): + client = FakeIngestClient([ZerobusIngestFailure("unknown column", retryable=False)]) + logger = _logger(client, batch_size=2) + + for request_id in ("a", "b"): + await logger.async_log_success_event(_event(request_id), None, None, None) + + await _settle(logger) + assert logger.log_queue == [] + + +@pytest.mark.asyncio +async def test_a_row_that_arrives_mid_flush_is_kept_for_the_next_one(): + client = FakeIngestClient() + logger = _logger(client, batch_size=1) + client.on_insert = lambda: logger.log_queue.append({"id": "late"}) + + await logger.async_log_success_event(_event("first"), None, None, None) + + await _settle(logger) + assert client.ids() == ("first",) + assert [row["id"] for row in logger.log_queue] == ["late"] + + +@pytest.mark.asyncio +async def test_the_queue_cap_holds_while_an_insert_is_in_flight(): + """A slow insert must not let the queue grow past max_queue_size, nor disturb the in-flight head.""" + insert_started = asyncio.Event() + finish_insert = asyncio.Event() + + class SlowClient: + batches: tuple[tuple[Row, ...], ...] = () + + async def insert(self, rows: Sequence[Row]) -> None: + insert_started.set() + await finish_insert.wait() + self.batches = (*self.batches, tuple(rows)) + + client = SlowClient() + logger = ZerobusLogger(params=ZerobusInitParams(batch_size=2), client=client) + logger.max_queue_size = 3 + + for request_id in ("a", "b"): + await logger.async_log_success_event(_event(request_id), None, None, None) + await insert_started.wait() + for request_id in ("c", "d", "e"): + await logger.async_log_success_event(_event(request_id), None, None, None) + finish_insert.set() + await _settle(logger) + + assert [[row["id"] for row in batch] for batch in client.batches] == [["a", "b"]] + assert [row["id"] for row in logger.log_queue] == ["c"] + + +@pytest.mark.asyncio +async def test_a_client_error_does_not_break_the_request_path(): + class ExplodingClient: + async def insert(self, rows: Sequence[Row]) -> None: + raise RuntimeError("bug") + + logger = ZerobusLogger(params=ZerobusInitParams(batch_size=1), client=ExplodingClient()) + + await logger.async_log_success_event(_event("a"), None, None, None) + await _settle(logger) + + assert [row["id"] for row in logger.log_queue] == ["a"] + + +@pytest.mark.asyncio +async def test_turn_off_message_logging_redacts_prompts_and_responses_but_keeps_the_rest(): + client = FakeIngestClient() + logger = _logger(client, batch_size=1, turn_off_message_logging=True) + + await logger.async_log_success_event( + _event("a", prompt_tokens=10, response={"choices": [{"message": {"content": "the secret answer"}}]}), + None, + None, + None, + ) + + await _settle(logger) + (row,) = client.batches[0] + assert row["id"] == "a" + assert row["prompt_tokens"] == 10 + assert '"hi"' not in str(row["messages"]) + assert "the secret answer" not in str(row["response"]) + + +def test_connection_comes_from_the_environment_the_proxy_ui_writes(monkeypatch): + monkeypatch.setenv("ZEROBUS_WORKSPACE_URL", WORKSPACE_URL) + monkeypatch.setenv("ZEROBUS_SERVER_ENDPOINT", SERVER_ENDPOINT) + monkeypatch.setenv("ZEROBUS_CLIENT_ID", "sp-id") + monkeypatch.setenv("ZEROBUS_CLIENT_SECRET", "sp-secret") + monkeypatch.setenv("ZEROBUS_TABLE_NAME", "main.litellm.traces") + + connection = connection_for(ZerobusInitParams()) + + assert connection.workspace_url == WORKSPACE_URL + assert connection.server_endpoint == SERVER_ENDPOINT + assert connection.workspace_id == "1234567890123456" + assert connection.client_id == "sp-id" + assert connection.client_secret == "sp-secret" + assert connection.table_name == "main.litellm.traces" + + +def test_config_yaml_params_win_over_the_environment(monkeypatch): + monkeypatch.setenv("ZEROBUS_TABLE_NAME", "env.schema.table") + monkeypatch.setenv("ZEROBUS_CLIENT_SECRET", "from-env") + + connection = connection_for( + ZerobusInitParams( + workspace_url=WORKSPACE_URL, + server_endpoint=SERVER_ENDPOINT, + client_id="sp-id", + client_secret="from-config", + table_name="cfg.schema.table", + ) + ) + + assert connection.table_name == "cfg.schema.table" + assert connection.client_secret == "from-config" + + +def test_a_secret_reference_in_config_yaml_is_resolved(monkeypatch): + monkeypatch.setenv("MY_SP_SECRET", "resolved-secret") + + connection = connection_for( + ZerobusInitParams( + workspace_url=WORKSPACE_URL, + server_endpoint=SERVER_ENDPOINT, + client_id="sp-id", + client_secret="os.environ/MY_SP_SECRET", + table_name="main.litellm.traces", + ) + ) + + assert connection.client_secret == "resolved-secret" + + +def test_a_missing_setting_names_the_env_var_to_set(monkeypatch): + monkeypatch.delenv("ZEROBUS_CLIENT_SECRET", raising=False) + + with pytest.raises(ValueError, match="ZEROBUS_CLIENT_SECRET"): + connection_for( + ZerobusInitParams( + workspace_url=WORKSPACE_URL, + server_endpoint=SERVER_ENDPOINT, + client_id="sp-id", + table_name="main.litellm.traces", + ) + ) + + +def test_a_table_that_is_not_fully_qualified_is_refused(): + with pytest.raises(ValueError, match=r"catalog\.schema\.table"): + connection_for( + ZerobusInitParams( + workspace_url=WORKSPACE_URL, + server_endpoint=SERVER_ENDPOINT, + client_id="sp-id", + client_secret="sp-secret", + table_name="traces", + ) + ) + + +def test_an_endpoint_without_a_workspace_id_is_refused(): + """The token's resource needs the numeric workspace id, which only the Zerobus hostname carries.""" + with pytest.raises(ValueError, match="ZEROBUS_SERVER_ENDPOINT"): + connection_for( + ZerobusInitParams( + workspace_url=WORKSPACE_URL, + server_endpoint=WORKSPACE_URL, + client_id="sp-id", + client_secret="sp-secret", + table_name="main.litellm.traces", + ) + ) + + +def test_a_misconfigured_logger_fails_at_startup_not_at_first_flush(monkeypatch): + for name in ("WORKSPACE_URL", "SERVER_ENDPOINT", "CLIENT_ID", "CLIENT_SECRET", "TABLE_NAME"): + monkeypatch.delenv(f"ZEROBUS_{name}", raising=False) + monkeypatch.setattr(litellm, "zerobus_params", None) + + with pytest.raises(ValueError, match="ZEROBUS_"): + ZerobusLogger() + + +def test_litellm_zerobus_params_configure_the_logger(monkeypatch): + monkeypatch.setattr( + litellm, + "zerobus_params", + { + "workspace_url": WORKSPACE_URL, + "server_endpoint": SERVER_ENDPOINT, + "client_id": "sp-id", + "client_secret": "sp-secret", + "table_name": "main.litellm.traces", + "batch_size": 7, + "flush_interval": 3, + }, + ) + + logger = ZerobusLogger() + + assert logger.batch_size == 7 + assert logger.flush_interval == 3 + assert logger.client.connection.table_name == "main.litellm.traces" + + +def test_the_client_is_kept_while_the_connection_is_unchanged_and_rebuilt_when_it_changes(monkeypatch): + """The client caches its token, so it must survive across flushes, yet a UI edit must take effect.""" + monkeypatch.setenv("ZEROBUS_WORKSPACE_URL", WORKSPACE_URL) + monkeypatch.setenv("ZEROBUS_SERVER_ENDPOINT", SERVER_ENDPOINT) + monkeypatch.setenv("ZEROBUS_CLIENT_ID", "sp-id") + monkeypatch.setenv("ZEROBUS_CLIENT_SECRET", "sp-secret") + monkeypatch.setenv("ZEROBUS_TABLE_NAME", "main.litellm.traces") + monkeypatch.setattr(litellm, "zerobus_params", None) + logger = ZerobusLogger() + + first = logger.client + unchanged = logger.client + monkeypatch.setenv("ZEROBUS_TABLE_NAME", "main.litellm.traces_v2") + rebuilt = logger.client + + assert unchanged is first + assert rebuilt is not first + assert rebuilt.connection.table_name == "main.litellm.traces_v2" + + +def test_callbacks_zerobus_builds_one_logger_and_reuses_it(monkeypatch): + """`litellm_settings.callbacks: ["zerobus"]` goes through litellm_logging, which must hand back one instance.""" + from litellm.litellm_core_utils import litellm_logging as logging_module + + monkeypatch.setenv("ZEROBUS_WORKSPACE_URL", WORKSPACE_URL) + monkeypatch.setenv("ZEROBUS_SERVER_ENDPOINT", SERVER_ENDPOINT) + monkeypatch.setenv("ZEROBUS_CLIENT_ID", "sp-id") + monkeypatch.setenv("ZEROBUS_CLIENT_SECRET", "sp-secret") + monkeypatch.setenv("ZEROBUS_TABLE_NAME", "main.litellm.traces") + monkeypatch.setattr(litellm, "zerobus_params", None) + monkeypatch.setattr(logging_module, "_in_memory_loggers", []) + + assert logging_module.get_custom_logger_compatible_class("zerobus") is None + + first = logging_module._init_custom_logger_compatible_class( + logging_integration="zerobus", internal_usage_cache=None, llm_router=None, custom_logger_init_args={} + ) + second = logging_module._init_custom_logger_compatible_class( + logging_integration="zerobus", internal_usage_cache=None, llm_router=None, custom_logger_init_args={} + ) + + assert isinstance(first, ZerobusLogger) + assert second is first + assert logging_module.get_custom_logger_compatible_class("zerobus") is first diff --git a/tests/unit/integrations/zerobus/test_zerobus_row.py b/tests/unit/integrations/zerobus/test_zerobus_row.py new file mode 100644 index 00000000000..b73c3bae48f --- /dev/null +++ b/tests/unit/integrations/zerobus/test_zerobus_row.py @@ -0,0 +1,139 @@ +import json + +from litellm.integrations.zerobus.row import TRACE_TABLE_COLUMNS, create_table_sql, trace_row + + +def _payload() -> dict[str, object]: + return { + "id": "chatcmpl-1", + "trace_id": "trace-1", + "session_id": "session-1", + "litellm_call_id": "call-1", + "call_type": "acompletion", + "status": "success", + "model": "gpt-4o", + "model_group": "gpt-4o-group", + "custom_llm_provider": "openai", + "api_base": "https://api.openai.com", + "stream": False, + "cache_hit": None, + "startTime": 1_700_000_000.25, + "endTime": 1_700_000_001.5, + "completionStartTime": 1_700_000_000.75, + "response_time": 1.25, + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + "response_cost": 0.0015, + "saved_cache_cost": 0.0, + "end_user": "end-user-1", + "requester_ip_address": "10.0.0.1", + "user_agent": "curl/8", + "request_tags": ["prod"], + "messages": [{"role": "user", "content": "hi"}], + "response": {"choices": [{"message": {"role": "assistant", "content": "hello"}}]}, + "error_str": None, + "error_information": None, + "metadata": { + "user_api_key_hash": "hash-1", + "user_api_key_alias": "alias-1", + "user_api_key_team_id": "team-1", + "user_api_key_team_alias": "team-alias-1", + "user_api_key_user_id": "user-1", + "user_api_key_org_id": "org-1", + }, + "model_parameters": {"temperature": 0.2}, + "hidden_params": {"response_cost": 0.0015}, + "guardrail_information": None, + "cost_breakdown": {"input_cost": 0.001, "output_cost": 0.0005}, + } + + +def test_every_row_has_exactly_the_documented_columns(): + """Zerobus rejects a record naming a column the table lacks, so the row and the DDL must agree.""" + assert tuple(trace_row(_payload())) == tuple(TRACE_TABLE_COLUMNS) + assert tuple(trace_row({})) == tuple(TRACE_TABLE_COLUMNS) + + +def test_scalars_land_in_their_columns(): + row = trace_row(_payload()) + + assert row["id"] == "chatcmpl-1" + assert row["trace_id"] == "trace-1" + assert row["status"] == "success" + assert row["model"] == "gpt-4o" + assert row["stream"] is False + assert row["prompt_tokens"] == 10 + assert row["total_tokens"] == 15 + assert row["response_cost"] == 0.0015 + assert row["end_user"] == "end-user-1" + + +def test_key_and_team_identity_is_lifted_out_of_metadata(): + """Filtering spend by team or key is the main query, so those live in their own columns.""" + row = trace_row(_payload()) + + assert row["api_key_hash"] == "hash-1" + assert row["api_key_alias"] == "alias-1" + assert row["team_id"] == "team-1" + assert row["team_alias"] == "team-alias-1" + assert row["user_id"] == "user-1" + assert row["org_id"] == "org-1" + + +def test_timestamps_become_epoch_microseconds(): + row = trace_row(_payload()) + + assert row["start_time"] == 1_700_000_000_250_000 + assert row["end_time"] == 1_700_000_001_500_000 + assert row["completion_start_time"] == 1_700_000_000_750_000 + + +def test_a_zero_timestamp_is_null_rather_than_1970(): + """LiteLLM leaves completionStartTime at 0 when there is no first token, which is not a real time.""" + row = trace_row({**_payload(), "completionStartTime": 0}) + + assert row["completion_start_time"] is None + + +def test_nested_fields_are_json_text_for_the_variant_columns(): + row = trace_row(_payload()) + + assert json.loads(str(row["messages"])) == [{"role": "user", "content": "hi"}] + assert json.loads(str(row["metadata"]))["user_api_key_team_id"] == "team-1" + assert json.loads(str(row["request_tags"])) == ["prod"] + assert json.loads(str(row["cost_breakdown"])) == {"input_cost": 0.001, "output_cost": 0.0005} + + +def test_missing_and_null_fields_are_null(): + row = trace_row({**_payload(), "messages": None, "guardrail_information": None}) + + assert row["messages"] is None + assert row["guardrail_information"] is None + assert row["error_str"] is None + assert row["cache_hit"] is None + + +def test_a_wrongly_typed_field_is_null_instead_of_a_rejected_record(): + """One odd payload must not poison the whole batch: the table type wins.""" + row = trace_row({**_payload(), "prompt_tokens": "ten", "stream": "yes", "startTime": "now"}) + + assert row["prompt_tokens"] is None + assert row["stream"] is None + assert row["start_time"] is None + + +def test_the_row_survives_a_json_round_trip_unchanged(): + row = trace_row(_payload()) + + assert json.loads(json.dumps(dict(row))) == dict(row) + + +def test_create_table_sql_declares_every_column_with_its_type(): + sql = create_table_sql("main.litellm.traces") + + assert sql.startswith("CREATE TABLE main.litellm.traces (") + assert " start_time TIMESTAMP," in sql + assert " messages VARIANT," in sql + assert " cost_breakdown VARIANT\n);" in sql + assert sql.count(",") == len(TRACE_TABLE_COLUMNS) - 1 diff --git a/tests/unit/litellm_core_utils/conftest.py b/tests/unit/litellm_core_utils/conftest.py new file mode 100644 index 00000000000..2a1e1f6382c --- /dev/null +++ b/tests/unit/litellm_core_utils/conftest.py @@ -0,0 +1,15 @@ +import importlib + +import pytest + +from tests.unit.litellm_core_utils.fake_secret_vault import FakeSecretVault + + +@pytest.fixture(autouse=True, scope="session") +def bundled_tiktoken_cache() -> None: + importlib.import_module("litellm.litellm_core_utils.default_encoding") + + +@pytest.fixture +def secret_vault_factory() -> type[FakeSecretVault]: + return FakeSecretVault diff --git a/tests/test_litellm/litellm_core_utils/event_loop_lag.py b/tests/unit/litellm_core_utils/event_loop_lag.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/event_loop_lag.py rename to tests/unit/litellm_core_utils/event_loop_lag.py diff --git a/tests/unit/litellm_core_utils/fake_secret_vault.py b/tests/unit/litellm_core_utils/fake_secret_vault.py new file mode 100644 index 00000000000..75e9d16e9ed --- /dev/null +++ b/tests/unit/litellm_core_utils/fake_secret_vault.py @@ -0,0 +1,67 @@ +from litellm.litellm_core_utils.cli_keyring import ( + KeyringDiscardsWrites, + KeyringUnreachable, + KeyringUnusable, + SecretErase, + SecretErased, + SecretFound, + SecretMissing, + SecretRead, + SecretStored, + SecretStranded, + SecretWrite, +) + + +class FakeSecretVault: + """In-memory stand-in for the OS keychain, injected wherever CLI credential storage is exercised. + + `available=False` models a keychain that is locked or has no backend, `writable=False` one that + refuses to store, `erasable=False` one that will not release what it already holds, and `failure` + picks which unusable state those report. `discards=True` is keyring's null backend, which answers + reads and erases like any other yet keeps nothing it is given, so only writes report it. + """ + + def __init__( + self, + blob: str | None = None, + *, + available: bool = True, + writable: bool = True, + erasable: bool = True, + discards: bool = False, + failure: KeyringUnusable = KeyringUnreachable(), + ) -> None: + self.blob: str | None = blob + self.available: bool = available + self.writable: bool = writable + self.erasable: bool = erasable + self.discards: bool = discards + self.failure: KeyringUnusable = failure + self.reads: int = 0 + self.writes: list[str] = [] + self.erases: int = 0 + + def read(self) -> SecretRead: + self.reads += 1 + if not self.available: + return self.failure + return SecretMissing() if self.blob is None else SecretFound(self.blob) + + def write(self, blob: str) -> SecretWrite: + self.writes.append(blob) + if not (self.available and self.writable): + return self.failure + if self.discards: + return KeyringDiscardsWrites() + self.blob = blob + return SecretStored() + + def erase(self) -> SecretErase: + self.erases += 1 + if not self.available: + return self.failure + if not self.erasable: + return SecretStranded() if self.blob is not None else SecretErased() + self.blob = None + return SecretErased() diff --git a/tests/test_litellm/rust_bridge/chat_completions/__init__.py b/tests/unit/litellm_core_utils/llm_cost_calc/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/chat_completions/__init__.py rename to tests/unit/litellm_core_utils/llm_cost_calc/__init__.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_azure_assistant_cost_tracking.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_guardrail_cost.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_openai_cache_write_cost.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_responses_cache_cost_breakdown.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking_dict_safety.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking_dict_safety.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking_dict_safety.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking_dict_safety.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_usage_object_transformation.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_usage_object_transformation.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_usage_object_transformation.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_usage_object_transformation.py diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py rename to tests/unit/litellm_core_utils/llm_cost_calc/test_zero_cost_diagnostic.py diff --git a/tests/test_litellm/litellm_core_utils/llm_response_utils/test_get_api_base.py b/tests/unit/litellm_core_utils/llm_response_utils/test_get_api_base.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/llm_response_utils/test_get_api_base.py rename to tests/unit/litellm_core_utils/llm_response_utils/test_get_api_base.py diff --git a/tests/test_litellm/litellm_core_utils/messages_with_counts.py b/tests/unit/litellm_core_utils/messages_with_counts.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/messages_with_counts.py rename to tests/unit/litellm_core_utils/messages_with_counts.py diff --git a/tests/test_litellm/rust_bridge/messages/__init__.py b/tests/unit/litellm_core_utils/prompt_templates/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/messages/__init__.py rename to tests/unit/litellm_core_utils/prompt_templates/__init__.py diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_bedrock_converse_strict_tools_opus_47_48.py b/tests/unit/litellm_core_utils/prompt_templates/test_bedrock_converse_strict_tools_opus_47_48.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/prompt_templates/test_bedrock_converse_strict_tools_opus_47_48.py rename to tests/unit/litellm_core_utils/prompt_templates/test_bedrock_converse_strict_tools_opus_47_48.py diff --git a/tests/test_litellm/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 similarity index 100% rename from tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py rename to tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py diff --git a/tests/test_litellm/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 similarity index 100% rename from tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py rename to tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py rename to tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py diff --git a/tests/test_litellm/rust_bridge/ocr/__init__.py b/tests/unit/litellm_core_utils/specialty_caches/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/ocr/__init__.py rename to tests/unit/litellm_core_utils/specialty_caches/__init__.py diff --git a/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py b/tests/unit/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py rename to tests/unit/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py diff --git a/tests/test_litellm/litellm_core_utils/test_agentic_followup_kwargs.py b/tests/unit/litellm_core_utils/test_agentic_followup_kwargs.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_agentic_followup_kwargs.py rename to tests/unit/litellm_core_utils/test_agentic_followup_kwargs.py diff --git a/tests/test_litellm/litellm_core_utils/test_anthropic_dedup_factory.py b/tests/unit/litellm_core_utils/test_anthropic_dedup_factory.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_anthropic_dedup_factory.py rename to tests/unit/litellm_core_utils/test_anthropic_dedup_factory.py diff --git a/tests/test_litellm/litellm_core_utils/test_api_route_to_call_types.py b/tests/unit/litellm_core_utils/test_api_route_to_call_types.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_api_route_to_call_types.py rename to tests/unit/litellm_core_utils/test_api_route_to_call_types.py diff --git a/tests/test_litellm/litellm_core_utils/test_audio_utils.py b/tests/unit/litellm_core_utils/test_audio_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_audio_utils.py rename to tests/unit/litellm_core_utils/test_audio_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_aws_partition.py b/tests/unit/litellm_core_utils/test_aws_partition.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_aws_partition.py rename to tests/unit/litellm_core_utils/test_aws_partition.py diff --git a/tests/test_litellm/litellm_core_utils/test_bedrock_converse_dedup_factory.py b/tests/unit/litellm_core_utils/test_bedrock_converse_dedup_factory.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_bedrock_converse_dedup_factory.py rename to tests/unit/litellm_core_utils/test_bedrock_converse_dedup_factory.py diff --git a/tests/test_litellm/litellm_core_utils/test_bug_report.py b/tests/unit/litellm_core_utils/test_bug_report.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_bug_report.py rename to tests/unit/litellm_core_utils/test_bug_report.py diff --git a/tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py b/tests/unit/litellm_core_utils/test_chat_completion_agentic_loop.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_chat_completion_agentic_loop.py rename to tests/unit/litellm_core_utils/test_chat_completion_agentic_loop.py diff --git a/tests/test_litellm/litellm_core_utils/test_classifier_logging.py b/tests/unit/litellm_core_utils/test_classifier_logging.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_classifier_logging.py rename to tests/unit/litellm_core_utils/test_classifier_logging.py diff --git a/tests/test_litellm/litellm_core_utils/test_cli_token_utils.py b/tests/unit/litellm_core_utils/test_cli_token_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_cli_token_utils.py rename to tests/unit/litellm_core_utils/test_cli_token_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_cloud_storage_security.py b/tests/unit/litellm_core_utils/test_cloud_storage_security.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_cloud_storage_security.py rename to tests/unit/litellm_core_utils/test_cloud_storage_security.py diff --git a/tests/test_litellm/litellm_core_utils/test_codestral_provider_routing.py b/tests/unit/litellm_core_utils/test_codestral_provider_routing.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_codestral_provider_routing.py rename to tests/unit/litellm_core_utils/test_codestral_provider_routing.py diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/unit/litellm_core_utils/test_core_helpers.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_core_helpers.py rename to tests/unit/litellm_core_utils/test_core_helpers.py diff --git a/tests/test_litellm/litellm_core_utils/test_coroutine_checker.py b/tests/unit/litellm_core_utils/test_coroutine_checker.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_coroutine_checker.py rename to tests/unit/litellm_core_utils/test_coroutine_checker.py diff --git a/tests/test_litellm/litellm_core_utils/test_dd_tracing.py b/tests/unit/litellm_core_utils/test_dd_tracing.py similarity index 85% rename from tests/test_litellm/litellm_core_utils/test_dd_tracing.py rename to tests/unit/litellm_core_utils/test_dd_tracing.py index b55ade5225d..30cae45e250 100644 --- a/tests/test_litellm/litellm_core_utils/test_dd_tracing.py +++ b/tests/unit/litellm_core_utils/test_dd_tracing.py @@ -55,18 +55,6 @@ def test_dd_tracer_when_package_not_exists(): assert result == "test" -def test_null_tracer_context_manager(): - """ - Test that the context manager works without raising exceptions when should_use_dd_tracer is False - """ - with patch("litellm.litellm_core_utils.dd_tracing.should_use_dd_tracer", False): - # Test that the context manager works without raising exceptions - with dd_tracer.trace("test_operation") as span: - # Test that we can call methods on the null span - span.finish() - assert True # If we get here without exceptions, the test passes - - def test_should_use_dd_tracer(): """ Test that the should_use_dd_tracer function works as expected diff --git a/tests/test_litellm/litellm_core_utils/test_decode_special_tokens.py b/tests/unit/litellm_core_utils/test_decode_special_tokens.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_decode_special_tokens.py rename to tests/unit/litellm_core_utils/test_decode_special_tokens.py diff --git a/tests/test_litellm/litellm_core_utils/test_dot_notation_indexing.py b/tests/unit/litellm_core_utils/test_dot_notation_indexing.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_dot_notation_indexing.py rename to tests/unit/litellm_core_utils/test_dot_notation_indexing.py diff --git a/tests/test_litellm/litellm_core_utils/test_duration_parser.py b/tests/unit/litellm_core_utils/test_duration_parser.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_duration_parser.py rename to tests/unit/litellm_core_utils/test_duration_parser.py diff --git a/tests/test_litellm/litellm_core_utils/test_error_normalization.py b/tests/unit/litellm_core_utils/test_error_normalization.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_error_normalization.py rename to tests/unit/litellm_core_utils/test_error_normalization.py diff --git a/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py rename to tests/unit/litellm_core_utils/test_exception_mapping_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_extract_base64_image.py b/tests/unit/litellm_core_utils/test_extract_base64_image.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_extract_base64_image.py rename to tests/unit/litellm_core_utils/test_extract_base64_image.py diff --git a/tests/test_litellm/litellm_core_utils/test_fallback_generalizations.py b/tests/unit/litellm_core_utils/test_fallback_generalizations.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_fallback_generalizations.py rename to tests/unit/litellm_core_utils/test_fallback_generalizations.py diff --git a/tests/test_litellm/litellm_core_utils/test_fallback_utils.py b/tests/unit/litellm_core_utils/test_fallback_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_fallback_utils.py rename to tests/unit/litellm_core_utils/test_fallback_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py b/tests/unit/litellm_core_utils/test_get_litellm_params.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_get_litellm_params.py rename to tests/unit/litellm_core_utils/test_get_litellm_params.py diff --git a/tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py b/tests/unit/litellm_core_utils/test_get_llm_provider_endpoint_match.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_get_llm_provider_endpoint_match.py rename to tests/unit/litellm_core_utils/test_get_llm_provider_endpoint_match.py diff --git a/tests/test_litellm/litellm_core_utils/test_get_llm_provider_logic.py b/tests/unit/litellm_core_utils/test_get_llm_provider_logic.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_get_llm_provider_logic.py rename to tests/unit/litellm_core_utils/test_get_llm_provider_logic.py diff --git a/tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py b/tests/unit/litellm_core_utils/test_get_model_cost_map.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_get_model_cost_map.py rename to tests/unit/litellm_core_utils/test_get_model_cost_map.py diff --git a/tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py b/tests/unit/litellm_core_utils/test_get_supported_openai_params.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_get_supported_openai_params.py rename to tests/unit/litellm_core_utils/test_get_supported_openai_params.py diff --git a/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py similarity index 96% rename from tests/test_litellm/litellm_core_utils/test_health_check_helpers.py rename to tests/unit/litellm_core_utils/test_health_check_helpers.py index 1cc96cb1256..c3478c0d5eb 100644 --- a/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -364,6 +364,26 @@ async def test_batch_health_check_uses_alist_batches_for_supported_providers(): mock_alist.assert_called_once() +@pytest.mark.asyncio +async def test_batch_health_check_hands_the_resolved_provider_to_alist_batches(): + filtered_model_params: Final = { + "model": "xai/grok-4.3", + "api_key": "sk-test", + "litellm_metadata": {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]}, + } + + with patch("litellm.alist_batches", new_callable=AsyncMock, return_value={}) as mock_alist: + await HealthCheckHelpers._batch_health_check( + custom_llm_provider="xai", + model_params={**filtered_model_params, "messages": []}, + filtered_model_params=filtered_model_params, + ) + + assert mock_alist.call_args.kwargs["custom_llm_provider"] == "xai" + assert mock_alist.call_args.kwargs["model"] == "xai/grok-4.3" + assert mock_alist.call_args.kwargs["api_key"] == "sk-test" + + @pytest.mark.asyncio async def test_batch_health_check_falls_back_to_acompletion_for_unsupported(): """Providers not in LIST_BATCHES_SUPPORTED_PROVIDERS fall back to acompletion.""" diff --git a/tests/test_litellm/litellm_core_utils/test_image_handling.py b/tests/unit/litellm_core_utils/test_image_handling.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_image_handling.py rename to tests/unit/litellm_core_utils/test_image_handling.py diff --git a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py b/tests/unit/litellm_core_utils/test_initialize_dynamic_callback_params.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py rename to tests/unit/litellm_core_utils/test_initialize_dynamic_callback_params.py diff --git a/tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py b/tests/unit/litellm_core_utils/test_internal_call_metadata.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_internal_call_metadata.py rename to tests/unit/litellm_core_utils/test_internal_call_metadata.py diff --git a/tests/test_litellm/litellm_core_utils/test_json_fragment_accumulator.py b/tests/unit/litellm_core_utils/test_json_fragment_accumulator.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_json_fragment_accumulator.py rename to tests/unit/litellm_core_utils/test_json_fragment_accumulator.py diff --git a/tests/test_litellm/litellm_core_utils/test_json_schema_validation.py b/tests/unit/litellm_core_utils/test_json_schema_validation.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_json_schema_validation.py rename to tests/unit/litellm_core_utils/test_json_schema_validation.py diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py similarity index 98% rename from tests/test_litellm/litellm_core_utils/test_litellm_logging.py rename to tests/unit/litellm_core_utils/test_litellm_logging.py index 3afa31cc801..d717718cba2 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -20,7 +20,7 @@ from openai._legacy_response import HttpxBinaryResponseContent import litellm from litellm._logging import session_id_var, trace_id_var -from litellm.constants import SENTRY_DENYLIST, SENTRY_PII_DENYLIST +from litellm.constants import REDACTED_BY_LITELLM, SENTRY_PII_DENYLIST from litellm.cost_calculator import ocr_batch_cost from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging @@ -357,108 +357,23 @@ def test_post_call_serializes_dict_with_datetime(logging_obj): assert "2026-05-11" in serialized -def test_sentry_sample_rate(monkeypatch): - existing_sample_rate = os.getenv("SENTRY_API_SAMPLE_RATE") - try: - # test with default value by removing the environment variable - if existing_sample_rate: - del os.environ["SENTRY_API_SAMPLE_RATE"] - - set_callbacks(["sentry"]) - # Check if the default sample rate is set to 1.0 - assert os.environ.get("SENTRY_API_SAMPLE_RATE") == "1.0" - - # test with custom value - monkeypatch.setenv("SENTRY_API_SAMPLE_RATE", "0.5") - - set_callbacks(["sentry"]) - # Check if the custom sample rate is set correctly - assert os.environ.get("SENTRY_API_SAMPLE_RATE") == "0.5" - except Exception as e: - print(f"Error: {e}") - finally: - # Restore the original environment variable - if existing_sample_rate: - monkeypatch.setenv("SENTRY_API_SAMPLE_RATE", existing_sample_rate) - else: - if "SENTRY_API_SAMPLE_RATE" in os.environ: - del os.environ["SENTRY_API_SAMPLE_RATE"] - - def test_sentry_environment(monkeypatch): - """Test that SENTRY_ENVIRONMENT is properly handled during Sentry initialization""" - existing_environment = os.getenv("SENTRY_ENVIRONMENT") - existing_dsn = os.getenv("SENTRY_DSN") + import sentry_sdk - # Create mock sentry_sdk module - mock_event_scrubber_instance = MagicMock() - mock_event_scrubber_cls = MagicMock(return_value=mock_event_scrubber_instance) - - mock_scrubber_module = MagicMock() - mock_scrubber_module.EventScrubber = mock_event_scrubber_cls - - mock_sentry_sdk = MagicMock() - mock_sentry_sdk.scrubber = mock_scrubber_module mock_init = MagicMock() - mock_sentry_sdk.init = mock_init + monkeypatch.setattr(sentry_sdk, "init", mock_init) + monkeypatch.setenv("SENTRY_DSN", "https://test@sentry.io/123456") + monkeypatch.delenv("SENTRY_ENVIRONMENT", raising=False) - # Inject mocks into sys.modules - sys.modules["sentry_sdk"] = mock_sentry_sdk - sys.modules["sentry_sdk.scrubber"] = mock_scrubber_module - - try: - # Set a mock DSN to allow Sentry initialization - monkeypatch.setenv("SENTRY_DSN", "https://test@sentry.io/123456") - - # Test with default value (no environment set) - if existing_environment: - del os.environ["SENTRY_ENVIRONMENT"] + set_callbacks(["sentry"]) + assert mock_init.call_args[1]["environment"] == "production" + for environment in ("development", "staging"): + monkeypatch.setenv("SENTRY_ENVIRONMENT", environment) mock_init.reset_mock() set_callbacks(["sentry"]) - # Check that init was called with default environment "production" mock_init.assert_called_once() - call_kwargs = mock_init.call_args[1] - assert call_kwargs["environment"] == "production" - - # Test with custom environment value - monkeypatch.setenv("SENTRY_ENVIRONMENT", "development") - - mock_init.reset_mock() - set_callbacks(["sentry"]) - # Check that init was called with custom environment "development" - mock_init.assert_called_once() - call_kwargs = mock_init.call_args[1] - assert call_kwargs["environment"] == "development" - - # Test with staging environment - monkeypatch.setenv("SENTRY_ENVIRONMENT", "staging") - - mock_init.reset_mock() - set_callbacks(["sentry"]) - # Check that init was called with custom environment "staging" - mock_init.assert_called_once() - call_kwargs = mock_init.call_args[1] - assert call_kwargs["environment"] == "staging" - - except Exception as e: - print(f"Error: {e}") - raise - finally: - # Restore the original environment variables - if existing_environment: - monkeypatch.setenv("SENTRY_ENVIRONMENT", existing_environment) - else: - if "SENTRY_ENVIRONMENT" in os.environ: - del os.environ["SENTRY_ENVIRONMENT"] - - if existing_dsn: - monkeypatch.setenv("SENTRY_DSN", existing_dsn) - else: - if "SENTRY_DSN" in os.environ: - del os.environ["SENTRY_DSN"] - - + assert mock_init.call_args[1]["environment"] == environment def test_use_custom_pricing_for_model(): from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model @@ -3100,37 +3015,34 @@ def test_speech_call_is_still_priced_from_input_characters(call_type): def test_sentry_event_scrubber_initialization(monkeypatch): - # Step 1: Create a fake sentry_sdk.scrubber module - mock_event_scrubber_instance = MagicMock() - mock_event_scrubber_cls = MagicMock(return_value=mock_event_scrubber_instance) + import sentry_sdk - mock_scrubber_module = MagicMock() - mock_scrubber_module.EventScrubber = mock_event_scrubber_cls - - # Step 2: Create a fake sentry_sdk module and insert into sys.modules - mock_sentry_sdk = MagicMock() - mock_sentry_sdk.scrubber = mock_scrubber_module mock_init = MagicMock() - mock_sentry_sdk.init = mock_init + monkeypatch.setattr(sentry_sdk, "init", mock_init) + monkeypatch.delenv("SENTRY_SEND_DEFAULT_PII", raising=False) - # Step 3: Inject both into sys.modules BEFORE import occurs - sys.modules["sentry_sdk"] = mock_sentry_sdk - sys.modules["sentry_sdk.scrubber"] = mock_scrubber_module - - # Step 4: Run the actual sentry setup code set_callbacks(["sentry"]) - # Step 5: Assert the EventScrubber was constructed correctly - mock_event_scrubber_cls.assert_called_once_with( - denylist=SENTRY_DENYLIST, - pii_denylist=SENTRY_PII_DENYLIST, - ) - - # Step 6: Assert the event_scrubber and PII args were passed mock_init.assert_called_once() call_args = mock_init.call_args[1] - assert call_args["event_scrubber"] == mock_event_scrubber_instance assert call_args["send_default_pii"] is False + assert call_args["event_scrubber"].recursive is True + assert {name.lower() for name in SENTRY_PII_DENYLIST} <= {name.lower() for name in call_args["event_scrubber"].denylist} + assert call_args["before_send"] is call_args["before_send_transaction"] + + +def test_sentry_send_default_pii_opt_in(monkeypatch): + import sentry_sdk + + mock_init = MagicMock() + monkeypatch.setattr(sentry_sdk, "init", mock_init) + monkeypatch.setenv("SENTRY_SEND_DEFAULT_PII", "true") + + set_callbacks(["sentry"]) + + call_args = mock_init.call_args[1] + assert call_args["send_default_pii"] is True + assert not {name.lower() for name in SENTRY_PII_DENYLIST} & {name.lower() for name in call_args["event_scrubber"].denylist} def test_get_masked_values(): @@ -5396,6 +5308,19 @@ def test_handle_anthropic_messages_response_logging_passes_model_response_throug assert logging_obj._handle_anthropic_messages_response_logging(result=model_response) is model_response +def test_anthropic_messages_logged_response_tolerates_a_stream_that_assembled_nothing(): + """A /v1/messages stream whose upstream yielded no chunks assembles to None; the spend + row must still land under the message id the caller was served instead of crashing.""" + logging_obj = _anthropic_messages_logging_obj() + logging_obj.record_streamed_anthropic_message_id("msg_served") + + result = logging_obj._anthropic_messages_logged_response(result=None) + + assert isinstance(result, ModelResponse) + assert result.id == "msg_served" + assert result.model == "openai/my-local" + + def test_handle_anthropic_messages_response_logging_degrades_on_unparseable_responses_payload(): """If the Responses translation raises (eg. empty output on an incomplete response), the row must still land: a minimal ModelResponse with model + usage is returned.""" @@ -6690,6 +6615,65 @@ def test_pre_call_redacts_and_masks_raw_request(logging_obj): assert "key=*****" in raw_api_base +_PRIVATE_RAW_REQUEST_ARGS: Final = { + "api_base": "https://api.openai.com/v1/chat/completions", + "headers": {}, + "complete_input_dict": {"messages": [{"role": "user", "content": "PRIVATE-PHRASE"}]}, +} + + +def _pre_call_with_raw_request_logging(logging_obj) -> dict: + metadata: Final = {"user_api_key_alias": "qa-key"} + logging_obj.model_call_details["litellm_params"] = {"metadata": metadata} + logging_obj.log_raw_request_response = True + logging_obj.pre_call(input="hi", api_key="", additional_args=_PRIVATE_RAW_REQUEST_ARGS) + return metadata + + +def _assert_raw_request_redacted_for_callbacks_only(logging_obj, metadata: dict) -> None: + assert metadata["raw_request"] == REDACTED_BY_LITELLM + typed_dict: Final = logging_obj.model_call_details["raw_request_typed_dict"] + assert typed_dict["raw_request_body"] == _PRIVATE_RAW_REQUEST_ARGS["complete_input_dict"] + assert typed_dict["error"] is None + + +def test_pre_call_raw_request_honors_turn_off_message_logging_set_after_import(logging_obj, monkeypatch): + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + + metadata = _pre_call_with_raw_request_logging(logging_obj) + + _assert_raw_request_redacted_for_callbacks_only(logging_obj, metadata) + + +def test_pre_call_raw_request_honors_per_request_turn_off_message_logging(logging_obj, monkeypatch): + monkeypatch.setattr(litellm, "turn_off_message_logging", False) + logging_obj.model_call_details["standard_callback_dynamic_params"] = {"turn_off_message_logging": True} + + metadata = _pre_call_with_raw_request_logging(logging_obj) + + _assert_raw_request_redacted_for_callbacks_only(logging_obj, metadata) + + +def test_debugging_log_honors_json_logs_set_after_import(logging_obj, monkeypatch): + monkeypatch.setattr(litellm, "json_logs", True) + logging_obj.litellm_request_debug = True + + with patch("litellm.litellm_core_utils.litellm_logging.verbose_logger.warning") as warning: + logging_obj._print_llm_call_debugging_log(api_base="https://api.openai.com/v1", headers={}, additional_args={}) + + assert "https://api.openai.com/v1" in warning.call_args.kwargs["extra"]["api_base"] + + +def test_debugging_log_with_json_logs_tolerates_missing_headers(logging_obj, monkeypatch): + monkeypatch.setattr(litellm, "json_logs", True) + logging_obj.litellm_request_debug = True + + with patch("litellm.litellm_core_utils.litellm_logging.verbose_logger.warning") as warning: + logging_obj._print_llm_call_debugging_log(api_base="https://api.openai.com/v1", headers=None, additional_args={}) + + assert "https://api.openai.com/v1" in warning.call_args.kwargs["extra"]["api_base"] + + def _streaming_logging_obj_with_callbacks(callbacks: list[CustomLogger]): import datetime @@ -7856,6 +7840,9 @@ _PUBLISHED_BATCH_RATES: Final = MappingProxyType( "output_cost_per_token_batches": 4.1e-6, "cache_read_input_token_cost_batches": 1.2e-7, "cache_creation_input_token_cost_batches": 1.3e-6, + "input_cost_per_token_above_200k_tokens_batches": 2.1e-6, + "output_cost_per_token_above_200k_tokens_batches": 5.1e-6, + "cache_read_input_token_cost_above_200k_tokens_batches": 2.2e-7, "input_cost_per_token_above_272k_tokens_batches": 3.1e-6, "output_cost_per_token_above_272k_tokens_batches": 7.1e-6, "cache_read_input_token_cost_above_272k_tokens_batches": 3.2e-7, @@ -7864,14 +7851,17 @@ _PUBLISHED_BATCH_RATES: Final = MappingProxyType( ) _PUBLISHED_INPUT_BATCH_KEYS: Final = ( "input_cost_per_token_batches", + "input_cost_per_token_above_200k_tokens_batches", "input_cost_per_token_above_272k_tokens_batches", "cache_read_input_token_cost_batches", + "cache_read_input_token_cost_above_200k_tokens_batches", "cache_read_input_token_cost_above_272k_tokens_batches", "cache_creation_input_token_cost_batches", "cache_creation_input_token_cost_above_272k_tokens_batches", ) _PUBLISHED_OUTPUT_BATCH_KEYS: Final = ( "output_cost_per_token_batches", + "output_cost_per_token_above_200k_tokens_batches", "output_cost_per_token_above_272k_tokens_batches", ) @@ -7960,22 +7950,22 @@ def test_batch_cost_calculator_bills_the_carried_output_tier_when_the_deployment ) +@pytest.mark.parametrize( + "tier_key", + ["input_cost_per_token_above_200k_tokens_batches", "input_cost_per_token_above_272k_tokens_batches"], +) def test_deployment_pricing_model_info_honors_a_tier_only_batch_override_over_the_published_flat_rates( - _published_batch_model: None, + _published_batch_model: None, tier_key: str ) -> None: from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info - info: Final = deployment_pricing_model_info( - _batch_deployment_id({"input_cost_per_token_above_272k_tokens_batches": 1e-3}), _PUBLISHED_BATCH_DEPLOYMENT - ) + info: Final = deployment_pricing_model_info(_batch_deployment_id({tier_key: 1e-3}), _PUBLISHED_BATCH_DEPLOYMENT) carried_keys: Final = tuple( - key - for key in (*_PUBLISHED_INPUT_BATCH_KEYS, *_PUBLISHED_OUTPUT_BATCH_KEYS) - if key != "input_cost_per_token_above_272k_tokens_batches" + key for key in (*_PUBLISHED_INPUT_BATCH_KEYS, *_PUBLISHED_OUTPUT_BATCH_KEYS) if key != tier_key ) assert info is not None - assert info["input_cost_per_token_above_272k_tokens_batches"] == 1e-3 + assert info[tier_key] == 1e-3 assert {key: info[key] for key in carried_keys} == {key: _PUBLISHED_BATCH_RATES[key] for key in carried_keys} diff --git a/tests/test_litellm/litellm_core_utils/test_llm_judge.py b/tests/unit/litellm_core_utils/test_llm_judge.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_llm_judge.py rename to tests/unit/litellm_core_utils/test_llm_judge.py diff --git a/tests/test_litellm/litellm_core_utils/test_llm_request_utils.py b/tests/unit/litellm_core_utils/test_llm_request_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_llm_request_utils.py rename to tests/unit/litellm_core_utils/test_llm_request_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_logging_utils.py b/tests/unit/litellm_core_utils/test_logging_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_logging_utils.py rename to tests/unit/litellm_core_utils/test_logging_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_logging_worker.py b/tests/unit/litellm_core_utils/test_logging_worker.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_logging_worker.py rename to tests/unit/litellm_core_utils/test_logging_worker.py diff --git a/tests/test_litellm/litellm_core_utils/test_max_streaming_duration.py b/tests/unit/litellm_core_utils/test_max_streaming_duration.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_max_streaming_duration.py rename to tests/unit/litellm_core_utils/test_max_streaming_duration.py diff --git a/tests/test_litellm/litellm_core_utils/test_model_param_helper.py b/tests/unit/litellm_core_utils/test_model_param_helper.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_model_param_helper.py rename to tests/unit/litellm_core_utils/test_model_param_helper.py diff --git a/tests/test_litellm/litellm_core_utils/test_model_response_utils.py b/tests/unit/litellm_core_utils/test_model_response_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_model_response_utils.py rename to tests/unit/litellm_core_utils/test_model_response_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_private_json.py b/tests/unit/litellm_core_utils/test_private_json.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_private_json.py rename to tests/unit/litellm_core_utils/test_private_json.py diff --git a/tests/test_litellm/litellm_core_utils/test_provider_affinity.py b/tests/unit/litellm_core_utils/test_provider_affinity.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_provider_affinity.py rename to tests/unit/litellm_core_utils/test_provider_affinity.py diff --git a/tests/test_litellm/litellm_core_utils/test_provider_specific_headers.py b/tests/unit/litellm_core_utils/test_provider_specific_headers.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_provider_specific_headers.py rename to tests/unit/litellm_core_utils/test_provider_specific_headers.py diff --git a/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py b/tests/unit/litellm_core_utils/test_ptu_pricing.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_ptu_pricing.py rename to tests/unit/litellm_core_utils/test_ptu_pricing.py diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_errors.py b/tests/unit/litellm_core_utils/test_realtime_errors.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_realtime_errors.py rename to tests/unit/litellm_core_utils/test_realtime_errors.py diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_realtime_streaming.py rename to tests/unit/litellm_core_utils/test_realtime_streaming.py diff --git a/tests/test_litellm/litellm_core_utils/test_redact_messages.py b/tests/unit/litellm_core_utils/test_redact_messages.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_redact_messages.py rename to tests/unit/litellm_core_utils/test_redact_messages.py diff --git a/tests/test_litellm/litellm_core_utils/test_request_timeout_resolver.py b/tests/unit/litellm_core_utils/test_request_timeout_resolver.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_request_timeout_resolver.py rename to tests/unit/litellm_core_utils/test_request_timeout_resolver.py diff --git a/tests/test_litellm/litellm_core_utils/test_retry_after_headers.py b/tests/unit/litellm_core_utils/test_retry_after_headers.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_retry_after_headers.py rename to tests/unit/litellm_core_utils/test_retry_after_headers.py diff --git a/tests/test_litellm/litellm_core_utils/test_safe_divide_seconds.py b/tests/unit/litellm_core_utils/test_safe_divide_seconds.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_safe_divide_seconds.py rename to tests/unit/litellm_core_utils/test_safe_divide_seconds.py diff --git a/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py b/tests/unit/litellm_core_utils/test_safe_json_dumps.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py rename to tests/unit/litellm_core_utils/test_safe_json_dumps.py diff --git a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py b/tests/unit/litellm_core_utils/test_sensitive_data_masker.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py rename to tests/unit/litellm_core_utils/test_sensitive_data_masker.py diff --git a/tests/unit/litellm_core_utils/test_sentry_scrubbing.py b/tests/unit/litellm_core_utils/test_sentry_scrubbing.py new file mode 100644 index 00000000000..9aae3999129 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_sentry_scrubbing.py @@ -0,0 +1,278 @@ +import hashlib +import json +import secrets +from collections.abc import Callable, Mapping +from functools import reduce +from typing import Final, cast + +import pytest +import sentry_sdk +from pydantic import JsonValue +from sentry_sdk.envelope import Envelope +from sentry_sdk.transport import Transport +from sentry_sdk.utils import event_from_exception + +from litellm.constants import LENGTH_OF_LITELLM_GENERATED_KEY, MINIMUM_CUSTOM_KEY_LENGTH +from litellm.litellm_core_utils.sentry_scrubbing import ( + FILTERED, + MAX_SCRUB_DEPTH, + build_key_pattern, + build_sentry_init_options, + build_string_scrubber, + scrub_json_strings, +) +from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth + +EMAIL: Final = "qa.user@example.com" +VIRTUAL_KEY: Final = "sk-virtual-key-under-test" +KEY_HASH: Final = hashlib.sha256(VIRTUAL_KEY.encode()).hexdigest() +MASTER_KEY: Final = "sk-master-key-under-test" +DATABASE_URL: Final = "postgresql://litellm:db-password-under-test@db.internal:5432/litellm" +PII_ON: Final = {"SENTRY_DSN": "https://key@sentry.example/1", "SENTRY_SEND_DEFAULT_PII": "true"} +PII_OFF: Final = {"SENTRY_DSN": "https://key@sentry.example/1"} + + +class RecordingTransport(Transport): + def __init__(self) -> None: + super().__init__() + self.last_envelope: Envelope | None = None + + def capture_envelope(self, envelope: Envelope) -> None: + self.last_envelope = envelope + + +def reject_request( + valid_token: UserAPIKeyAuth, + user_obj: LiteLLM_UserTable, + general_settings: Mapping[str, str], + data: Mapping[str, Mapping[str, str]], + raw_headers: Mapping[str, str], +) -> None: + raise RuntimeError(f"key {valid_token.token} owned by {user_obj.user_email} was rejected") + + +def raise_with_identity_locals() -> None: + reject_request( + valid_token=UserAPIKeyAuth(token=KEY_HASH, key_name="sk-...test", user_id=EMAIL, user_email=EMAIL), + user_obj=LiteLLM_UserTable(user_id=EMAIL, user_email=EMAIL, user_role="internal_user"), + general_settings={"master_key": MASTER_KEY, "database_url": DATABASE_URL}, + data={"metadata": {"user_api_key_hash": KEY_HASH, "user_api_key_user_email": EMAIL}}, + raw_headers={"authorization": f"Bearer {VIRTUAL_KEY}", "x-api-key": VIRTUAL_KEY, "content-type": "application/json"}, + ) + + +def raise_with_source_context_named_locals() -> None: + metadata: Final = {"context_line": f"Bearer {VIRTUAL_KEY}", "pre_context": [EMAIL], "post_context": [KEY_HASH]} + stacktrace: Final = {"frames": [{"context_line": MASTER_KEY, "pre_context": [EMAIL]}]} + raise RuntimeError(f"rejected with {len(metadata)} metadata fields and {len(stacktrace)} stack fields") + + +def capture_serialized_event(env: Mapping[str, str], raiser: Callable[[], None] = raise_with_identity_locals) -> str: + transport: Final = RecordingTransport() + client: Final = sentry_sdk.Client(transport=transport, **build_sentry_init_options(env)) + try: + raiser() + except RuntimeError as error: + event, hint = event_from_exception(error, client_options=client.options) + client.capture_event(event, hint=hint) + assert transport.last_envelope is not None + return json.dumps(transport.last_envelope.items[0].payload.json) + + +def innermost_frame_vars(serialized: str) -> dict[str, JsonValue]: + event: Final = json.loads(serialized) + frames: Final = event["exception"]["values"][0]["stacktrace"]["frames"] + return frames[-1]["vars"] + + +def test_default_event_carries_no_email_hash_or_secret_anywhere() -> None: + serialized: Final = capture_serialized_event(PII_OFF) + assert EMAIL not in serialized + assert KEY_HASH not in serialized + assert MASTER_KEY not in serialized + assert VIRTUAL_KEY not in serialized + assert "db-password-under-test" not in serialized + frame_vars: Final = innermost_frame_vars(serialized) + assert frame_vars["raw_headers"] == {"authorization": FILTERED, "x-api-key": FILTERED, "content-type": "'application/json'"} + assert f"token='{FILTERED}'" in frame_vars["valid_token"] + assert f"user_id='{FILTERED}'" in frame_vars["valid_token"] + assert f"user_email='{FILTERED}'" in frame_vars["user_obj"] + assert frame_vars["general_settings"] == {"master_key": FILTERED, "database_url": FILTERED} + assert frame_vars["data"] == {"metadata": {"user_api_key_hash": FILTERED, "user_api_key_user_email": FILTERED}} + assert "key_name='sk-...test'" in frame_vars["valid_token"] + assert "user_role='internal_user'" in frame_vars["user_obj"] + + +def test_source_context_lines_are_left_readable() -> None: + frames: Final = json.loads(capture_serialized_event(PII_OFF))["exception"]["values"][0]["stacktrace"]["frames"] + source_lines: Final = tuple( + line + for frame in frames + for line in (*frame.get("pre_context", []), frame.get("context_line", ""), *frame.get("post_context", [])) + ) + assert any("token=KEY_HASH" in line for line in source_lines) + assert not any(FILTERED in line for line in source_lines) + + +def test_source_context_names_outside_stack_frames_are_scrubbed() -> None: + serialized: Final = capture_serialized_event(PII_OFF, raise_with_source_context_named_locals) + assert VIRTUAL_KEY not in serialized + assert MASTER_KEY not in serialized + assert EMAIL not in serialized + assert KEY_HASH not in serialized + frame_vars: Final = innermost_frame_vars(serialized) + assert frame_vars["metadata"] == { + "context_line": f"'Bearer {FILTERED}'", + "pre_context": [f"'{FILTERED}'"], + "post_context": [f"'{FILTERED}'"], + } + assert frame_vars["stacktrace"] == {"frames": [{"context_line": f"'{FILTERED}'", "pre_context": [f"'{FILTERED}'"]}]} + innermost_frame: Final = json.loads(serialized)["exception"]["values"][0]["stacktrace"]["frames"][-1] + assert "raise RuntimeError" in innermost_frame["context_line"] + assert FILTERED not in json.dumps(innermost_frame["pre_context"]) + + +def test_default_event_keeps_the_exception_message_shape() -> None: + serialized: Final = capture_serialized_event(PII_OFF) + message: Final = json.loads(serialized)["exception"]["values"][0]["value"] + assert message == f"key {FILTERED} owned by {FILTERED} was rejected" + + +def test_pii_opt_in_keeps_identifiers_and_still_scrubs_secrets() -> None: + serialized: Final = capture_serialized_event(PII_ON) + frame_vars: Final = innermost_frame_vars(serialized) + assert f"user_id='{EMAIL}'" in frame_vars["valid_token"] + assert f"user_email='{EMAIL}'" in frame_vars["user_obj"] + assert frame_vars["data"] == { + "metadata": {"user_api_key_hash": f"'{KEY_HASH}'", "user_api_key_user_email": f"'{EMAIL}'"} + } + assert f"token='{FILTERED}'" in frame_vars["valid_token"] + assert frame_vars["general_settings"] == {"master_key": FILTERED, "database_url": FILTERED} + assert frame_vars["raw_headers"] == {"authorization": FILTERED, "x-api-key": FILTERED, "content-type": "'application/json'"} + assert MASTER_KEY not in serialized + assert VIRTUAL_KEY not in serialized + assert "db-password-under-test" not in serialized + + +def test_transaction_events_are_scrubbed_too() -> None: + transport: Final = RecordingTransport() + client: Final = sentry_sdk.Client(transport=transport, **build_sentry_init_options(PII_OFF)) + client.capture_event( + { + "type": "transaction", + "transaction": "/user/info", + "contexts": {"trace": {"trace_id": "a" * 32, "span_id": "b" * 16}}, + "spans": [{"description": f"lookup {EMAIL} by {KEY_HASH}", "span_id": "c" * 16, "trace_id": "a" * 32}], + } + ) + assert transport.last_envelope is not None + serialized: Final = json.dumps(transport.last_envelope.items[0].payload.json) + assert EMAIL not in serialized + assert KEY_HASH not in serialized + assert f"lookup {FILTERED} by {FILTERED}" in serialized + + +@pytest.mark.parametrize( + ("text", "expected"), + [ + ( + "UserAPIKeyAuth(token='abc', key_alias='team-a', user_id=None)", + f"UserAPIKeyAuth(token='{FILTERED}', key_alias='team-a', user_id=None)", + ), + ('{"api_key": "sk-1", "model": "gpt-5"}', f'{{"api_key": "{FILTERED}", "model": "gpt-5"}}'), + ("{'user_id': 'u-1', 'max_budget': 5}", f"{{'user_id': '{FILTERED}', 'max_budget': 5}}"), + ("Config(OPENAI_API_KEY=sk-live, timeout=10)", f"Config(OPENAI_API_KEY='{FILTERED}', timeout=10)"), + ("lookup for somebody@example.com failed", f"lookup for {FILTERED} failed"), + (f"hashed key {KEY_HASH} not found", f"hashed key {FILTERED} not found"), + ("request id 0123456789abcdef0123456789abcdef stays", "request id 0123456789abcdef0123456789abcdef stays"), + ("monkey=banana", "monkey=banana"), + ( + "{'x-api-key': 'k-1', 'cookie': 'session=abc', 'content-type': 'application/json'}", + f"{{'x-api-key': '{FILTERED}', 'cookie': '{FILTERED}', 'content-type': 'application/json'}}", + ), + ( + "headers={'x-tenant-key': 'sk-custom-header-key-0123456789'} key_name='sk-...6789'", + f"headers={{'x-tenant-key': '{FILTERED}'}} key_name='sk-...6789'", + ), + ( + "master_key={'value': 'not-a-litellm-key'} timeout=10", + f"master_key='{FILTERED}' timeout=10", + ), + ( + "credentials=[{'value': ('deep', 'secret')}], model='gpt-5'", + f"credentials='{FILTERED}', model='gpt-5'", + ), + ], +) +def test_string_scrubber_rewrites_field_and_value_forms(text: str, expected: str) -> None: + assert build_string_scrubber(send_default_pii=False)(text) == expected + + +def test_bare_key_floor_follows_the_custom_key_minimum() -> None: + scrub: Final = build_string_scrubber(send_default_pii=False) + shortest_key: Final = "sk-" + "a" * (MINIMUM_CUSTOM_KEY_LENGTH - len("sk-")) + assert scrub(f"label={shortest_key} model=gpt-5") == f"label={FILTERED} model=gpt-5" + assert scrub(f"label={shortest_key[:-1]} model=gpt-5") == f"label={shortest_key[:-1]} model=gpt-5" + + +def test_key_pattern_floor_never_exceeds_a_generated_key() -> None: + generated_key: Final = "sk-" + secrets.token_urlsafe(LENGTH_OF_LITELLM_GENERATED_KEY) + stricter_custom_minimum: Final = len(generated_key) + 10 + assert build_key_pattern(stricter_custom_minimum, LENGTH_OF_LITELLM_GENERATED_KEY).fullmatch(generated_key) + assert build_key_pattern(stricter_custom_minimum, LENGTH_OF_LITELLM_GENERATED_KEY).fullmatch(generated_key[:-1]) is None + + +def test_json_walk_fails_closed_past_the_depth_cap() -> None: + scrub: Final = build_string_scrubber(send_default_pii=False) + nested: Final = reduce(lambda inner, _: [inner], range(MAX_SCRUB_DEPTH + 1), cast("JsonValue", "api_key=sk-1")) + assert FILTERED in json.dumps(scrub_json_strings(nested, scrub)) + assert "sk-1" not in json.dumps(scrub_json_strings(nested, scrub)) + assert scrub_json_strings([["api_key=sk-1"]], scrub) == [[f"api_key='{FILTERED}'"]] + + +def test_string_scrubber_with_pii_on_only_scrubs_secrets() -> None: + scrub: Final = build_string_scrubber(send_default_pii=True) + assert scrub(f"user_id='{EMAIL}', token='{KEY_HASH}', email {EMAIL} hash {KEY_HASH}") == ( + f"user_id='{EMAIL}', token='{FILTERED}', email {EMAIL} hash {KEY_HASH}" + ) + assert scrub(f"headers={{'authorization': 'Bearer {VIRTUAL_KEY}'}} sent {VIRTUAL_KEY}") == ( + f"headers={{'authorization': '{FILTERED}'}} sent {FILTERED}" + ) + + +@pytest.mark.parametrize( + ("env", "expected"), + [ + ({}, False), + ({"SENTRY_SEND_DEFAULT_PII": "true"}, True), + ({"SENTRY_SEND_DEFAULT_PII": "True"}, True), + ({"SENTRY_SEND_DEFAULT_PII": "false"}, False), + ({"SENTRY_SEND_DEFAULT_PII": "yes please"}, False), + ], +) +def test_send_default_pii_comes_from_the_environment(env: Mapping[str, str], expected: bool) -> None: + assert build_sentry_init_options(env)["send_default_pii"] is expected + + +def test_init_options_read_dsn_rates_and_environment() -> None: + options: Final = build_sentry_init_options( + { + "SENTRY_DSN": "https://key@sentry.example/7", + "SENTRY_API_TRACE_RATE": "0.25", + "SENTRY_API_SAMPLE_RATE": "0.5", + "SENTRY_ENVIRONMENT": "staging", + } + ) + assert options["dsn"] == "https://key@sentry.example/7" + assert options["traces_sample_rate"] == 0.25 + assert options["sample_rate"] == 0.5 + assert options["environment"] == "staging" + assert options["event_scrubber"].recursive is True + + +def test_init_options_defaults() -> None: + options: Final = build_sentry_init_options({}) + assert options["dsn"] is None + assert options["traces_sample_rate"] == 1.0 + assert options["sample_rate"] == 1.0 + assert options["environment"] == "production" diff --git a/tests/test_litellm/litellm_core_utils/test_served_output_texts.py b/tests/unit/litellm_core_utils/test_served_output_texts.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_served_output_texts.py rename to tests/unit/litellm_core_utils/test_served_output_texts.py diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_cursor.py b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_cursor.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_cursor.py rename to tests/unit/litellm_core_utils/test_streaming_chunk_builder_cursor.py diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_server_tool_use.py b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_server_tool_use.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_server_tool_use.py rename to tests/unit/litellm_core_utils/test_streaming_chunk_builder_server_tool_use.py diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py rename to tests/unit/litellm_core_utils/test_streaming_chunk_builder_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py similarity index 99% rename from tests/test_litellm/litellm_core_utils/test_streaming_handler.py rename to tests/unit/litellm_core_utils/test_streaming_handler.py index 3af79c709cc..6557811b530 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -4900,7 +4900,7 @@ class TestStableStreamingResponseId: @pytest.mark.asyncio async def test_async_stream_without_usage_counts_tokens_off_the_event_loop(): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_overhead.py b/tests/unit/litellm_core_utils/test_streaming_overhead.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_streaming_overhead.py rename to tests/unit/litellm_core_utils/test_streaming_overhead.py diff --git a/tests/test_litellm/litellm_core_utils/test_thread_pool_executor.py b/tests/unit/litellm_core_utils/test_thread_pool_executor.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_thread_pool_executor.py rename to tests/unit/litellm_core_utils/test_thread_pool_executor.py diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py new file mode 100644 index 00000000000..b9818d0d347 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -0,0 +1,1565 @@ +#### What this tests #### +# This tests litellm.token_counter.token_counter() function +import asyncio +import base64 +import importlib +import struct +import threading +import time +import traceback +from concurrent.futures import Future, wait +from typing import Final +from unittest.mock import MagicMock + +import anyio.to_thread +import pytest +import tiktoken + +from unittest.mock import AsyncMock, patch + +import litellm +from litellm import decode, encode, get_modified_max_tokens +from litellm import token_counter as token_counter_old +import litellm.constants +from litellm.constants import TOKEN_COUNTER_MAX_CONCURRENT_COUNTS +from litellm.litellm_core_utils.asyncify import asyncify +from litellm.litellm_core_utils.token_counter import ( + _get_exact_count_function, + _get_extrapolating_count_function, + _get_tiktoken_count_function, + calculate_img_tokens, + get_image_dimensions, + high_detail_image_token_upper_bound, + image_dimensions_from_bytes, + offload_token_count, +) +from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new +from tests.large_text import text +from tests.unit.litellm_core_utils.event_loop_lag import ( + assert_loop_stayed_free, + timed_with_loop_lags, + warm_tokenizer, +) +from tests.unit.litellm_core_utils.messages_with_counts import ( + MESSAGES_TEXT, + MESSAGES_WITH_IMAGES, + MESSAGES_WITH_TOOLS, +) + + +def token_counter_both_assert_same(**args): + new = token_counter_new(**args) + old = token_counter_old(**args) + assert new == old, f"New token counter {new} does not match old token counter {old}" + return new + + +## Choose which token_counter the test will use. + +# token_counter = token_counter_new +# token_counter = token_counter_old +token_counter = token_counter_both_assert_same + + +def test_token_counter_basic(): + assert ( + token_counter( + model="claude-2", + messages=[ + { + "role": "user", + "content": "This is a long message that definitely exceeds the token limit.", + } + ], + ) + == 19 + ) + + +def test_token_counter_large_repeated_text_is_fast(): + messages = [{"role": "user", "content": [{"type": "text", "text": "A" * 1024 * 1024}]}] + + start_time = time.perf_counter() + tokens = token_counter_new(model="us.anthropic.claude-sonnet-4-6", messages=messages) + elapsed = time.perf_counter() - start_time + + assert elapsed < 2, f"Token counting took too long: {elapsed:.2f}s" + assert tokens > 0 + + +@pytest.mark.parametrize( + "text", + [ + "Short text", + "This is a normal message with punctuation, numbers, and a few words.", + ], +) +def test_token_counter_short_text_matches_tiktoken(text): + encoding = tiktoken.get_encoding("cl100k_base") + expected = len(encoding.encode(text, disallowed_special=())) + + assert token_counter_new(model="us.anthropic.claude-sonnet-4-6", text=text) == expected + + +def test_token_counter_default_encoding_matches_cl100k(): + encoding: Final = tiktoken.get_encoding("cl100k_base") + expected: Final = len(encoding.encode("hello world", disallowed_special=())) + + assert token_counter_new(model=None, text="hello world") == expected + + +def test_token_counter_text_over_chunk_boundary_stays_close_to_tiktoken(): + text = ("The quick brown fox jumps over the lazy dog. " * 30)[:1025] + encoding = tiktoken.get_encoding("cl100k_base") + expected = len(encoding.encode(text, disallowed_special=())) + + actual = token_counter_new(model="us.anthropic.claude-sonnet-4-6", text=text) + + assert abs(actual - expected) <= 4 + + +@pytest.mark.parametrize( + "configured", + ["0", "-1", "-1024", "not-an-int", "", " ", "999999999", "inf", "1e9"], +) +def test_invalid_chunk_size_config_stays_usable(monkeypatch, configured): + """A misconfigured chunk size must not raise, count zero, or restore the quadratic encode cost.""" + monkeypatch.setenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS", configured) + try: + reloaded = importlib.reload(litellm.constants) + chunk_size = reloaded.TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS + assert 1 <= chunk_size <= reloaded.TIKTOKEN_ENCODE_MAX_CHUNK_SIZE_CHARS + + encoding = tiktoken.get_encoding("cl100k_base") + count_tokens = _get_tiktoken_count_function( + lambda text: len(encoding.encode(text, disallowed_special=())), + chunk_size=chunk_size, + ) + assert count_tokens("The quick brown fox jumps over the lazy dog. " * 40) > 0 + finally: + monkeypatch.delenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS") + importlib.reload(litellm.constants) + + +def test_valid_chunk_size_config_is_honoured(monkeypatch): + monkeypatch.setenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS", "2048") + try: + assert importlib.reload(litellm.constants).TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS == 2048 + finally: + monkeypatch.delenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS") + importlib.reload(litellm.constants) + + +async def test_huggingface_count_in_a_worker_thread_leaves_the_event_loop_free(): + warm_tokenizer("claude-fable-5") + + tokens, took, lags = await timed_with_loop_lags( + lambda: asyncify(token_counter_new)(model="claude-fable-5", text=text * 100) + ) + + assert tokens > 0 + assert_loop_stayed_free(took, lags) + + +@pytest.mark.parametrize("max_exact_chars", [64, 1_000, 2_500]) +def test_count_above_the_cap_samples_the_whole_string_and_scales(max_exact_chars: int): + count_exactly: Final = MagicMock(side_effect=lambda chunk: chunk.count("a") + len(chunk)) + front_heavy: Final = "a" * 1_000 + "b" * 4_000 + exact: Final = 1_000 + len(front_heavy) + + estimate: Final = _get_extrapolating_count_function(count_exactly, max_exact_chars=max_exact_chars)(front_heavy) + + assert abs(estimate - exact) <= exact // 100 + assert sum(len(call.args[0]) for call in count_exactly.call_args_list) <= max_exact_chars + + +def test_count_at_or_below_the_cap_is_exact(): + count_exactly: Final = MagicMock(side_effect=len) + + assert _get_extrapolating_count_function(count_exactly, max_exact_chars=5_000)("a" * 5_000) == 5_000 + assert count_exactly.call_args_list == [(("a" * 5_000,),)] + + +class _SlowEncoder: + def __init__(self) -> None: + self._lock: Final = threading.Lock() + self.in_flight = 0 + self.peak_in_flight = 0 + + def encode_batch_fast(self, texts: list[str]) -> list[list[int]]: + with self._lock: + self.in_flight += 1 + self.peak_in_flight = max(self.peak_in_flight, self.in_flight) + time.sleep(0.1) + with self._lock: + self.in_flight -= 1 + return [[0] * len(text) for text in texts] + + +@pytest.mark.asyncio +async def test_offloaded_counts_do_not_borrow_from_the_shared_thread_pool(): + encoder: Final = _SlowEncoder() + count: Final = _get_exact_count_function(None, {"type": "huggingface_tokenizer", "tokenizer": encoder}) + shared_pool: Final = anyio.to_thread.current_default_thread_limiter() + burst: Final = 2 * TOKEN_COUNTER_MAX_CONCURRENT_COUNTS + + async def shared_pool_borrowed_until_done(counting: asyncio.Future[list[int]]) -> tuple[int, ...]: + if counting.done(): + return () + await asyncio.sleep(0.01) + return (shared_pool.borrowed_tokens, *await shared_pool_borrowed_until_done(counting)) + + counting: Final = asyncio.ensure_future(asyncio.gather(*(offload_token_count(count)("abc") for _ in range(burst)))) + borrowed: Final = await shared_pool_borrowed_until_done(counting) + + assert await counting == [3] * burst + assert len(borrowed) > 1 and max(borrowed) == 0 + assert 1 < encoder.peak_in_flight <= TOKEN_COUNTER_MAX_CONCURRENT_COUNTS + + +def _count_in_a_fresh_event_loop(text: str, result: Future[int]) -> None: + def slow_count(counted: str) -> int: + time.sleep(0.1) + return len(counted) + + result.set_result(asyncio.run(offload_token_count(slow_count)(text))) + + +def test_offloaded_counts_finish_in_every_event_loop_that_shares_the_process(): + loops: Final = 2 * TOKEN_COUNTER_MAX_CONCURRENT_COUNTS + results: Final = tuple(Future[int]() for _ in range(loops)) + threads: Final = tuple( + threading.Thread(target=_count_in_a_fresh_event_loop, args=("a" * size, result), daemon=True) + for size, result in enumerate(results, start=1) + ) + for thread in threads: + thread.start() + + _, pending = wait(results, timeout=5) + + assert not pending + assert tuple(result.result() for result in results) == tuple(range(1, loops + 1)) + + +@pytest.mark.parametrize( + ("configured", "expected"), + [("8", 8), ("0", 4), ("not-an-int", 4)], +) +def test_max_concurrent_counts_config_is_honoured(monkeypatch: pytest.MonkeyPatch, configured: str, expected: int): + monkeypatch.setenv("TOKEN_COUNTER_MAX_CONCURRENT_COUNTS", configured) + try: + assert importlib.reload(litellm.constants).TOKEN_COUNTER_MAX_CONCURRENT_COUNTS == expected + finally: + monkeypatch.delenv("TOKEN_COUNTER_MAX_CONCURRENT_COUNTS") + importlib.reload(litellm.constants) + + +def test_token_counter_applies_the_default_cap(): + max_exact_chars: Final = litellm.constants.TOKEN_COUNTER_MAX_EXACT_CHARS + prose: Final = ("The quick brown fox jumps over the lazy dog. " * (max_exact_chars // 45 + 1))[:max_exact_chars] + over_the_cap: Final = prose + "a" * 200_000 + exact: Final = _get_exact_count_function("gpt-5.6")(over_the_cap) + + estimate: Final = token_counter_new(model="gpt-5.6", text=over_the_cap) + + assert estimate != exact + assert abs(estimate - exact) <= exact // 100 + + +@pytest.mark.parametrize( + ("configured", "expected"), + [("2048", 2048), ("0", 4_000_000), ("not-an-int", 4_000_000)], +) +def test_max_exact_chars_config_is_honoured(monkeypatch: pytest.MonkeyPatch, configured: str, expected: int): + monkeypatch.setenv("TOKEN_COUNTER_MAX_EXACT_CHARS", configured) + try: + assert importlib.reload(litellm.constants).TOKEN_COUNTER_MAX_EXACT_CHARS == expected + finally: + monkeypatch.delenv("TOKEN_COUNTER_MAX_EXACT_CHARS") + importlib.reload(litellm.constants) + + +def test_token_counter_with_prefix(): + messages = [ + {"role": "user", "content": "Who won the world cup in 2022?"}, + {"role": "assistant", "content": "Argentina", "prefix": True}, + ] + tokens = token_counter(model="gpt-3.5-turbo", messages=messages) + assert tokens == 22, f"Expected 22 tokens, got {tokens}" + + +def test_token_counter_normal_plus_function_calling(): + messages = [ + {"role": "system", "content": "System prompt"}, + {"role": "user", "content": "content1"}, + {"role": "assistant", "content": "content2"}, + {"role": "user", "content": "conten3"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_E0lOb1h6qtmflUyok4L06TgY", + "function": { + "arguments": '{"query":"search query","domain":"google.ca","gl":"ca","hl":"en"}', + "name": "SearchInternet", + }, + "type": "function", + } + ], + }, + { + "tool_call_id": "call_E0lOb1h6qtmflUyok4L06TgY", + "role": "tool", + "name": "SearchInternet", + "content": "tool content", + }, + ] + tokens = token_counter(model="gpt-3.5-turbo", messages=messages) + assert tokens == 80 + + +# test_token_counter_normal_plus_function_calling() + + +def test_token_counter_legacy_function_call_counts_arguments(): + """ + Regression for VERIA-492 (Token-counter function_call bypass). + + The legacy OpenAI assistant `function_call` field carries arbitrary text in + `arguments`. Before the fix, `_count_messages` had no branch for + `function_call` and fell through to the unsupported-key `continue`, so an + assistant turn could smuggle unlimited text past `token_counter` and the + proxy `/utils/token_counter` endpoint (and downstream pre-call budget / + `get_modified_max_tokens` math). After the fix it must be counted the + same as the equivalent `tool_calls` payload. + """ + long_arg = "A" * 4000 + fc_messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": None, + "function_call": {"name": "search", "arguments": long_arg}, + }, + ] + tc_messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "search", "arguments": long_arg}, + } + ], + }, + ] + fc_tokens = token_counter(model="gpt-3.5-turbo", messages=fc_messages) + tc_tokens = token_counter(model="gpt-3.5-turbo", messages=tc_messages) + assert fc_tokens == tc_tokens, ( + f"function_call arguments must count like tool_calls arguments; " + f"got function_call={fc_tokens}, tool_calls={tc_tokens}" + ) + assert fc_tokens > 500, f"4000-char arguments payload must contribute real tokens, got {fc_tokens}" + + +@pytest.mark.parametrize( + "message_count_pair", + MESSAGES_TEXT, +) +def test_token_counter_textonly(message_count_pair): + counted_tokens = token_counter( + model="gpt-35-turbo", messages=[message_count_pair["message"]] + ) + assert counted_tokens == message_count_pair["count"] + + +@pytest.mark.parametrize( + "message_count_pair", + MESSAGES_TEXT, +) +def test_token_counter_count_response_tokens(message_count_pair): + counted_tokens = token_counter( + model="gpt-35-turbo", + messages=[message_count_pair["message"]], + count_response_tokens=True, + ) + # 3 tokens are not added because of count_response_tokens=True + expected = message_count_pair["count"] - 3 + assert counted_tokens == expected + + +@pytest.mark.parametrize( + "message_count_pair", + MESSAGES_WITH_IMAGES, +) +def test_token_counter_with_images(message_count_pair): + counted_tokens = token_counter( + model="gpt-4o", messages=[message_count_pair["message"]] + ) + assert counted_tokens == message_count_pair["count"] + + +@pytest.mark.parametrize( + "message_count_pair", + MESSAGES_WITH_TOOLS, +) +def test_token_counter_with_tools(message_count_pair): + counted_tokens = token_counter( + model="gpt-35-turbo", + messages=[message_count_pair["system_message"]], + tools=message_count_pair["tools"], + tool_choice=message_count_pair["tool_choice"], + ) + expected_tokens = message_count_pair["count"] + actual_diff = counted_tokens - expected_tokens + + if "count-tolerate" in message_count_pair: + if message_count_pair["count-tolerate"] == counted_tokens: + pass # expected + else: + tolerated_diff = message_count_pair["count-tolerate"] - expected_tokens + assert ( + actual_diff <= tolerated_diff + ), f"Expected {expected_tokens} tokens, got {counted_tokens}. Counted tokens is only allowed to be off by {tolerated_diff} in the over-counting direction." + if actual_diff != tolerated_diff: + raise NeedsToleranceUpdateError( + f"SOMETHING BROKEN GOT FIXED! THIS is good! Adjust 'count-tolerate' from {message_count_pair['count-tolerate']} to {counted_tokens}" + ) + + else: + assert ( + expected_tokens == counted_tokens + ), f"Expected {expected_tokens} tokens, got {counted_tokens}." + + +class NeedsToleranceUpdateError(Exception): + """Custom exception to mark tests that have improved""" + + pass + + +# test_tokenizers() + + +def test_encoding_and_decoding(): + try: + sample_text = "Hellö World, this is my input string!" + # openai encoding + decoding + openai_tokens = encode(model="gpt-3.5-turbo", text=sample_text) + openai_text = decode(model="gpt-3.5-turbo", tokens=openai_tokens) + + assert openai_text == sample_text + + # claude encoding + decoding + claude_tokens = encode(model="claude-3-5-haiku-20241022", text=sample_text) + + claude_text = decode(model="claude-3-5-haiku-20241022", tokens=claude_tokens) + + assert claude_text == sample_text + + # cohere encoding + decoding + cohere_tokens = encode(model="command-nightly", text=sample_text) + cohere_text = decode(model="command-nightly", tokens=cohere_tokens) + + assert cohere_text == sample_text + + # llama2 encoding + decoding + llama2_tokens = encode(model="meta-llama/Llama-2-7b-chat", text=sample_text) + llama2_text = decode(model="meta-llama/Llama-2-7b-chat", tokens=llama2_tokens) + + assert llama2_text == sample_text + except Exception as e: + pytest.fail(f"An exception occured: {e}\n{traceback.format_exc()}") + + +# test_encoding_and_decoding() + + +# test_gpt_vision_token_counting() + + +@pytest.mark.parametrize( + "model", + [ + "gpt-4-vision-preview", + "gpt-4o", + "claude-3-opus-20240229", + "command-nightly", + "mistral/mistral-tiny", + ], +) +def test_load_test_token_counter(model): + """ + Token count large prompt 100 times. + + Assert time taken is < 1.5s. + """ + import tiktoken + + messages = [{"role": "user", "content": text}] * 10 + + start_time = time.time() + for _ in range(10): + _ = token_counter(model=model, messages=messages) + # enc.encode("".join(m["content"] for m in messages)) + + end_time = time.time() + + total_time = end_time - start_time + print("model={}, total test time={}".format(model, total_time)) + assert total_time < 10, f"Total encoding time > 10s, {total_time}" + + +@pytest.mark.parametrize( + "model, base_model, input_tokens, user_max_tokens, expected_value", + [ + ("random-model", "random-model", 1024, 1024, 1024), + ("gpt-3.5-turbo", "gpt-3.5-turbo", 4000, 5000, 4096), # model max output = 4096 + ], +) +def test_get_modified_max_tokens( + model, base_model, input_tokens, user_max_tokens, expected_value +): + """ + - Test when max_output is not known => expect user_max_tokens + - Test when max_output == max_input, + - input > max_output, no max_tokens => expect None + - input + max_tokens > max_output => expect remainder + - input + max_tokens < max_output => expect max_tokens + - Test when max_tokens > max_output => expect max_output + """ + args = locals() + import litellm + + litellm.token_counter = MagicMock() + + def _mock_token_counter(*args, **kwargs): + return input_tokens + + litellm.token_counter.side_effect = _mock_token_counter + print(f"_mock_token_counter: {_mock_token_counter()}") + messages = [{"role": "user", "content": "Hello world!"}] + + calculated_value = get_modified_max_tokens( + model=model, + base_model=base_model, + messages=messages, + user_max_tokens=user_max_tokens, + buffer_perc=0, + buffer_num=0, + ) + + if expected_value is None: + assert calculated_value is None + else: + assert ( + calculated_value == expected_value + ), "Got={}, Expected={}, Params={}".format( + calculated_value, expected_value, args + ) + + +def test_empty_tools(): + messages = [{"role": "user", "content": "hey, how's it going?", "tool_calls": None}] + + result = token_counter( + messages=messages, + ) + + print(result) + + +@pytest.mark.skip( + reason="Skipping this test temporarily because it relies on a function being called that I am removing." +) +def test_gpt_4o_token_counter(): + with patch.object( + litellm.utils, "openai_token_counter", new=MagicMock() + ) as mock_client: + token_counter( + model="gpt-4o-2024-05-13", messages=[{"role": "user", "content": "Hey!"}] + ) + + mock_client.assert_called() + + +@pytest.mark.parametrize( + "img_url", + [ + "https://example.com/test-image.png", + "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAL0AAAC9CAMAAADRCYwCAAAAh1BMVEX///8AAAD8/Pz5+fkEBAT39/cJCQn09PRNTU3y8vIMDAwzMzPe3t7v7+8QEBCOjo7FxcXR0dHn5+elpaWGhoYYGBivr686OjocHBy0tLQtLS1TU1PY2Ni6urpaWlpERER3d3ecnJxoaGiUlJRiYmIlJSU4ODhBQUFycnKAgIDBwcFnZ2chISE7EjuwAAAI/UlEQVR4nO1caXfiOgz1bhJIyAJhX1JoSzv8/9/3LNlpYd4rhX6o4/N8Z2lKM2cURZau5JsQEhERERERERERERERERERERHx/wBjhDPC3OGN8+Cc5JeMuheaETSdO8vZFyCScHtmz2CsktoeMn7rLM1u3h0PMAEhyYX7v/Q9wQvoGdB0hlbzm45lEq/wd6y6G9aezvBk9AXwp1r3LHJIRsh6s2maxaJpmvqgvkC7WFS3loUnaFJtKRVUCEoV/RpCnHRvAsesVQ1hw+vd7Mpo+424tLs72NplkvQgcdrsvXkW/zJWqH/fA0FT84M/xnQJt4to3+ZLuanbM6X5lfXKHosO9COgREqpCR5i86pf2zPS7j9tTj+9nO7bQz3+xGEyGW9zqgQ1tyQ/VsxEDvce/4dcUPNb5OD9yXvR4Z2QisuP0xiGWPnemgugU5q/troHhGEjIF5sTOyW648aC0TssuaaCEsYEIkGzjWXOp3A0vVsf6kgRyqaDk+T7DIVWrb58b2tT5xpUucKwodOD/5LbrZC1ws6YSaBZJ/8xlh+XZSYXaMJ2ezNqjB3IPXuehPcx2U6b4t1dS/xNdFzguUt8ie7arnPeyCZroxLHzGgGdqVcspwafizPWEXBee+9G1OaufGdvNng/9C+gwgZ3PH3r87G6zXTZ5D5De2G2DeFoANXfbACkT+fxBQ22YFsTTJF9hjFVO6VbqxZXko4WJ8s52P4PnuxO5KRzu0/hlix1ySt8iXjgaQ+4IHPA9nVzNkdduM9LFT/Aacj4FtKrHA7iAw602Vnht6R8Vq1IOS+wNMKLYqayAYfRuufQPGeGb7sZogQQoLZrGPgZ6KoYn70Iw30O92BNEDpvwouCFn6wH2uS+EhRb3WF/HObZk3HuxfRQM3Y/Of/VH0n4MKNHZDiZvO9+m/ABALfkOcuar/7nOo7B95ACGVAFaz4jMiJwJhdaHBkySmzlGTu82gr6FSTik2kJvLnY9nOd/D90qcH268m3I/cgI1xg1maE5CuZYaWLH+UHANCIck0yt7Mx5zBm5vVHXHwChsZ35kKqUpmo5Svq5/fzfAI5g2vDtFPYo1HiEA85QrDeGm9g//LG7K0scO3sdpj2CBDgCa+0OFs0bkvVgnnM/QBDwllOMm+cN7vMSHlB7Uu4haHKaTwgGkv8tlK+hP8fzmFuK/RQTpaLPWvbd58yWIo66HHM0OsPoPhVqmtaEVL7N+wYcTLTbb0DLdgp23Eyy2VYJ2N7bkLFAAibtoLPe5sLt6Oa2bvU+zyeMa8wrixO0gRTn9tO9NCSThTLGqcqtsDvphlfmx/cPBZVvw24jg1LE2lPuEo35Mhi58U0I/Ga8n5w+NS8i34MAQLos5B1u0xL1ZvCVYVRw/Fs2q53KLaXJMWwOZZ/4MPYV19bAHmgGDKB6f01xoeJKFbl63q9J34KdaVNPJWztQyRkzA3KNs1AdAEDowMxh10emXTCx75CkurtbY/ZpdNDGdsn2UcHKHsQ8Ai3WZi48IfkvtjOhsLpuIRSKZTX9FA4o+0d6o/zOWqQzVJMynL9NsxhSJOaourq6nBVQBueMSyubsX2xHrmuABZN2Ns9jr5nwLFlLF/2R6atjW/67Yd11YQ1Z+kA9Zk9dPTM/o6dVo6HHVgC0JR8oUfmI93T9u3gvTG94bAH02Y5xeqRcjuwnKCK6Q2+ajl8KXJ3GSh22P3Zfx6S+n008ROhJn+JRIUVu6o7OXl8w1SeyhuqNDwNI7SjbK08QrqPxS95jy4G7nCXVq6G3HNu0LtK5J0e226CfC005WKK9sVvfxI0eUbcnzutfhWe3rpZHM0nZ/ny/N8tanKYlQ6VEW5Xuym8yV1zZX58vwGhZp/5tFfhybZabdbrQYOs8F+xEhmPsb0/nki6kIyVvzZzUASiOrTfF+Sj9bXC7DoJxeiV8tjQL6loSd0yCx7YyB6rPdLx31U2qCG3F/oXIuDuqd6LFO+4DNIJuxFZqSsU0ea88avovFnWKRYFYRQDfCfcGaBCLn4M4A1ntJ5E57vicwqq2enaZEF5nokCYu9TbKqCC5yCDfL+GhLxT4w4xEJs+anqgou8DOY2q8FMryjb2MehC1dRJ9s4g9NXeTwPkWON4RH+FhIe0AWR/S9ekvQ+t70XHeimGF78LzuU7d7PwrswdIG2VpgF8C53qVQsTDtBJc4CdnkQPbnZY9mbPdDFra3PCXBBQ5QBn2aQqtyhvlyYM4Hb2/mdhsxCUen04GZVvIJZw5PAamMOmjzq8Q+dzAKLXDQ3RUZItWsg4t7W2DP+JDrJDymoMH7E5zQtuEpG03GTIjGCW3LQqOYEsXgFc78x76NeRwY6SNM+IfQoh6myJKRBIcLYxZcwscJ/gI2isTBty2Po9IkYzP0/SS4hGlxRjFAG5z1Jt1LckiB57yWvo35EaolbvA+6fBa24xodL2YjsPpTnj3JgJOqhcgOeLVsYYwoK0wjY+m1D3rGc40CukkaHnkEjarlXrF1B9M6ECQ6Ow0V7R7N4G3LfOHAXtymoyXOb4QhaYHJ/gNBJUkxclpSs7DNcgWWDDmM7Ke5MJpGuioe7w5EOvfTunUKRzOh7G2ylL+6ynHrD54oQO3//cN3yVO+5qMVsPZq0CZIOx4TlcJ8+Vz7V5waL+7WekzUpRFMTnnTlSCq3X5usi8qmIleW/rit1+oQZn1WGSU/sKBYEqMNh1mBOc6PhK8yCfKHdUNQk8o/G19ZPTs5MYfai+DLs5vmee37zEyyH48WW3XA6Xw6+Az8lMhci7N/KleToo7PtTKm+RA887Kqc6E9dyqL/QPTugzMHLbLZtJKqKLFfzVWRNJ63c+95uWT/F7R0U5dDVvuS409AJXhJvD0EwWaWdW8UN11u/7+umaYjT8mJtzZwP/MD4r57fihiHlC5fylHfaqnJdro+Dr7DajvO+vi2EwyD70s8nCH71nzIO1l5Zl+v1DMCb5ebvCMkGHvobXy/hPumGLyX0218/3RyD1GRLOuf9u/OGQyDmto32yMiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIv7GP8YjWPR/czH2AAAAAElFTkSuQmCC", + ], +) +def test_img_url_token_counter(img_url, monkeypatch): + """ + Verify get_image_dimensions returns valid (width, height) for both an + HTTPS URL and a base64 data URI. The HTTPS branch is exercised with a + mocked HTTP fetch so the test is hermetic - it can't break when a + third-party image URL goes away. + """ + import base64 + from litellm.litellm_core_utils.token_counter import get_image_dimensions + + # Minimal valid 1x1 PNG, served by the mocked safe_get for the URL case. + _tiny_png = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII=" + ) + + if img_url.startswith(("http://", "https://")): + + class _FakeResponse: + headers = {"Content-Length": str(len(_tiny_png))} + + def read(self): + return _tiny_png + + monkeypatch.setattr( + "litellm.litellm_core_utils.token_counter.safe_get", + lambda client, url, **kw: _FakeResponse(), + ) + + width, height = get_image_dimensions(data=img_url) + + print(width, height) + + assert width is not None + assert height is not None + + +def test_token_encode_disallowed_special(): + encode(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>") + token_counter(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>") + + +def test_token_counter(): + try: + messages = [{"role": "user", "content": "hi how are you what time is it"}] + tokens = token_counter(model="gpt-3.5-turbo", messages=messages) + print("gpt-35-turbo") + print(tokens) + assert tokens > 0 + + tokens = token_counter(model="claude-2", messages=messages) + print("claude-2") + print(tokens) + assert tokens > 0 + + tokens = token_counter(model="gemini/chat-bison", messages=messages) + print("gemini/chat-bison") + print(tokens) + assert tokens > 0 + + tokens = token_counter(model="ollama/llama2", messages=messages) + print("ollama/llama2") + print(tokens) + assert tokens > 0 + + tokens = token_counter(model="anthropic.claude-instant-v1", messages=messages) + print("anthropic.claude-instant-v1") + print(tokens) + assert tokens > 0 + except Exception as e: + pytest.fail(f"Error occurred: {e}") + + +import unittest + +from litellm.utils import _load_huggingface_tokenizer, _select_tokenizer_helper, claude_json_str, encoding + +# Clear the cache at module load to ensure clean state +_load_huggingface_tokenizer.cache_clear() + + +class TestTokenizerSelection(unittest.TestCase): + def setUp(self): + """Clear the LRU cache before each test method. + + The HuggingFace tokenizers behind _select_tokenizer_helper are cached with + @lru_cache, which can cause cache hits from previous tests when running with + --dist=loadscope (tests from same file run on same worker). + """ + _load_huggingface_tokenizer.cache_clear() + + @patch("litellm.utils.tokenizer_dispatch.from_pretrained") + def test_llama3_tokenizer_api_failure(self, mock_from_pretrained): + # Setup mock to raise an error + mock_from_pretrained.side_effect = Exception("Failed to load tokenizer") + + # Test with llama-3 model + result = _select_tokenizer_helper("llama-3-7b") + + # Verify the attempt to load Llama-3 tokenizer + mock_from_pretrained.assert_called_once_with("Xenova/llama-3-tokenizer") + + # Verify fallback to OpenAI tokenizer + self.assertEqual(result["type"], "openai_tokenizer") + self.assertEqual(result["tokenizer"], encoding) + + @patch("litellm.utils.tokenizer_dispatch.from_pretrained") + def test_cohere_tokenizer_api_failure(self, mock_from_pretrained): + # Setup mock to raise an error + mock_from_pretrained.side_effect = Exception("Failed to load tokenizer") + + # Add Cohere model to the list for testing + litellm.cohere_models = ["command-r-v1"] + + # Test with Cohere model + result = _select_tokenizer_helper("command-r-v1") + + # Verify the attempt to load Cohere tokenizer + mock_from_pretrained.assert_called_once_with( + "Xenova/c4ai-command-r-v01-tokenizer" + ) + + # Verify fallback to OpenAI tokenizer + self.assertEqual(result["type"], "openai_tokenizer") + self.assertEqual(result["tokenizer"], encoding) + + @patch("litellm.utils.tokenizer_dispatch.anthropic") + def test_claude_tokenizer_api_failure(self, mock_anthropic): + # Setup mock to raise an error + mock_anthropic.side_effect = Exception("Failed to load tokenizer") + + # Add Claude model to the list for testing + litellm.anthropic_models = ["claude-2"] + + # Test with Claude model + result = _select_tokenizer_helper("claude-2") + + # Verify the attempt to load Claude tokenizer + mock_anthropic.assert_called_once_with() + + # Verify fallback to OpenAI tokenizer + self.assertEqual(result["type"], "openai_tokenizer") + self.assertEqual(result["tokenizer"], encoding) + + @patch("litellm.utils.tokenizer_dispatch.from_pretrained") + def test_llama2_tokenizer_api_failure(self, mock_from_pretrained): + # Setup mock to raise an error + mock_from_pretrained.side_effect = Exception("Failed to load tokenizer") + + # Test with Llama-2 model + result = _select_tokenizer_helper("llama-2-7b") + + # Verify the attempt to load Llama-2 tokenizer + mock_from_pretrained.assert_called_once_with( + "hf-internal-testing/llama-tokenizer" + ) + + # Verify fallback to OpenAI tokenizer + self.assertEqual(result["type"], "openai_tokenizer") + self.assertEqual(result["tokenizer"], encoding) + + @patch("litellm.utils._return_huggingface_tokenizer") + def test_disable_hf_tokenizer_download(self, mock_return_huggingface_tokenizer): + monkeypatch = pytest.MonkeyPatch() + monkeypatch.setattr(litellm, "disable_hf_tokenizer_download", True) + try: + result = _select_tokenizer_helper("grok-32r22r") + mock_return_huggingface_tokenizer.assert_not_called() + assert result["type"] == "openai_tokenizer" + assert result["tokenizer"] == encoding + finally: + monkeypatch.undo() + + +def test_token_counter_with_anthropic_tool_use(): + """ + Test that _count_anthropic_content() correctly handles tool_use blocks. + + Validates that: + - 'name' field is counted (string) + - 'input' field is counted (dict serialized to string) + - Metadata fields ('type', 'id') are skipped + """ + messages = [ + {"role": "user", "content": "What's the weather in San Francisco?"}, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "I'll check the weather for you."}, + { + "type": "tool_use", + "id": "toolu_01234567890", # Should be skipped + "name": "get_weather", # Should be counted + "input": { # Should be counted (serialized) + "location": "San Francisco, CA", + "unit": "fahrenheit", + }, + }, + ], + }, + ] + + tokens = token_counter(model="gpt-3.5-turbo", messages=messages) + assert tokens > 0, f"Expected positive token count, got {tokens}" + # Should count: user message + "I'll check" text + "get_weather" name + input dict + assert ( + tokens > 15 + ), f"Expected reasonable token count for message with tool_use, got {tokens}" + + +def test_token_counter_with_anthropic_tool_result(): + """ + Test that _count_anthropic_content() correctly handles tool_result blocks. + + Validates that: + - 'content' field (when string) is counted + - Metadata fields ('type', 'tool_use_id') are skipped + - Full conversation with tool_use → tool_result flow works + """ + messages = [ + {"role": "user", "content": "What's the weather in San Francisco?"}, + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_01234567890", + "name": "get_weather", + "input": {"location": "San Francisco, CA"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_01234567890", # Should be skipped + "content": "The weather in San Francisco is 65°F and sunny.", # Should be counted + } + ], + }, + ] + + tokens = token_counter(model="gpt-3.5-turbo", messages=messages) + assert tokens > 0, f"Expected positive token count, got {tokens}" + assert ( + tokens > 25 + ), f"Expected reasonable token count for conversation with tool_result, got {tokens}" + + +def test_token_counter_with_nested_tool_result(): + """ + Test that _count_anthropic_content() recursively handles nested content lists. + + Validates that: + - tool_result with 'content' as a list (not string) is handled + - Nested content blocks are recursively counted via _count_content_list() + - TypedDict inference correctly identifies list fields + """ + messages = [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_01234567890", + "content": [ # Nested list - should recursively count + { + "type": "text", + "text": "The weather in San Francisco is 65°F and sunny.", + }, + {"type": "text", "text": "UV index is moderate."}, + ], + } + ], + } + ] + + tokens = token_counter(model="gpt-3.5-turbo", messages=messages) + assert tokens > 0, f"Expected positive token count, got {tokens}" + # Should count both nested text blocks + assert ( + tokens > 15 + ), f"Expected reasonable token count for nested tool_result, got {tokens}" + + +def test_token_counter_tool_use_and_result_combined(): + """ + Test dynamic field inference with multiple tool_use and tool_result blocks. + + Validates that: + - Multiple tool_use blocks in same message are handled + - Multiple tool_result blocks in same message are handled + - skip_fields correctly filters metadata across all blocks + - Full realistic conversation flow works end-to-end + """ + messages = [ + { + "role": "user", + "content": "What's the weather in San Francisco and New York?", + }, + { + "role": "assistant", + "content": [ + { + "type": "text", + "text": "I'll check the weather in both cities for you.", + }, + { + "type": "tool_use", + "id": "toolu_01A", + "name": "get_weather", + "input": {"location": "San Francisco, CA"}, + }, + { + "type": "tool_use", + "id": "toolu_01B", + "name": "get_weather", + "input": {"location": "New York, NY"}, + }, + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_01A", + "content": "San Francisco: 65°F, sunny", + }, + { + "type": "tool_result", + "tool_use_id": "toolu_01B", + "content": "New York: 45°F, cloudy", + }, + ], + }, + { + "role": "assistant", + "content": "The weather in San Francisco is 65°F and sunny, while New York is cooler at 45°F and cloudy.", + }, + ] + + tokens = token_counter(model="gpt-3.5-turbo", messages=messages) + assert tokens > 0, f"Expected positive token count, got {tokens}" + # Should count all text, tool names, inputs, and results + assert ( + tokens > 60 + ), f"Expected substantial token count for full tool conversation, got {tokens}" + + +def test_token_counter_with_image_url(): + """ + Test that _count_image_tokens() correctly handles image_url content blocks. + + Validates that: + - image_url as dict with 'url' and 'detail' is handled + - image_url as string is handled + - 'detail' field validation works ('low', 'high', 'auto') + - calculate_img_tokens is called with correct parameters + """ + # Test with dict format (detail: low) + messages_dict = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What's in this image?"}, + { + "type": "image_url", + "image_url": { + "url": "https://example.com/image.jpg", + "detail": "low", # Should use low token count (85 base tokens) + }, + }, + ], + } + ] + + tokens_dict = token_counter( + model="gpt-3.5-turbo", + messages=messages_dict, + use_default_image_token_count=True, # Avoid actual HTTP request + ) + assert tokens_dict > 0, f"Expected positive token count, got {tokens_dict}" + assert tokens_dict > 85, f"Expected at least base image tokens, got {tokens_dict}" + + # Test with string format (defaults to auto/low) + messages_str = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": "https://example.com/image.jpg", # String format + } + ], + } + ] + + tokens_str = token_counter( + model="gpt-3.5-turbo", messages=messages_str, use_default_image_token_count=True + ) + assert ( + tokens_str > 0 + ), f"Expected positive token count for string image_url, got {tokens_str}" + + # Test invalid detail value raises error + messages_invalid = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": { + "url": "https://example.com/image.jpg", + "detail": "invalid", # Should raise ValueError + }, + } + ], + } + ] + + with pytest.raises(ValueError, match="Invalid detail value") as exc_info: + token_counter(model="gpt-3.5-turbo", messages=messages_invalid) + e = exc_info.value + assert "Invalid detail value" in str( + e + ), f"Expected detail validation error, got: {e}" + + +def test_token_counter_with_thinking_content(): + """ + Test that _count_content_list() correctly handles Claude's extended thinking content blocks. + + Validates that: + - 'thinking' content type is recognized and counted + - 'thinking' text field is counted + - 'signature' field is skipped (opaque signature blob) + - Full conversation with thinking blocks works + """ + messages = [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Analyze this complex problem: who came first, chicken or egg", + } + ], + }, + { + "role": "assistant", + "content": [ + { + "type": "thinking", + "thinking": "This is actually a fascinating question that touches on philosophy, biology, and semantics. Let me break this down: The egg came first from an evolutionary biology perspective.", + "signature": "EqcLCkYICxgCKkCrqu6lP...", # Should be skipped + }, + { + "type": "text", + "text": "# The Chicken-or-Egg Question: A Multi-Layered Answer\n\n## **The Short Answer: The Egg Came First**", + }, + ], + }, + {"role": "user", "content": [{"type": "text", "text": "Thanks"}]}, + ] + + tokens = token_counter( + model="anthropic/claude-sonnet-4-5-20250929", messages=messages + ) + assert tokens > 0, f"Expected positive token count, got {tokens}" + # Should count: user message + thinking text + response text + "Thanks" + # The thinking text alone is ~30 tokens, plus other content should be > 50 total + assert ( + tokens > 50 + ), f"Expected substantial token count for message with thinking, got {tokens}" + + # Test that thinking block without 'thinking' field doesn't crash (edge case) + messages_no_thinking = [ + { + "role": "assistant", + "content": [ + { + "type": "thinking", + # No 'thinking' field - should count as 0 tokens + "signature": "EqcLCkYICxgCKkCrqu6lP...", + }, + {"type": "text", "text": "Response"}, + ], + } + ] + + tokens_no_thinking = token_counter( + model="anthropic/claude-sonnet-4-5-20250929", messages=messages_no_thinking + ) + assert ( + tokens_no_thinking > 0 + ), f"Expected positive token count even with empty thinking, got {tokens_no_thinking}" + # Should only count "Response" and message overhead + assert ( + tokens_no_thinking < 15 + ), f"Expected minimal token count for empty thinking block, got {tokens_no_thinking}" + + +def test_token_counter_with_redacted_thinking_content(): + """ + A replayed redacted_thinking block (Anthropic redacted reasoning, or the /v1/messages bridge's stand-in + for a reasoning item with no summary) counts zero tokens for its encrypted payload, like a thinking + block with no text. It used to raise, which made is_prompt_caching_valid_prompt return False and the + prompt_caching pre-call check stop pinning the deployment that held the cached prefix. + """ + model = "anthropic/claude-sonnet-4-5-20250929" + reply = {"type": "text", "text": "Draw from the box labeled Mixed, because that label must be wrong."} + redacted_block = {"type": "redacted_thinking", "data": "EqQBCkYIBRgCKkBjZ2xhc3M" * 30} + user_turn = {"role": "user", "content": [{"type": "text", "text": "Which box do you draw from?"}]} + follow_up = {"role": "user", "content": [{"type": "text", "text": "Restate that in one sentence."}]} + + without_block = [user_turn, {"role": "assistant", "content": [reply]}, follow_up] + with_block = [user_turn, {"role": "assistant", "content": [redacted_block, reply]}, follow_up] + + assert token_counter(model=model, messages=with_block) == token_counter(model=model, messages=without_block) + +def test_token_counter_with_tool_reference_block(): + """ + Regression test: a message containing an Anthropic tool-search + `tool_reference` content block must NOT raise. + + Before the fix, token_counter raised + `Invalid content item type: tool_reference`. On the streaming + anthropic_messages proxy path this nulled response_cost and caused the + SpendLogs row to be dropped, silently undercounting cost. token_counter + must instead count the referenced tool name and return a positive count. + """ + messages = [ + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Let me look up the right tool."}, + {"type": "tool_reference", "tool_name": "search_knowledge_base"}, + ], + } + ] + + # Must not raise, and must produce a positive token count. + tokens = token_counter_new( + model="anthropic/claude-sonnet-4-5-20250929", messages=messages + ) + assert tokens > 0, f"Expected positive token count, got {tokens}" + + # A tool_reference with no/empty tool_name must also be handled gracefully. + messages_empty = [ + { + "role": "assistant", + "content": [{"type": "tool_reference", "tool_name": ""}], + } + ] + tokens_empty = token_counter_new( + model="anthropic/claude-sonnet-4-5-20250929", messages=messages_empty + ) + assert tokens_empty >= 0 + + +def test_count_content_list_rejects_unknown_type(): + """ + An unrecognized content block type must raise, and the error message must + enumerate the supported types (including `tool_reference`). This pins the + catch-all contract so a future block type isn't silently dropped. + """ + from litellm.litellm_core_utils.token_counter import _count_content_list + + with pytest.raises(ValueError, match='Error getting number of tokens from content list: Invalid') as exc_info: + _count_content_list( + count_function=len, + content_list=[{"type": "totally_unknown_block"}], + use_default_image_token_count=False, + default_token_count=None, + ) + + message = str(exc_info.value) + assert "Invalid content item type: totally_unknown_block" in message + assert "tool_reference" in message + + +@pytest.mark.parametrize( + "source", + [ + {"type": "base64", "media_type": "image/png", "data": "iVBORw0KGgo="}, + {"type": "url", "url": "https://example.com/image.png"}, + {"type": "file", "file_id": "file-abc123"}, + ], + ids=["base64", "url", "file"], +) +def test_token_counter_with_anthropic_image_block(source: dict[str, str]): + """Anthropic `image` blocks must count for every source variant, not raise `Invalid content item type` (which the router's context-window pre-call check swallows into an unfiltered dispatch).""" + from litellm.constants import DEFAULT_IMAGE_TOKEN_COUNT + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + {"type": "image", "source": source}, + ], + } + ] + + tokens = token_counter( + model="anthropic/claude-sonnet-4-5-20250929", + messages=messages, + use_default_image_token_count=True, + ) + assert tokens > DEFAULT_IMAGE_TOKEN_COUNT, ( + f"Expected the image block to contribute tokens, got {tokens}" + ) + + +def test_anthropic_image_block_matches_equivalent_image_url(): + """An Anthropic `image` block prices identically to the OpenAI `image_url` carrying the same bytes.""" + anthropic_messages = [ + { + "role": "user", + "content": [ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "iVBORw0KGgo=", + }, + } + ], + } + ] + openai_messages = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + } + ], + } + ] + + anthropic_tokens = token_counter( + model="anthropic/claude-sonnet-4-5-20250929", messages=anthropic_messages + ) + openai_tokens = token_counter( + model="anthropic/claude-sonnet-4-5-20250929", messages=openai_messages + ) + assert anthropic_tokens == openai_tokens + + +def test_anthropic_image_block_nested_in_tool_result(): + """An `image` block nested in a `tool_result.content` list is counted through the same recursion.""" + messages = [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_01", + "content": [ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": "iVBORw0KGgo=", + }, + } + ], + } + ], + } + ] + + tokens = token_counter( + model="anthropic/claude-sonnet-4-5-20250929", + messages=messages, + use_default_image_token_count=True, + ) + assert tokens > 0 + + +@pytest.mark.parametrize( + ("source", "expected"), + [ + ({"type": "base64", "media_type": "image/jpeg", "data": "/9j/4AAQ"}, "data:image/jpeg;base64,/9j/4AAQ"), + ({"type": "url", "url": "https://example.com/image.png"}, "https://example.com/image.png"), + ({"type": "file", "file_id": "file-abc123"}, ""), + ], + ids=["base64", "url", "file"], +) +def test_anthropic_image_source_resolves_to_what_the_image_pricer_reads(source: dict[str, str], expected: str): + """base64 sources become a data URI, url sources pass through, file sources resolve to an empty string.""" + from litellm.litellm_core_utils.token_counter import _anthropic_image_source_data + + assert _anthropic_image_source_data(source) == expected + + +def test_anthropic_image_block_with_empty_base64_data(): + """A base64 source with empty `data` prices as an image rather than raising.""" + from litellm.litellm_core_utils.token_counter import _count_content_list + + tokens = _count_content_list( + count_function=len, + content_list=[ + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": ""}} + ], + use_default_image_token_count=False, + default_token_count=None, + ) + assert tokens > 0 + + +def test_anthropic_image_block_without_source_raises(): + """An `image` block with no `source` raises, matching the OpenAI `image_url`-without-`url` behavior.""" + from litellm.litellm_core_utils.token_counter import _count_content_list + + with pytest.raises(ValueError, match="Error getting number of tokens from content list"): + _count_content_list( + count_function=len, + content_list=[{"type": "image"}], + use_default_image_token_count=False, + default_token_count=None, + ) + + # ... and `default_token_count`, the caller's opt-out from raising, still wins. + assert ( + _count_content_list( + count_function=len, + content_list=[{"type": "image"}], + use_default_image_token_count=False, + default_token_count=7, + ) + == 7 + ) + + +def _count_user_content(content: list[dict]) -> int: + from litellm.litellm_core_utils.token_counter import token_counter + + return token_counter( + model="anthropic/claude-fable-5", + messages=[{"role": "user", "content": content}], + use_default_image_token_count=True, + ) + + +@pytest.mark.parametrize( + "source", + [ + {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"}, + {"type": "url", "url": "https://example.com/report.pdf"}, + {"type": "file", "file_id": "file-abc123"}, + ], + ids=["base64", "url", "file"], +) +def test_anthropic_document_block_with_opaque_source_is_priced_like_an_image(source: dict[str, str]): + """A `document` whose bytes can't be tokenized locally is priced like an `image`, not raised on.""" + prompt = {"type": "text", "text": "Summarize this file."} + + assert _count_user_content([prompt, {"type": "document", "source": source}]) == _count_user_content( + [prompt, {"type": "image", "source": source}] + ) + + +def test_anthropic_document_block_text_sources_count_their_text(): + """`text` and `content` document sources count the text they carry, as inline text blocks would.""" + prompt = {"type": "text", "text": "Summarize this file."} + body = {"type": "text", "text": "Revenue grew eleven percent while churn fell to two percent."} + picture = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "iVBORw0KGgo="}} + + text_source = {"type": "document", "source": {"type": "text", "media_type": "text/plain", "data": body["text"]}} + assert _count_user_content([prompt, text_source]) == _count_user_content([prompt, body]) + + string_content = {"type": "document", "source": {"type": "content", "content": body["text"]}} + assert _count_user_content([prompt, string_content]) == _count_user_content([prompt, body]) + + block_content = {"type": "document", "source": {"type": "content", "content": [body, picture]}} + assert _count_user_content([prompt, block_content]) == _count_user_content([prompt, body, picture]) + + +def test_anthropic_document_title_and_context_add_their_tokens(): + prompt = {"type": "text", "text": "Summarize this file."} + source = {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"} + described = {"type": "document", "source": source, "title": "Q3 board packet", "context": "Shared by finance"} + + assert _count_user_content([prompt, described]) == _count_user_content( + [ + prompt, + {"type": "text", "text": "Q3 board packet"}, + {"type": "text", "text": "Shared by finance"}, + {"type": "document", "source": source}, + ] + ) + + +def test_openai_file_block_prices_like_the_equivalent_anthropic_document(): + """An inline `file` is a `document` in the chat-completions dialect, so it must price identically, not raise. + + Before the fix `file` was missing from the content-block match even though `ChatCompletionFileObject` + is in the union this counter accepts, so every local count of a Responses `input_file` raised + `Invalid content item type: file` and surfaced as a 500 on /v1/responses/input_tokens. + """ + prompt = {"type": "text", "text": "Summarize this file."} + inline_file = { + "type": "file", + "file": {"filename": "report.pdf", "file_data": "data:application/pdf;base64,JVBERi0xLjQK"}, + } + document = { + "type": "document", + "title": "report.pdf", + "source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0xLjQK"}, + } + + assert _count_user_content([prompt, inline_file]) == _count_user_content([prompt, document]) + assert _count_user_content([prompt, inline_file]) > _count_user_content([prompt]) + + +def test_openai_file_block_without_inline_bytes_counts_what_it_carries(): + """A `file` block naming an uploaded file has no bytes to price, so it adds only the filename's tokens.""" + prompt = {"type": "text", "text": "Summarize this file."} + + by_id = {"type": "file", "file": {"file_id": "file-abc123"}} + assert _count_user_content([prompt, by_id]) == _count_user_content([prompt]) + + named = {"type": "file", "file": {"file_id": "file-abc123", "filename": "report.pdf"}} + assert _count_user_content([prompt, named]) == _count_user_content( + [prompt, {"type": "text", "text": "report.pdf"}] + ) + + +def _png_data_url(width: int, height: int) -> str: + ihdr = b"\x89PNG\r\n\x1a\n" + (13).to_bytes(4, "big") + b"IHDR" + width.to_bytes(4, "big") + height.to_bytes(4, "big") + return "data:image/png;base64," + base64.b64encode(ihdr + b"\x08\x06\x00\x00\x00").decode() + + +@pytest.mark.parametrize(("width", "height"), [(1, 1), (768, 768), (2000, 768), (768, 2000), (4096, 4096), (8000, 3072)]) +def test_high_detail_image_token_upper_bound_covers_every_image_size(width: int, height: int) -> None: + assert calculate_img_tokens(_png_data_url(width, height), mode="high") <= high_detail_image_token_upper_bound() + + +def test_high_detail_image_token_upper_bound_is_reached_by_the_largest_high_res_image() -> None: + assert calculate_img_tokens(_png_data_url(2000, 768), mode="high") == high_detail_image_token_upper_bound() + assert calculate_img_tokens(_png_data_url(1, 1), mode="high") < high_detail_image_token_upper_bound() + + +def _png_bytes(width: int, height: int) -> bytes: + return ( + b"\x89PNG\r\n\x1a\n" + + (13).to_bytes(4, "big") + + b"IHDR" + + struct.pack(">II", width, height) + + b"\x08\x02\x00\x00\x00" + ) + + +def _gif_bytes(width: int, height: int) -> bytes: + return b"GIF89a" + struct.pack(" bytes: + app: Final = b"".join( + b"\xff\xe0" + struct.pack(">H", 16) + b"JFIF\x00\x01\x01\x00\x00\x01\x00\x01\x00\x00" + for _ in range(app_segments) + ) + sof: Final = ( + b"\xff" + sof_marker + struct.pack(">HBHHB", 17, 8, height, width, 3) + b"\x01\x22\x00\x02\x11\x01\x03\x11\x01" + ) + return b"\xff\xd8" + app + sof + + +def _webp_bytes(chunk: bytes, payload: bytes) -> bytes: + body: Final = chunk + struct.pack(" bytes: + return _webp_bytes(b"VP8 ", b"\x00\x00\x00\x9d\x01\x2a" + struct.pack(" bytes: + return _webp_bytes(b"VP8L", b"\x2f" + struct.pack(" bytes: + return _webp_bytes(b"VP8X", b"\x00" * 4 + (width - 1).to_bytes(3, "little") + (height - 1).to_bytes(3, "little")) + + +@pytest.mark.parametrize( + ("image", "expected"), + [ + pytest.param(_png_bytes(1024, 768), (1024, 768), id="png"), + pytest.param(_gif_bytes(100, 50), (100, 50), id="gif"), + pytest.param(_jpeg_bytes(800, 600, b"\xc0", 1), (800, 600), id="jpeg-baseline"), + pytest.param(_jpeg_bytes(640, 480, b"\xc2", 3), (640, 480), id="jpeg-progressive-after-app-segments"), + pytest.param(_webp_vp8_bytes(640, 480), (640, 480), id="webp-vp8"), + pytest.param(_webp_vp8l_bytes(320, 240), (320, 240), id="webp-vp8l"), + pytest.param(_webp_vp8x_bytes(1920, 1080), (1920, 1080), id="webp-vp8x"), + ], +) +def test_image_dimensions_from_bytes_reads_each_header_format(image: bytes, expected: tuple[int, int]) -> None: + assert image_dimensions_from_bytes(image) == expected + + +@pytest.mark.parametrize( + "image", + [ + pytest.param(b"", id="empty"), + pytest.param(b"BM" + b"\x00" * 30, id="unknown-format"), + pytest.param(_webp_bytes(b"ALPH", b"\x00" * 16), id="webp-without-an-image-chunk"), + pytest.param(b"\x89PNG\r\n\x1a\n\x00\x00", id="png-truncated-before-ihdr"), + pytest.param(b"\xff\xd8\xff\xe0\x00\x10JFIF", id="jpeg-truncated-inside-app0"), + pytest.param(b"\xff\xd8\xff\xe0\x00\x04\x00\x00", id="jpeg-ends-before-sof"), + ], +) +def test_image_dimensions_from_bytes_returns_none_for_unreadable_headers(image: bytes) -> None: + assert image_dimensions_from_bytes(image) is None + + +def _jpeg_sof(width: int, height: int) -> bytes: + return b"\xff\xc0" + struct.pack(">HBHHB", 17, 8, height, width, 3) + b"\x01\x22\x00\x02\x11\x01\x03\x11\x01" + + +@pytest.mark.parametrize( + "image", + [ + pytest.param(b"\xff\xd8" + b"\xff\xe0\x00\x02" * 1025 + _jpeg_sof(800, 600), id="too-many-segments"), + pytest.param(b"\xff\xd8" + b"\xff" * 2000 + _jpeg_sof(800, 600)[1:], id="too-many-fill-bytes"), + pytest.param(b"\xff\xd8\xff\xe0\x00\x00\x02" + _jpeg_sof(800, 600), id="segment-length-below-two"), + ], +) +def test_image_dimensions_from_bytes_gives_up_on_pathological_jpeg_headers(image: bytes) -> None: + assert image_dimensions_from_bytes(image) is None + + +def test_image_dimensions_from_bytes_still_reads_a_jpeg_with_many_real_segments() -> None: + image: Final = b"\xff\xd8" + b"\xff\xe0\x00\x02" * 1000 + b"\xff" * 64 + _jpeg_sof(800, 600)[1:] + + assert image_dimensions_from_bytes(image) == (800, 600) + + +@pytest.mark.parametrize( + "header", + [ + pytest.param(b"\x89PNG\r\n\x1a\n\x00\x00", id="png-truncated"), + pytest.param(b"\xff\xd8\xff\xe0\x00\x10JFIF", id="jpeg-truncated"), + ], +) +def test_get_image_dimensions_still_raises_for_a_truncated_header(header: bytes) -> None: + with pytest.raises((struct.error, TypeError)): + get_image_dimensions(data="data:image/png;base64," + base64.b64encode(header).decode()) + + +@pytest.mark.parametrize( + "image", + [ + pytest.param(b"BM" + b"\x00" * 30, id="unknown-format"), + pytest.param(b"\xff\xd8" + b"\xff\xe0\x00\x02" * 1025 + _jpeg_sof(800, 600), id="pathological-jpeg"), + ], +) +def test_get_image_dimensions_falls_back_to_the_default_size_for_a_header_it_cannot_read(image: bytes) -> None: + assert get_image_dimensions(data="data:image/png;base64," + base64.b64encode(image).decode()) == ( + litellm.constants.DEFAULT_IMAGE_WIDTH, + litellm.constants.DEFAULT_IMAGE_HEIGHT, + ) diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter_tool.py b/tests/unit/litellm_core_utils/test_token_counter_tool.py similarity index 93% rename from tests/test_litellm/litellm_core_utils/test_token_counter_tool.py rename to tests/unit/litellm_core_utils/test_token_counter_tool.py index 9f8c1070a47..f61b7d335c1 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter_tool.py +++ b/tests/unit/litellm_core_utils/test_token_counter_tool.py @@ -5,8 +5,8 @@ import pytest # Use the same token_counter as the main test. -from tests.test_litellm.litellm_core_utils.test_token_counter import token_counter -from tests.test_litellm.litellm_core_utils.test_token_counter_tool_data import * +from tests.unit.litellm_core_utils.test_token_counter import token_counter +from tests.unit.litellm_core_utils.test_token_counter_tool_data import * @pytest.mark.parametrize( diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter_tool_data.py b/tests/unit/litellm_core_utils/test_token_counter_tool_data.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_token_counter_tool_data.py rename to tests/unit/litellm_core_utils/test_token_counter_tool_data.py diff --git a/tests/unit/litellm_core_utils/test_tokenizer.py b/tests/unit/litellm_core_utils/test_tokenizer.py new file mode 100644 index 00000000000..a9005ff6a86 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_tokenizer.py @@ -0,0 +1,411 @@ +import copy +import os +import pickle +import subprocess +import sys +from pathlib import Path +from typing import Final, Literal + +import pytest +import tiktoken +from tokenizers import Tokenizer as ReferenceTokenizer + +import litellm +from litellm.caching._embedding_router import truncate_embedding_input +from litellm.litellm_core_utils.tokenizer import HuggingFaceTokenizer, OpenAIEncoding +from litellm.utils import claude_json_str +from tests.unit.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON + + +OFFLINE_ENCODINGS: Final = ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "o200k_harmony") +UNICODE_TEXTS: Final = ("hello world", "café 漢字 🙂", "", "a\ud800b", "\ud83d\ude42", "🙂\ud83d\ude42\udfff", " " * 64) + + +@pytest.mark.parametrize("name", OFFLINE_ENCODINGS) +@pytest.mark.parametrize("text", UNICODE_TEXTS) +def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None: + assert_openai_encoding_matches_python(name, text) + + +def assert_openai_encoding_matches_python(name: str, text: str) -> None: + reference: Final = tiktoken.get_encoding(name) + encoding: Final = OpenAIEncoding.from_tiktoken(name) + expected: Final = reference.encode(text) + + assert encoding.encode(text) == expected + assert encoding.count(text) == len(expected) + assert encoding.encode_batch([text], num_threads=2) == reference.encode_batch([text], num_threads=2) + assert encoding.encode_ordinary_batch([text]) == reference.encode_ordinary_batch([text]) + assert encoding.decode_batch([expected]) == reference.decode_batch([expected]) + assert encoding.decode_bytes_batch([expected]) == reference.decode_bytes_batch([expected]) + + +@pytest.mark.parametrize("allowed", (frozenset(), frozenset({"<|endoftext|>"}), "all")) +@pytest.mark.parametrize("disallowed", (frozenset(), frozenset({"<|fim_prefix|>"}), "all")) +def test_openai_special_token_options_match_python( + allowed: frozenset[str] | Literal["all"], disallowed: frozenset[str] | Literal["all"] +) -> None: + reference: Final = tiktoken.get_encoding("cl100k_base") + encoding: Final = OpenAIEncoding.from_tiktoken(reference.name) + text: Final = "hello<|endoftext|><|fim_prefix|>world" + allowed_set: Final = reference.special_tokens_set if allowed == "all" else allowed + disallowed_set: Final = reference.special_tokens_set - allowed_set if disallowed == "all" else disallowed + if any(token in text for token in disallowed_set): + with pytest.raises(ValueError, match="disallowed special token"): + encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed) + return + assert encoding.encode(text, allowed_special=allowed, disallowed_special=disallowed) == reference.encode( + text, allowed_special=allowed, disallowed_special=disallowed + ) + assert encoding.special_tokens_set == reference.special_tokens_set + assert encoding.eot_token == reference.eot_token + + +@pytest.mark.parametrize("errors", ("replace", "ignore", "backslashreplace", "strict")) +def test_openai_partial_token_decoding_preserves_error_policy(errors: str) -> None: + reference: Final = tiktoken.get_encoding("cl100k_base") + encoding: Final = OpenAIEncoding.from_tiktoken(reference.name) + tokens: Final = reference.encode("🙂")[:1] + assert encoding.decode_bytes(tokens) == reference.decode_bytes(tokens) + if errors == "strict": + with pytest.raises(UnicodeDecodeError): + encoding.decode(tokens, errors=errors) + return + assert encoding.decode(tokens, errors=errors) == reference.decode(tokens, errors=errors) + assert encoding.decode_tokens_bytes(tokens) == reference.decode_tokens_bytes(tokens) + + +def test_public_encoding_and_semantic_cache_preserve_truncated_unicode() -> None: + reference: Final = tiktoken.get_encoding(litellm.encoding.name) + text: Final = "🙂" + tokens: Final = reference.encode(text) + + assert litellm.encoding.encode(text, disallowed_special=()) == tokens + assert litellm.encoding.encode_batch([text]) == [tokens] + assert litellm.decode(tokens=tokens[:1]) == reference.decode(tokens[:1]) + assert truncate_embedding_input(text, "", 1) == reference.decode(tokens[:1]) + + +@pytest.mark.parametrize("add_special_tokens", (True, False)) +def test_huggingface_encoding_preserves_result_fields_and_serialization(add_special_tokens: bool) -> None: + reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) + tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON) + expected: Final = reference.encode("Hello World", add_special_tokens=add_special_tokens) + actual: Final = tokenizer.encode("Hello World", add_special_tokens=add_special_tokens) + + assert (actual.ids, actual.tokens, actual.type_ids, actual.offsets, actual.word_ids, actual.sequence_ids) == ( + expected.ids, + expected.tokens, + expected.type_ids, + expected.offsets, + expected.word_ids, + expected.sequence_ids, + ) + assert (actual.attention_mask, actual.special_tokens_mask, actual.n_sequences, len(actual)) == ( + expected.attention_mask, + expected.special_tokens_mask, + expected.n_sequences, + len(expected), + ) + assert copy.deepcopy(actual).ids == expected.ids + assert pickle.loads(pickle.dumps(actual)).offsets == expected.offsets + assert tokenizer.decode(actual.ids, skip_special_tokens=False) == reference.decode( + expected.ids, skip_special_tokens=False + ) + + +def test_huggingface_character_offsets_and_pretokenized_pairs_match_python() -> None: + reference: Final = ReferenceTokenizer.from_str(claude_json_str) + tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str) + text: Final = "café 漢字 🙂" + actual: Final = tokenizer.encode(text) + expected: Final = reference.encode(text) + + assert actual.offsets == expected.offsets + assert actual.ids == expected.ids + assert ( + tokenizer.encode(["hello", "world"], ["again"], is_pretokenized=True).ids + == reference.encode(["hello", "world"], ["again"], is_pretokenized=True).ids + ) + + +def test_huggingface_batches_apply_padding_across_inputs() -> None: + reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) + reference.enable_padding(pad_id=0, pad_token="[UNK]") + tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str()) + inputs: Final = ["Hello", ("Hello World", "World")] + expected: Final = reference.encode_batch(inputs) + actual: Final = tokenizer.encode_batch(inputs) + fast: Final = tokenizer.encode_batch_fast(inputs) + + assert [(item.ids, item.attention_mask, item.offsets) for item in actual] == [ + (item.ids, item.attention_mask, item.offsets) for item in expected + ] + assert [item.ids for item in fast] == [item.ids for item in expected] + assert tokenizer.decode_batch([item.ids for item in actual]) == reference.decode_batch( + [item.ids for item in expected] + ) + + +def test_caller_supplied_huggingface_tokenizer_preserves_public_encode_and_count() -> None: + tokenizer: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) + custom: Final = {"type": "huggingface_tokenizer", "tokenizer": tokenizer} + expected: Final = tokenizer.encode("Hello World").ids + + assert litellm.encode(text="Hello World", custom_tokenizer=custom) == expected + assert litellm.token_counter(text="Hello World", custom_tokenizer=custom) == len(expected) + assert litellm.decode(tokens=expected, custom_tokenizer=custom) == "Hello World" + + +def test_caller_supplied_tiktoken_treats_special_spellings_as_text() -> None: + tokenizer: Final = tiktoken.get_encoding("cl100k_base") + custom: Final = {"type": "openai_tokenizer", "tokenizer": tokenizer} + text: Final = "<|endoftext|>" + + assert litellm.encode(text=text, custom_tokenizer=custom) == tokenizer.encode(text, disallowed_special=()) + + +def test_public_tokenizer_objects_survive_pickle_and_deepcopy(tmp_path: Path) -> None: + custom: Final = litellm.create_tokenizer(TOKENIZER_JSON) + tokenizer: Final = custom["tokenizer"] + path: Final = tmp_path / "tokenizer.json" + tokenizer.save(str(path)) + + assert copy.deepcopy(custom)["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids + assert ( + pickle.loads(pickle.dumps(custom))["tokenizer"].encode("Hello World").ids == tokenizer.encode("Hello World").ids + ) + assert HuggingFaceTokenizer.from_file(str(path)).encode("Hello World").ids == tokenizer.encode("Hello World").ids + assert copy.deepcopy(litellm.encoding).encode("hello") == litellm.encoding.encode("hello") + assert pickle.loads(pickle.dumps(litellm.encoding)).encode("hello") == litellm.encoding.encode("hello") + + +@pytest.mark.parametrize("offline", ("0", "1")) +def test_hub_loader_preserves_environment_auth_cache_and_offline(tmp_path: Path, offline: str) -> None: + script: Final = """ +import json +import sys +from pathlib import Path +sys.path.insert(0, sys.argv[1]) +import httpx +import huggingface_hub +from huggingface_hub.errors import LocalEntryNotFoundError +import litellm +payload = sys.argv[2].encode() +offline = sys.argv[3] == "1" +observed = [] +def handle(request): + assert not offline, "offline loading issued a request" + if request.url.path.endswith("/tokenizer.json"): + observed.append(request.headers.get("authorization")) + if request.headers.get("authorization") != "Bearer audit-fixture-token": + return httpx.Response(401) + return httpx.Response(200, headers={"content-length": str(len(payload)), "etag": '"fixture"', "x-repo-commit": "a" * 40}, content=payload if request.method == "GET" else b"") +if not offline: + huggingface_hub.set_client_factory(lambda: httpx.Client(transport=httpx.MockTransport(handle))) +try: + tokenizer = litellm.create_pretrained_tokenizer("test-fixture/tokenizer")["tokenizer"] +except LocalEntryNotFoundError: + assert offline + assert observed == [] +else: + assert not offline + assert "Bearer audit-fixture-token" in observed + assert tokenizer.decode(tokenizer.encode("Hello World").ids) == "Hello World" + assert tuple(Path(sys.argv[4]).rglob("tokenizer.json")) +print("compatible") +""" + result: Final = subprocess.run( + [ + sys.executable, + "-I", + "-c", + script, + str(Path(litellm.__file__).parent.parent), + TOKENIZER_JSON, + offline, + str(tmp_path / "cache"), + ], + capture_output=True, + text=True, + timeout=30, + env={ + **os.environ, + "HF_HOME": str(tmp_path / "home"), + "HF_HUB_CACHE": str(tmp_path / "cache"), + "HF_ENDPOINT": "http://127.0.0.1:9", + "HF_TOKEN": "audit-fixture-token", + "HF_HUB_OFFLINE": offline, + "HF_HUB_DISABLE_IMPLICIT_TOKEN": "0", + "LITELLM_LOCAL_MODEL_COST_MAP": "True", + }, + ) + assert result.returncode == 0, result.stdout + result.stderr + assert result.stdout.strip() == "compatible" + + +@pytest.mark.parametrize("rust", (None, "0", "1")) +def test_tokenization_without_native_extension_stays_offline(tmp_path: Path, rust: str | None) -> None: + script: Final = """ +import importlib.abc +import sys +sys.path.insert(0, sys.argv[1]) +def reject_network(event, args): + if event == "socket.connect": + raise AssertionError("tokenizer attempted a network connection") +sys.addaudithook(reject_network) +class Block(importlib.abc.MetaPathFinder): + def find_spec(self, fullname, path=None, target=None): + if fullname == "litellm.rust_bridge._native": + raise ImportError("native extension is unavailable") +sys.meta_path.insert(0, Block()) +import litellm +from litellm.rust_bridge.tokenizer import get_encoding +import tiktoken +from tokenizers import Tokenizer +assert isinstance(litellm.encoding, tiktoken.Encoding) +for name in ("cl100k_base", "o200k_base", "o200k_harmony", "p50k_base", "p50k_edit"): + encoding = get_encoding(name) + text = "offline café 漢字 🙂" + " " * 64 + assert encoding.decode(encoding.encode(text)) == text +ids = litellm.encode(text="hello world") +assert litellm.decode(tokens=ids) == "hello world" +assert litellm.token_counter(model=None, text="hello world") == len(ids) +custom = litellm.create_tokenizer(sys.argv[2]) +assert isinstance(custom["tokenizer"], Tokenizer) +custom["tokenizer"].enable_padding(pad_id=0, pad_token="[UNK]") +assert litellm.decode(tokens=litellm.encode(text="Hello World", custom_tokenizer=custom), custom_tokenizer=custom) == "Hello World" +print("compatible") +""" + result: Final = subprocess.run( + [sys.executable, "-I", "-c", script, str(Path(litellm.__file__).parent.parent), TOKENIZER_JSON], + capture_output=True, + text=True, + timeout=30, + cwd=tmp_path, + env={ + **{key: value for key, value in os.environ.items() if key != "LITELLM_RUST"}, + **({"LITELLM_RUST": rust} if rust is not None else {}), + "LITELLM_LOCAL_MODEL_COST_MAP": "True", + "TIKTOKEN_CACHE_DIR": str(tmp_path / "unused-tokenizer-cache"), + }, + ) + assert result.returncode == 0, result.stdout + result.stderr + assert result.stdout.strip() == "compatible" + assert not (tmp_path / "unused-tokenizer-cache").exists() + + +@pytest.mark.parametrize("is_pretokenized", (False, True)) +def test_huggingface_batch_sequence_containers_match_python(is_pretokenized: bool) -> None: + reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) + tokenizer: Final = HuggingFaceTokenizer.from_str(TOKENIZER_JSON) + inputs: Final = [["Hello", "World"], ("Hello", "World")] + actual: Final = tokenizer.encode_batch(inputs, is_pretokenized=is_pretokenized) + expected: Final = reference.encode_batch(inputs, is_pretokenized=is_pretokenized) + assert [(item.ids, item.type_ids, item.sequence_ids) for item in actual] == [ + (item.ids, item.type_ids, item.sequence_ids) for item in expected + ] + + +@pytest.mark.parametrize("name", ("cl100k_base", "o200k_base", "p50k_edit")) +def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: + assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name) + + +def assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: + reference: Final = tiktoken.get_encoding(name) + encoding: Final = OpenAIEncoding.from_tiktoken(name) + text: Final = "hello fanta" + + assert repr(encoding) == repr(reference) == f"" + assert (encoding.name, encoding.n_vocab, encoding.max_token_value) == ( + reference.name, + reference.n_vocab, + reference.max_token_value, + ) + assert encoding.token_byte_values() == reference.token_byte_values() + assert encoding.encode_single_token("hello") == reference.encode_single_token("hello") + assert encoding.encode_single_token(b"<|endoftext|>") == reference.eot_token + assert [encoding.is_special_token(token) for token in (0, reference.eot_token)] == [False, True] + assert encoding.decode_with_offsets(reference.encode(text)) == reference.decode_with_offsets(reference.encode(text)) + assert encoding.encode_to_numpy(text).tolist() == reference.encode_to_numpy(text).tolist() + stable, completions = encoding.encode_with_unstable(text) + expected_stable, expected_completions = reference.encode_with_unstable(text) + assert (stable, sorted(completions)) == (expected_stable, sorted(expected_completions)) + with pytest.raises(KeyError): + encoding.encode_single_token("<|not-a-token|>") + + +def test_huggingface_tokenizer_exposes_the_tokenizers_vocabulary_surface() -> None: + reference: Final = ReferenceTokenizer.from_str(TOKENIZER_JSON) + reference.enable_padding(pad_id=0, pad_token="[UNK]", length=4) + reference.enable_truncation(max_length=3, stride=1, strategy="only_first", direction="left") + tokenizer: Final = HuggingFaceTokenizer.from_str(reference.to_str()) + + assert tokenizer.token_to_id("Hello") == reference.token_to_id("Hello") == 1 + assert tokenizer.id_to_token(3) == reference.id_to_token(3) == "[BOS]" + assert tokenizer.id_to_token(99) is None + assert tokenizer.get_vocab() == reference.get_vocab() + assert tokenizer.get_vocab(with_added_tokens=False) == reference.get_vocab(with_added_tokens=False) + assert tokenizer.get_vocab_size() == reference.get_vocab_size() == 4 + assert tokenizer.get_vocab_size(with_added_tokens=False) == reference.get_vocab_size(with_added_tokens=False) + added: Final = tokenizer.get_added_tokens_decoder() + expected_added: Final = reference.get_added_tokens_decoder() + assert {token_id: str(token) for token_id, token in added.items()} == { + token_id: str(token) for token_id, token in expected_added.items() + } + assert added[3].special == expected_added[3].special + assert tokenizer.num_special_tokens_to_add(False) == reference.num_special_tokens_to_add(False) == 1 + assert tokenizer.num_special_tokens_to_add(True) == reference.num_special_tokens_to_add(True) == 0 + assert tokenizer.padding == reference.padding + assert tokenizer.truncation == reference.truncation + assert tokenizer.encode_special_tokens == reference.encode_special_tokens is False + assert HuggingFaceTokenizer.from_buffer(TOKENIZER_JSON.encode()).encode("Hello").ids == [3, 1] + assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).padding is None + assert HuggingFaceTokenizer.from_str(TOKENIZER_JSON).truncation is None + + +def test_huggingface_encoding_exposes_the_tokenizers_lookup_and_mutation_surface() -> None: + reference: Final = ReferenceTokenizer.from_str(claude_json_str) + tokenizer: Final = HuggingFaceTokenizer.from_str(claude_json_str) + text: Final = "hello wide world" + actual: Final = tokenizer.encode(text, "again") + expected: Final = reference.encode(text, "again") + + lookups: Final = ( + lambda encoding: [encoding.token_to_chars(index) for index in range(len(encoding))], + lambda encoding: [encoding.token_to_word(index) for index in range(len(encoding))], + lambda encoding: [encoding.token_to_sequence(index) for index in range(len(encoding))], + lambda encoding: [encoding.char_to_token(position) for position in range(len(text))], + lambda encoding: [encoding.char_to_word(position) for position in range(len(text))], + lambda encoding: [encoding.char_to_token(position, 1) for position in range(5)], + lambda encoding: [encoding.word_to_tokens(word) for word in range(3)], + lambda encoding: [encoding.word_to_chars(word) for word in range(3)], + lambda encoding: [encoding.word_to_tokens(0, 1), encoding.word_to_chars(0, 1)], + ) + for lookup in lookups: + assert lookup(actual) == lookup(expected) + assert repr(actual) == repr(expected) + + actual.truncate(4, stride=1, direction="left") + expected.truncate(4, stride=1, direction="left") + assert (actual.ids, [item.ids for item in actual.overflowing]) == ( + expected.ids, + [item.ids for item in expected.overflowing], + ) + actual.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="") + expected.pad(6, direction="left", pad_id=7, pad_type_id=1, pad_token="") + assert (actual.ids, actual.attention_mask, actual.type_ids, actual.tokens) == ( + expected.ids, + expected.attention_mask, + expected.type_ids, + expected.tokens, + ) + actual.set_sequence_id(3) + expected.set_sequence_id(3) + assert actual.sequence_ids == expected.sequence_ids + merged: Final = type(actual).merge([actual, tokenizer.encode("more")]) + assert merged.ids == type(expected).merge([expected, reference.encode("more")]).ids + assert merged.offsets == type(expected).merge([expected, reference.encode("more")]).offsets + with pytest.raises(ValueError, match="direction"): + actual.pad(8, direction="sideways") diff --git a/tests/test_litellm/litellm_core_utils/test_tool_search_spend_logging.py b/tests/unit/litellm_core_utils/test_tool_search_spend_logging.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_tool_search_spend_logging.py rename to tests/unit/litellm_core_utils/test_tool_search_spend_logging.py diff --git a/tests/test_litellm/litellm_core_utils/test_url_utils.py b/tests/unit/litellm_core_utils/test_url_utils.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_url_utils.py rename to tests/unit/litellm_core_utils/test_url_utils.py diff --git a/tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py b/tests/unit/litellm_core_utils/test_xai_oauth_routing.py similarity index 100% rename from tests/test_litellm/litellm_core_utils/test_xai_oauth_routing.py rename to tests/unit/litellm_core_utils/test_xai_oauth_routing.py diff --git a/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_compact.py b/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_compact.py index d835db63d83..5e2956b532a 100644 --- a/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_compact.py +++ b/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_compact.py @@ -2733,7 +2733,7 @@ def test_build_summary_messages_keeps_midturn_system_correction_in_place(): async def test_threshold_check_counts_tokens_off_the_event_loop(monkeypatch): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py b/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py index a21c22cf5fa..9fad6ca5e66 100644 --- a/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py +++ b/tests/unit/llms/anthropic/experimental_pass_through/context_management/test_dispatcher.py @@ -133,7 +133,7 @@ async def test_malformed_edit_entries_are_skipped(): async def test_sync_editor_counts_tokens_off_the_event_loop(): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/unit/llms/bedrock/chat/test_invoke_handler.py b/tests/unit/llms/bedrock/chat/test_invoke_handler.py index 466e9b4fda8..ed8b7023977 100644 --- a/tests/unit/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/unit/llms/bedrock/chat/test_invoke_handler.py @@ -1,5 +1,6 @@ import base64 import binascii +import itertools import datetime import json import struct @@ -14,10 +15,13 @@ import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.bedrock.chat.invoke_handler import ( + AmazonOpenAICompatibleStreamDecoder, AWSEventStreamDecoder, make_call, make_sync_call, ) +from litellm.exceptions import MidStreamFallbackError +from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.utils import ModelResponseStream @@ -799,3 +803,125 @@ async def test_moonshot_invoke_async_stream_yields_openai_shaped_chunks(_aws_tes ) _assert_moonshot_stream_content([chunk async for chunk in stream]) + + +def _truncated_frame() -> bytes: + return _bedrock_event_stream_frame(_openai_stream_chunk({"role": "assistant"}))[:-8] + + +def _event_stream_headers() -> httpx.Headers: + return httpx.Headers({"content-type": "application/vnd.amazon.eventstream", "x-amzn-RequestId": "req-empty-1"}) + + +_UNDECODABLE_STREAM_BODIES: Final = ( + pytest.param(b"", id="empty"), + pytest.param(b"\x00\x00\x00\x05", id="shorter-than-a-prelude"), + pytest.param(_truncated_frame(), id="truncated-first-message"), +) + + +def _assert_no_events_error(error: BedrockError, body: bytes) -> None: + assert error.status_code == 502 + assert "HTTP 200" in error.message + assert "decoded to no events" in error.message + assert f"{len(body)} bytes received" in error.message + assert "application/vnd.amazon.eventstream" in error.message + assert "req-empty-1" in error.message + assert f"first bytes={body[:200]!r}" in error.message + + +@pytest.mark.parametrize("body", _UNDECODABLE_STREAM_BODIES) +def test_iter_bytes_raises_when_a_200_body_decodes_to_no_events(body: bytes) -> None: + decoder: Final = AWSEventStreamDecoder(model="us.moonshotai.kimi-k3") + + with pytest.raises(BedrockError) as exc_info: + list(decoder.iter_bytes(iter([body]), response_headers=_event_stream_headers())) + + _assert_no_events_error(exc_info.value, body) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", _UNDECODABLE_STREAM_BODIES) +async def test_aiter_bytes_raises_when_a_200_body_decodes_to_no_events(body: bytes) -> None: + async def _chunks() -> AsyncIterator[bytes]: + yield body + + decoder: Final = AWSEventStreamDecoder(model="us.moonshotai.kimi-k3") + + with pytest.raises(BedrockError) as exc_info: + _ = [chunk async for chunk in decoder.aiter_bytes(_chunks(), response_headers=_event_stream_headers())] + + _assert_no_events_error(exc_info.value, body) + + +def test_iter_bytes_raises_when_the_stream_ends_mid_message() -> None: + decoder: Final = AmazonOpenAICompatibleStreamDecoder(model="moonshot.kimi-k2-thinking", sync_stream=True) + stream: Final = decoder.iter_bytes(iter([_MOONSHOT_RAW_STREAM, _truncated_frame()])) + + chunks: Final = list(itertools.islice(stream, 4)) + with pytest.raises(BedrockError) as exc_info: + next(stream) + + _assert_moonshot_stream_content(chunks) + assert exc_info.value.status_code == 502 + assert f"{len(_truncated_frame())} undecoded bytes after 4 events" in exc_info.value.message + assert "first bytes=" not in exc_info.value.message + + +def test_iter_bytes_yields_a_complete_stream_without_raising() -> None: + decoder: Final = AmazonOpenAICompatibleStreamDecoder(model="moonshot.kimi-k2-thinking", sync_stream=True) + + chunks: Final = list(decoder.iter_bytes(iter([_MOONSHOT_RAW_STREAM[:100], _MOONSHOT_RAW_STREAM[100:]]))) + + _assert_moonshot_stream_content(chunks) + + +def _assert_empty_stream_surfaced_as_bad_gateway(error: MidStreamFallbackError) -> None: + assert error.status_code == 502 + assert error.is_pre_first_chunk is True + assert isinstance(error.original_exception, litellm.BadGatewayError) + assert "decoded to no events" in str(error) + assert "req-empty-1" in str(error) + + +def test_converse_stream_with_an_empty_200_body_raises_instead_of_an_empty_turn(_aws_test_credentials: None) -> None: + response: Final = MagicMock(status_code=200, headers=_event_stream_headers()) + response.iter_bytes = lambda chunk_size=None: iter([b""]) + client: Final = HTTPHandler() + client.post = MagicMock(return_value=response) + + with pytest.raises(MidStreamFallbackError) as exc_info: + list( + litellm.completion( + model="bedrock/us.moonshotai.kimi-k3", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + ) + ) + + _assert_empty_stream_surfaced_as_bad_gateway(exc_info.value) + + +@pytest.mark.asyncio +async def test_async_converse_stream_with_an_empty_200_body_raises_instead_of_an_empty_turn( + _aws_test_credentials: None, +) -> None: + async def _aiter_bytes(chunk_size: int | None = None) -> AsyncIterator[bytes]: + yield b"" + + response: Final = MagicMock(status_code=200, headers=_event_stream_headers()) + response.aiter_bytes = _aiter_bytes + client: Final = AsyncHTTPHandler() + client.post = AsyncMock(return_value=response) + + stream: Final = await litellm.acompletion( + model="bedrock/us.moonshotai.kimi-k3", + messages=[{"role": "user", "content": "hi"}], + stream=True, + client=client, + ) + with pytest.raises(MidStreamFallbackError) as exc_info: + _ = [chunk async for chunk in stream] + + _assert_empty_stream_surfaced_as_bad_gateway(exc_info.value) diff --git a/tests/unit/llms/test_polling_url_origin_match.py b/tests/unit/llms/test_polling_url_origin_match.py index ab5f41c757f..2df35131e3d 100644 --- a/tests/unit/llms/test_polling_url_origin_match.py +++ b/tests/unit/llms/test_polling_url_origin_match.py @@ -18,7 +18,7 @@ import pytest # Azure DALL-E sync + async paths route through ``assert_same_origin`` # the same way as the case below. The helper itself is unit-tested in -# ``tests/test_litellm/litellm_core_utils/test_url_utils.py``. +# ``tests/unit/litellm_core_utils/test_url_utils.py``. # ── Black Forest Labs polling ───────────────────────────────────────────────── diff --git a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_handler.py b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_handler.py index e0f0b7e5c0b..9b7cd127b83 100644 --- a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_handler.py +++ b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_handler.py @@ -208,6 +208,7 @@ class TestVertexAIFilesHandler: assert service_account == "/model/sa.json" def test_resolve_read_gcs_config_falls_back_to_env(self, monkeypatch): + monkeypatch.delenv("GCS_BATCH_BUCKET_NAME", raising=False) monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket") monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json") @@ -216,6 +217,40 @@ class TestVertexAIFilesHandler: assert bucket == "env-default-bucket" assert service_account == "/env/sa.json" + def test_resolve_read_gcs_config_prefers_batch_env_over_logging_env(self, monkeypatch): + monkeypatch.setenv("GCS_BATCH_BUCKET_NAME", "batch-bucket") + monkeypatch.setenv("GCS_BUCKET_NAME", "logging-bucket") + + bucket, _ = self.handler._resolve_read_gcs_config(litellm_params={}, vertex_credentials=None) + + assert bucket == "batch-bucket" + + def test_resolve_read_gcs_config_prefers_per_model_bucket_over_batch_env(self, monkeypatch): + monkeypatch.setenv("GCS_BATCH_BUCKET_NAME", "batch-bucket") + + bucket, _ = self.handler._resolve_read_gcs_config( + litellm_params={"gcs_bucket_name": "my-model-bucket"}, + vertex_credentials=None, + ) + + assert bucket == "my-model-bucket" + + def test_resolve_read_gcs_config_prefers_gcs_bucket_name_over_legacy(self): + bucket, _ = self.handler._resolve_read_gcs_config( + litellm_params={"gcs_bucket_name": "my-model-bucket", "bucket_name": "legacy-bucket"}, + vertex_credentials=None, + ) + + assert bucket == "my-model-bucket" + + def test_resolve_read_gcs_config_accepts_legacy_bucket_name_alone(self): + bucket, _ = self.handler._resolve_read_gcs_config( + litellm_params={"bucket_name": "legacy-bucket"}, + vertex_credentials=None, + ) + + assert bucket == "legacy-bucket" + def test_resolve_read_gcs_config_serializes_dict_credentials(self, monkeypatch): monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) diff --git a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index 7434eae72a4..6f18a391f7b 100644 --- a/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/unit/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -1186,10 +1186,21 @@ class TestConfiguredBucketNameResolution: assert config._get_configured_bucket_name({"gcs_bucket_name": "new", "bucket_name": "legacy"}) == "new" def test_should_fall_back_to_env(self, config, monkeypatch): + monkeypatch.delenv("GCS_BATCH_BUCKET_NAME", raising=False) monkeypatch.setenv("GCS_BUCKET_NAME", "env-bucket") assert config._get_configured_bucket_name({}) == "env-bucket" + def test_should_prefer_batch_env_over_logging_env(self, config, monkeypatch): + monkeypatch.setenv("GCS_BATCH_BUCKET_NAME", "batch-bucket") + monkeypatch.setenv("GCS_BUCKET_NAME", "logging-bucket") + assert config._get_configured_bucket_name({}) == "batch-bucket" + + def test_should_prefer_litellm_params_over_batch_env(self, config, monkeypatch): + monkeypatch.setenv("GCS_BATCH_BUCKET_NAME", "batch-bucket") + assert config._get_configured_bucket_name({"gcs_bucket_name": "per-model-bucket"}) == "per-model-bucket" + def test_should_raise_when_no_bucket_anywhere(self, config, monkeypatch): + monkeypatch.delenv("GCS_BATCH_BUCKET_NAME", raising=False) monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) with pytest.raises(ValueError, match="GCS bucket_name is required"): config._get_configured_bucket_name({}) diff --git a/tests/test_litellm/rust_bridge/responses/__init__.py b/tests/unit/llms/vertex_ai/rag_engine/__init__.py similarity index 100% rename from tests/test_litellm/rust_bridge/responses/__init__.py rename to tests/unit/llms/vertex_ai/rag_engine/__init__.py diff --git a/tests/unit/llms/vertex_ai/rag_engine/test_ingestion.py b/tests/unit/llms/vertex_ai/rag_engine/test_ingestion.py new file mode 100644 index 00000000000..3acabc4d14e --- /dev/null +++ b/tests/unit/llms/vertex_ai/rag_engine/test_ingestion.py @@ -0,0 +1,76 @@ +import asyncio +import sys +from types import ModuleType, SimpleNamespace + +import litellm +from litellm.litellm_core_utils.get_litellm_params import get_litellm_params +from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig +from litellm.llms.vertex_ai.rag_engine.ingestion import VertexAIRAGIngestion + + +def _ingestion_for_bucket(bucket: str) -> VertexAIRAGIngestion: + return VertexAIRAGIngestion( + { + "vector_store": { + "custom_llm_provider": "vertex_ai", + "vector_store_id": "corpus-123", + "vertex_project": "test-project", + "vertex_location": "us-central1", + "gcs_bucket": bucket, + } + } + ) + + +def test_upload_lands_in_the_corpus_bucket_when_batch_bucket_env_is_set(monkeypatch): + monkeypatch.setenv("GCS_BATCH_BUCKET_NAME", "batch-bucket") + monkeypatch.setenv("GCS_BUCKET_NAME", "logging-bucket") + resolver = VertexAIFilesConfig() + + async def acreate_file_through_real_bucket_resolver(**kwargs): + bucket = resolver._get_configured_bucket_name(get_litellm_params(**kwargs)) + return SimpleNamespace(id=f"gs://{bucket}/{kwargs['file'][0]}") + + monkeypatch.setattr(litellm, "acreate_file", acreate_file_through_real_bucket_resolver) + + uri = asyncio.run(_ingestion_for_bucket("rag-bucket")._upload_file_to_gcs(b"doc", "doc.txt", "text/plain")) + + assert uri == "gs://rag-bucket/doc.txt" + + +def _vertexai_sdk_stub(import_calls: list[dict[str, object]]) -> ModuleType: + rag = ModuleType("vertexai.rag") + rag.TransformationConfig = lambda chunking_config: chunking_config + rag.ChunkingConfig = lambda chunk_size, chunk_overlap: (chunk_size, chunk_overlap) + + def import_files(**kwargs): + import_calls.append(kwargs) + return SimpleNamespace(imported_rag_files_count=1) + + rag.import_files = import_files + vertexai = ModuleType("vertexai") + vertexai.init = lambda project, location: None + vertexai.rag = rag + return vertexai + + +def test_ingest_runs_end_to_end_through_the_base_pipeline(monkeypatch): + monkeypatch.setenv("GCS_BATCH_BUCKET_NAME", "batch-bucket") + resolver = VertexAIFilesConfig() + import_calls: list[dict[str, object]] = [] + stub = _vertexai_sdk_stub(import_calls) + monkeypatch.setitem(sys.modules, "vertexai", stub) + monkeypatch.setitem(sys.modules, "vertexai.rag", stub.rag) + + async def acreate_file_through_real_bucket_resolver(**kwargs): + bucket = resolver._get_configured_bucket_name(get_litellm_params(**kwargs)) + return SimpleNamespace(id=f"gs://{bucket}/{kwargs['file'][0]}") + + monkeypatch.setattr(litellm, "acreate_file", acreate_file_through_real_bucket_resolver) + + result = asyncio.run(_ingestion_for_bucket("rag-bucket").ingest(file_data=("doc.txt", b"doc", "text/plain"))) + + assert (result["status"], result["vector_store_id"], result["file_id"]) == ("completed", "corpus-123", "gs://rag-bucket/doc.txt") + assert [(c["corpus_name"], c["paths"]) for c in import_calls] == [ + ("projects/test-project/locations/us-central1/ragCorpora/corpus-123", ["gs://rag-bucket/doc.txt"]) + ] diff --git a/tests/unit/llms/xai/batches/__init__.py b/tests/unit/llms/xai/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/xai/batches/test_xai_batches_handler.py b/tests/unit/llms/xai/batches/test_xai_batches_handler.py new file mode 100644 index 00000000000..6dcdf06e7ab --- /dev/null +++ b/tests/unit/llms/xai/batches/test_xai_batches_handler.py @@ -0,0 +1,344 @@ +import json +from typing import Final + +import httpx +import pytest +import respx + +import litellm +from litellm.llms.xai.batches.transformation import XAIBatchesError +from litellm.types.utils import LiteLLMBatch + +API_BASE: Final = "https://api.x.ai" +KEY: Final = "xai-test-key" + + +@pytest.fixture(autouse=True) +def _httpx_transport_so_respx_can_intercept(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +_XAI_BATCH: Final = { + "batch_id": "batch_1", + "name": "litellm-batch", + "create_time": "2026-09-23", + "expire_time": "2026-10-23", + "cancel_time": None, + "cancel_by_xai_message": None, + "state": {"num_requests": 2, "num_pending": 0, "num_success": 2, "num_error": 0, "num_cancelled": 0}, + "input_file_id": "file_1", +} + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@respx.mock +async def test_create_batch_posts_input_file_id_with_bearer_auth(sync_mode: bool) -> None: + route: Final = respx.post(f"{API_BASE}/v1/batches").respond(200, json=_XAI_BATCH) + + kwargs: Final = { + "completion_window": "24h", + "endpoint": "/v1/embeddings", + "input_file_id": "file_1", + "custom_llm_provider": "xai", + "api_key": KEY, + "api_base": API_BASE, + } + batch: Final = litellm.create_batch(**kwargs) if sync_mode else await litellm.acreate_batch(**kwargs) + + assert isinstance(batch, LiteLLMBatch) + request: Final = route.calls.last.request + assert request.headers["authorization"] == f"Bearer {KEY}" + assert json.loads(request.content) == {"name": "litellm-batch", "input_file_id": "file_1"} + assert (batch.id, batch.endpoint, batch.status, batch.output_file_id) == ( + "batch_1", + "/v1/embeddings", + "completed", + "batch_1", + ) + + +@pytest.mark.parametrize( + "endpoint", + [ + "/v1/chat/completions", + "/v1/embeddings", + "/v1/completions", + "/v1/responses", + "/v1/ocr", + "/v1/images/generations", + "/v1/images/edits", + "/v1/videos/generations", + "/v1/videos", + "/v1/videos/edits", + "/v1/videos/extensions", + ], +) +@respx.mock +async def test_create_batch_keeps_image_and_video_endpoints_on_the_batch(endpoint: str) -> None: + respx.post(f"{API_BASE}/v1/batches").respond(200, json=_XAI_BATCH) + + batch: Final = await litellm.acreate_batch( + completion_window="24h", + endpoint=endpoint, + input_file_id="file_1", + custom_llm_provider="xai", + api_key=KEY, + api_base=API_BASE, + ) + + assert isinstance(batch, LiteLLMBatch) + assert batch.endpoint == endpoint + assert json.loads(respx.calls.last.request.content) == {"name": "litellm-batch", "input_file_id": "file_1"} + + +@respx.mock +async def test_retrieve_after_a_non_chat_create_reports_chat() -> None: + respx.post(f"{API_BASE}/v1/batches").respond(200, json=_XAI_BATCH) + respx.get(f"{API_BASE}/v1/batches/batch_1").respond(200, json=_XAI_BATCH) + + created: Final = await litellm.acreate_batch( + completion_window="24h", + endpoint="/v1/embeddings", + input_file_id="file_1", + custom_llm_provider="xai", + api_key=KEY, + api_base=API_BASE, + ) + retrieved: Final = await litellm.aretrieve_batch( + batch_id="batch_1", custom_llm_provider="xai", api_key=KEY, api_base=API_BASE + ) + + assert isinstance(created, LiteLLMBatch) and isinstance(retrieved, LiteLLMBatch) + assert (created.endpoint, retrieved.endpoint) == ("/v1/embeddings", "/v1/chat/completions") + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@respx.mock +async def test_retrieve_batch_reads_native_batch_route(sync_mode: bool) -> None: + respx.get(f"{API_BASE}/v1/batches/batch_1").respond( + 200, json={**_XAI_BATCH, "state": {"num_requests": 2, "num_pending": 2}} + ) + + kwargs: Final = {"batch_id": "batch_1", "custom_llm_provider": "xai", "api_key": KEY, "api_base": API_BASE} + batch: Final = litellm.retrieve_batch(**kwargs) if sync_mode else await litellm.aretrieve_batch(**kwargs) + + assert isinstance(batch, LiteLLMBatch) + assert (batch.status, batch.output_file_id, batch.input_file_id, batch.endpoint) == ( + "in_progress", + None, + "file_1", + "/v1/chat/completions", + ) + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@respx.mock +async def test_cancel_batch_uses_colon_cancel_route(sync_mode: bool) -> None: + route: Final = respx.post(f"{API_BASE}/v1/batches/batch_1:cancel").respond( + 200, json={**_XAI_BATCH, "cancel_time": "2026-09-23", "state": {}} + ) + + kwargs: Final = {"batch_id": "batch_1", "custom_llm_provider": "xai", "api_key": KEY, "api_base": API_BASE} + batch: Final = litellm.cancel_batch(**kwargs) if sync_mode else await litellm.acancel_batch(**kwargs) + + assert route.called + assert isinstance(batch, LiteLLMBatch) + assert (batch.status, batch.endpoint) == ("cancelled", "/v1/chat/completions") + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@respx.mock +async def test_list_batches_forwards_cursor_and_returns_openai_list(sync_mode: bool) -> None: + route: Final = respx.get(f"{API_BASE}/v1/batches").respond( + 200, json={"batches": [_XAI_BATCH], "pagination_token": "next"} + ) + + kwargs: Final = {"custom_llm_provider": "xai", "api_key": KEY, "api_base": API_BASE, "after": "cur", "limit": 5} + listed: Final = litellm.list_batches(**kwargs) if sync_mode else await litellm.alist_batches(**kwargs) + + assert dict(route.calls.last.request.url.params) == {"limit": "5", "pagination_token": "cur"} + assert listed.object == "list" + assert [(b.id, b.endpoint) for b in listed.data] == [("batch_1", "/v1/chat/completions")] + assert (listed.has_more, listed.next_page_token) == (True, "next") + + +@respx.mock +async def test_list_batches_treats_empty_pagination_token_as_last_page() -> None: + respx.get(f"{API_BASE}/v1/batches").respond(200, json={"batches": [_XAI_BATCH], "pagination_token": ""}) + + listed: Final = await litellm.alist_batches(custom_llm_provider="xai", api_key=KEY, api_base=API_BASE) + + assert (listed.has_more, listed.next_page_token) == (False, None) + assert [batch.endpoint for batch in listed.data] == ["/v1/chat/completions"] + + +@respx.mock +async def test_file_content_stops_paging_on_empty_pagination_token() -> None: + route: Final = respx.get(f"{API_BASE}/v1/batches/batch_1/results").respond( + 200, + json={ + "results": [{"batch_request_id": "r1", "batch_result": {"error": {"code": 3, "message": "boom"}}}], + "pagination_token": "", + }, + ) + + content: Final = await litellm.afile_content( + file_id="batch_1", custom_llm_provider="xai", api_key=KEY, api_base=API_BASE + ) + + assert route.call_count == 1 + assert len(content.content.decode().splitlines()) == 1 + + +@pytest.mark.parametrize("operation", ["create", "retrieve", "cancel", "list", "file_content"]) +@respx.mock +async def test_batch_calls_fall_back_to_litellm_xai_key(operation: str, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("XAI_API_KEY", raising=False) + monkeypatch.setattr(litellm, "xai_key", "configured-xai-key") + monkeypatch.setattr(litellm, "api_key", "generic-key-must-not-be-used") + routes: Final = { + "create": respx.post(f"{API_BASE}/v1/batches").respond(200, json=_XAI_BATCH), + "retrieve": respx.get(f"{API_BASE}/v1/batches/batch_1").respond(200, json=_XAI_BATCH), + "cancel": respx.post(f"{API_BASE}/v1/batches/batch_1:cancel").respond(200, json=_XAI_BATCH), + "list": respx.get(f"{API_BASE}/v1/batches").respond( + 200, json={"batches": [_XAI_BATCH], "pagination_token": None} + ), + "file_content": respx.get(f"{API_BASE}/v1/batches/batch_1/results").respond( + 200, json={"results": [], "pagination_token": None} + ), + } + + if operation == "create": + await litellm.acreate_batch( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id="file_1", + custom_llm_provider="xai", + api_base=API_BASE, + ) + elif operation == "retrieve": + await litellm.aretrieve_batch(batch_id="batch_1", custom_llm_provider="xai", api_base=API_BASE) + elif operation == "cancel": + await litellm.acancel_batch(batch_id="batch_1", custom_llm_provider="xai", api_base=API_BASE) + elif operation == "list": + await litellm.alist_batches(custom_llm_provider="xai", api_base=API_BASE) + else: + await litellm.afile_content(file_id="batch_1", custom_llm_provider="xai", api_base=API_BASE) + + assert routes[operation].calls.last.request.headers["authorization"] == "Bearer configured-xai-key" + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@respx.mock +async def test_file_content_of_a_batch_id_walks_every_results_page(sync_mode: bool) -> None: + def _page(request: httpx.Request) -> httpx.Response: + token: Final = request.url.params.get("pagination_token") + if token is None: + return httpx.Response( + 200, + json={ + "results": [ + { + "batch_request_id": "r1", + "batch_result": {"response": {"chat_get_completion": {"id": "c1", "choices": []}}}, + } + ], + "pagination_token": "r1", + }, + ) + assert token == "r1" + return httpx.Response( + 200, + json={ + "results": [ + {"batch_request_id": "r2", "batch_result": {"error": {"code": 3, "message": "boom"}}}, + ], + "pagination_token": None, + }, + ) + + route: Final = respx.get(f"{API_BASE}/v1/batches/batch_1/results").mock(side_effect=_page) + + kwargs: Final = {"file_id": "batch_1", "custom_llm_provider": "xai", "api_key": KEY, "api_base": API_BASE} + content: Final = litellm.file_content(**kwargs) if sync_mode else await litellm.afile_content(**kwargs) + + assert route.call_count == 2 + assert [dict(c.request.url.params) for c in route.calls] == [ + {"limit": "1000"}, + {"limit": "1000", "pagination_token": "r1"}, + ] + assert [json.loads(line) for line in content.content.decode().splitlines()] == [ + { + "id": "batch_req_r1", + "custom_id": "r1", + "response": {"status_code": 200, "request_id": "c1", "body": {"id": "c1", "choices": []}}, + "error": None, + }, + {"id": "batch_req_r2", "custom_id": "r2", "response": None, "error": {"code": "3", "message": "boom"}}, + ] + + +@respx.mock +async def test_file_content_unwraps_image_and_video_result_bodies() -> None: + respx.get(f"{API_BASE}/v1/batches/batch_1/results").respond( + 200, + json={ + "results": [ + { + "batch_request_id": "img", + "batch_result": { + "response": {"image_generation": {"data": [{"url": "https://cdn.example/img.png"}]}} + }, + }, + { + "batch_request_id": "vid", + "batch_result": { + "response": {"video_generation": {"id": "vid_1", "url": "https://cdn.example/clip.mp4"}} + }, + }, + ], + "pagination_token": None, + }, + ) + + content: Final = await litellm.afile_content( + file_id="batch_1", custom_llm_provider="xai", api_key=KEY, api_base=API_BASE + ) + + assert [json.loads(line)["response"]["body"] for line in content.content.decode().splitlines()] == [ + {"data": [{"url": "https://cdn.example/img.png"}]}, + {"id": "vid_1", "url": "https://cdn.example/clip.mp4"}, + ] + + +@respx.mock +async def test_missing_xai_key_is_a_401_before_any_request(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("XAI_API_KEY", raising=False) + monkeypatch.setattr(litellm, "xai_key", None) + monkeypatch.setattr(litellm, "api_key", "generic-key-must-not-be-used") + route: Final = respx.post(f"{API_BASE}/v1/batches").respond(200, json=_XAI_BATCH) + + with pytest.raises(XAIBatchesError) as exc: + await litellm.acreate_batch( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id="file_1", + custom_llm_provider="xai", + api_base=API_BASE, + ) + + assert exc.value.status_code == 401 + assert route.called is False + + +@respx.mock +async def test_upstream_error_surfaces_status_code_and_body() -> None: + respx.get(f"{API_BASE}/v1/batches/batch_missing").respond(404, json={"code": "404", "error": "not found"}) + + with pytest.raises(XAIBatchesError) as exc: + await litellm.aretrieve_batch( + batch_id="batch_missing", custom_llm_provider="xai", api_key=KEY, api_base=API_BASE + ) + + assert exc.value.status_code == 404 + assert "not found" in exc.value.message diff --git a/tests/unit/llms/xai/batches/test_xai_batches_transformation.py b/tests/unit/llms/xai/batches/test_xai_batches_transformation.py new file mode 100644 index 00000000000..5f2bb6a33ce --- /dev/null +++ b/tests/unit/llms/xai/batches/test_xai_batches_transformation.py @@ -0,0 +1,224 @@ +import json +from typing import Final + +import pytest + +from litellm.llms.xai.batches.transformation import ( + XAIBatch, + XAIBatchesError, + XAIBatchList, + XAIBatchResult, + XAIBatchResultsPage, + get_xai_api_base, + results_to_openai_jsonl, + to_create_batch_body, + to_litellm_batch, + to_openai_batch_list, + xai_batches_url, +) +from litellm.types.llms.openai import CreateBatchRequest + +SEPT_23_2026_UTC: Final = 1790121600 + + +def _xai_batch(**overrides: object) -> XAIBatch: + return XAIBatch.model_validate( + { + "batch_id": "batch_9bdf", + "name": "nightly", + "create_time": "2026-09-23", + "expire_time": "2026-10-23", + "cancel_time": None, + "cancel_by_xai_message": None, + "state": {"num_requests": 2, "num_pending": 0, "num_success": 2, "num_error": 0, "num_cancelled": 0}, + "input_file_id": "file_07", + **overrides, + } + ) + + +def test_completed_batch_exposes_batch_id_as_output_file_and_maps_counts() -> None: + batch: Final = to_litellm_batch(_xai_batch()) + + assert batch.model_dump(exclude_none=True) == { + "id": "batch_9bdf", + "object": "batch", + "endpoint": "/v1/chat/completions", + "input_file_id": "file_07", + "completion_window": "24h", + "status": "completed", + "created_at": SEPT_23_2026_UTC, + "expires_at": SEPT_23_2026_UTC + 30 * 86400, + "output_file_id": "batch_9bdf", + "request_counts": {"total": 2, "completed": 2, "failed": 0}, + "metadata": {"name": "nightly"}, + } + + +def test_pending_requests_mean_in_progress_and_no_output_file() -> None: + batch: Final = to_litellm_batch( + _xai_batch(state={"num_requests": 3, "num_pending": 1, "num_success": 1, "num_error": 1, "num_cancelled": 0}) + ) + + assert (batch.status, batch.output_file_id) == ("in_progress", None) + assert batch.request_counts is not None + assert batch.request_counts.model_dump() == {"total": 3, "completed": 1, "failed": 1} + + +def test_empty_batch_is_still_validating() -> None: + assert to_litellm_batch(_xai_batch(state={})).status == "validating" + + +def test_batch_cancelled_by_xai_validation_is_failed_with_the_message() -> None: + batch: Final = to_litellm_batch( + _xai_batch( + state={}, + cancel_time="2026-09-23T10:00:00Z", + cancel_by_xai_message="JSONL file validation failed: Model grok-nope is not supported", + ) + ) + + assert batch.status == "failed" + assert batch.failed_at == SEPT_23_2026_UTC + 10 * 3600 + assert batch.cancelled_at is None + assert batch.errors is not None and batch.errors.data is not None + assert [e.message for e in batch.errors.data] == ["JSONL file validation failed: Model grok-nope is not supported"] + + +def test_batch_cancelled_by_caller_is_cancelled() -> None: + batch: Final = to_litellm_batch(_xai_batch(cancel_time="2026-09-23")) + + assert (batch.status, batch.cancelled_at, batch.errors) == ("cancelled", SEPT_23_2026_UTC, None) + + +@pytest.mark.parametrize( + "endpoint", + [ + "/v1/images/generations", + "/v1/images/edits", + "/v1/videos/generations", + "/v1/videos/edits", + "/v1/videos/extensions", + ], +) +def test_create_body_accepts_image_and_video_endpoints(endpoint: str) -> None: + body: Final = to_create_batch_body( + CreateBatchRequest(completion_window="24h", endpoint=endpoint, input_file_id="file_07") + ) + + assert dict(body) == {"name": "litellm-batch", "input_file_id": "file_07"} + + +def test_create_body_uses_input_file_id_and_metadata_name() -> None: + body: Final = to_create_batch_body( + CreateBatchRequest( + completion_window="24h", endpoint="/v1/chat/completions", input_file_id="file_07", metadata={"name": "n1"} + ) + ) + + assert dict(body) == {"name": "n1", "input_file_id": "file_07"} + + +def test_create_body_without_input_file_id_is_a_400() -> None: + with pytest.raises(XAIBatchesError) as exc: + to_create_batch_body(CreateBatchRequest(completion_window="24h", endpoint="/v1/chat/completions")) + + assert exc.value.status_code == 400 + + +def test_results_render_as_openai_output_jsonl_with_errors_per_line() -> None: + page: Final = XAIBatchResultsPage.model_validate( + { + "results": [ + { + "batch_request_id": "r1", + "batch_result": { + "response": { + "chat_get_completion": {"id": "c1", "object": "chat.completion", "choices": [], "usage": {}} + } + }, + }, + {"batch_request_id": "r2", "batch_result": {"error": {"code": 3, "message": "bad model"}}}, + {"batch_request_id": "r3", "batch_result": {}}, + ], + "pagination_token": None, + } + ) + + lines: Final = [json.loads(line) for line in results_to_openai_jsonl(page.results).decode().splitlines()] + + assert lines == [ + { + "id": "batch_req_r1", + "custom_id": "r1", + "response": { + "status_code": 200, + "request_id": "c1", + "body": {"id": "c1", "object": "chat.completion", "choices": [], "usage": {}}, + }, + "error": None, + }, + {"id": "batch_req_r2", "custom_id": "r2", "response": None, "error": {"code": "3", "message": "bad model"}}, + { + "id": "batch_req_r3", + "custom_id": "r3", + "response": None, + "error": {"code": "request_failed", "message": "xAI returned no response for this request"}, + }, + ] + + +@pytest.mark.parametrize( + ("response_key", "body"), + [ + ("responses", {"id": "resp_1", "output": []}), + ("image_generation", {"created": 1, "data": [{"url": "https://cdn.example/img.png"}]}), + ("video_generation", {"id": "vid_1", "url": "https://cdn.example/clip.mp4"}), + ], +) +def test_result_unwraps_the_single_response_key_into_the_openai_body( + response_key: str, body: dict[str, object] +) -> None: + result: Final = XAIBatchResult.model_validate( + {"batch_request_id": "r", "batch_result": {"response": {response_key: body}}} + ) + + line: Final = json.loads(results_to_openai_jsonl((result,)).decode()) + assert line["response"]["body"] == body + assert line["response"]["request_id"] == body.get("id") + assert response_key not in line["response"]["body"] + + +def test_retrieve_and_list_report_chat_because_xai_has_no_batch_endpoint() -> None: + retrieved: Final = to_litellm_batch(_xai_batch()) + listed: Final = to_openai_batch_list(XAIBatchList.model_validate({"batches": [_xai_batch().model_dump()]})) + + assert retrieved.endpoint == "/v1/chat/completions" + assert [batch.endpoint for batch in listed.data] == ["/v1/chat/completions"] + assert retrieved.metadata == {"name": "nightly"} + + +def test_list_page_maps_to_openai_list_with_cursor_flags() -> None: + page: Final = XAIBatchList.model_validate( + {"batches": [_xai_batch().model_dump(), _xai_batch(batch_id="batch_2").model_dump()], "pagination_token": "t"} + ) + + listed: Final = to_openai_batch_list(page) + + assert (listed.object, listed.first_id, listed.last_id, listed.has_more, listed.next_page_token) == ( + "list", + "batch_9bdf", + "batch_2", + True, + "t", + ) + assert [b.id for b in listed.data] == ["batch_9bdf", "batch_2"] + + +@pytest.mark.parametrize( + "api_base", ["https://api.x.ai", "https://api.x.ai/", "https://api.x.ai/v1", "https://api.x.ai/v1/"] +) +def test_api_base_never_doubles_the_v1_segment(api_base: str) -> None: + assert get_xai_api_base(api_base) == "https://api.x.ai" + assert xai_batches_url(api_base, "batch_1", ":cancel") == "https://api.x.ai/v1/batches/batch_1:cancel" + assert xai_batches_url(api_base) == "https://api.x.ai/v1/batches" diff --git a/tests/unit/llms/xai/files/__init__.py b/tests/unit/llms/xai/files/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/xai/files/test_xai_files_transformation.py b/tests/unit/llms/xai/files/test_xai_files_transformation.py new file mode 100644 index 00000000000..5a7d86bdfb7 --- /dev/null +++ b/tests/unit/llms/xai/files/test_xai_files_transformation.py @@ -0,0 +1,144 @@ +from typing import Final + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.types.llms.openai import OpenAIFileObject + +API_BASE: Final = "https://api.x.ai" +KEY: Final = "xai-test-key" + + +@pytest.fixture(autouse=True) +def _httpx_transport_so_respx_can_intercept(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + +_XAI_FILE: Final = { + "bytes": 337, + "created_at": 1790197740, + "expires_at": None, + "filename": "batch.jsonl", + "id": "file_07", + "object": "file", + "purpose": "", +} + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@respx.mock +async def test_create_file_uploads_multipart_to_xai_and_reports_batch_purpose(sync_mode: bool) -> None: + route: Final = respx.post(f"{API_BASE}/v1/files").respond(200, json=_XAI_FILE) + + kwargs: Final = { + "file": ("batch.jsonl", b'{"custom_id":"r1"}\n', "application/jsonl"), + "purpose": "batch", + "custom_llm_provider": "xai", + "api_key": KEY, + "api_base": API_BASE, + } + created: Final = litellm.create_file(**kwargs) if sync_mode else await litellm.acreate_file(**kwargs) + + request: Final = route.calls.last.request + assert request.headers["authorization"] == f"Bearer {KEY}" + assert request.headers["content-type"].startswith("multipart/form-data") + assert b'filename="batch.jsonl"' in request.content + assert b'{"custom_id":"r1"}' in request.content + assert created.model_dump(exclude_none=True) == { + "id": "file_07", + "bytes": 337, + "created_at": 1790197740, + "filename": "batch.jsonl", + "object": "file", + "purpose": "batch", + "status": "uploaded", + } + + +@respx.mock +async def test_file_content_of_an_uploaded_file_downloads_original_bytes() -> None: + respx.get(f"{API_BASE}/v1/files/file_07/content").respond(200, content=b'{"custom_id":"r1"}\n') + + content: Final = await litellm.afile_content( + file_id="file_07", custom_llm_provider="xai", api_key=KEY, api_base=API_BASE + ) + + assert content.content == b'{"custom_id":"r1"}\n' + + +@respx.mock +async def test_delete_file_maps_xai_deleted_object() -> None: + respx.delete(f"{API_BASE}/v1/files/file_07").respond(200, json={"id": "file_07", "deleted": True, "object": "file"}) + + deleted: Final = await litellm.afile_delete( + file_id="file_07", custom_llm_provider="xai", api_key=KEY, api_base=API_BASE + ) + + assert deleted.model_dump() == {"id": "file_07", "deleted": True, "object": "file"} + + +@respx.mock +async def test_create_file_falls_back_to_litellm_xai_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("XAI_API_KEY", raising=False) + monkeypatch.setattr(litellm, "xai_key", "configured-xai-key") + monkeypatch.setattr(litellm, "api_key", "generic-key-must-not-be-used") + route: Final = respx.post(f"{API_BASE}/v1/files").respond(200, json=_XAI_FILE) + + await litellm.acreate_file( + file=("batch.jsonl", b'{"custom_id":"r1"}\n', "application/jsonl"), + purpose="batch", + custom_llm_provider="xai", + api_base=API_BASE, + ) + + assert route.calls.last.request.headers["authorization"] == "Bearer configured-xai-key" + + +@respx.mock +async def test_list_files_reads_data_array() -> None: + respx.get(f"{API_BASE}/v1/files").respond(200, json={"data": [_XAI_FILE], "pagination_token": None}) + + listed: Final = await litellm.afile_list(custom_llm_provider="xai", api_key=KEY, api_base=API_BASE) + + files: Final = TypeAdapter(tuple[OpenAIFileObject, ...]).validate_python(listed) + assert [f.id for f in files] == ["file_07"] + + +@respx.mock +async def test_list_files_walks_every_page_by_pagination_token() -> None: + route: Final = respx.get(f"{API_BASE}/v1/files").mock( + side_effect=[ + httpx.Response(200, json={"data": [_XAI_FILE], "pagination_token": "file_07"}), + httpx.Response(200, json={"data": [{**_XAI_FILE, "id": "file_08"}], "pagination_token": "file_08"}), + httpx.Response(200, json={"data": [], "pagination_token": "file_08"}), + ] + ) + + listed: Final = await litellm.afile_list(custom_llm_provider="xai", api_key=KEY, api_base=API_BASE) + + files: Final = TypeAdapter(tuple[OpenAIFileObject, ...]).validate_python(listed) + assert [f.id for f in files] == ["file_07", "file_08"] + assert [call.request.url.params.get("pagination_token") for call in route.calls] == [None, "file_07", "file_08"] + + +async def _retrieve_file(sync_mode: bool, file_id: str) -> None: + if sync_mode: + litellm.file_retrieve(file_id=file_id, custom_llm_provider="xai", api_key=KEY, api_base=API_BASE) + return + await litellm.afile_retrieve(file_id=file_id, custom_llm_provider="xai", api_key=KEY, api_base=API_BASE) + + +@pytest.mark.parametrize("sync_mode", [True, False]) +@respx.mock +async def test_retrieve_file_maps_xai_not_found_to_a_404_error(sync_mode: bool) -> None: + respx.get(f"{API_BASE}/v1/files/file_gone").respond(404, json={"code": "not-found", "error": "File not found"}) + + with pytest.raises(BaseLLMException) as raised: + await _retrieve_file(sync_mode, "file_gone") + + assert raised.value.status_code == 404 + assert "File not found" in str(raised.value) diff --git a/tests/unit/llms/xai/test_xai_chat_transformation.py b/tests/unit/llms/xai/test_xai_chat_transformation.py index 3fd666e4f50..704d0061103 100644 --- a/tests/unit/llms/xai/test_xai_chat_transformation.py +++ b/tests/unit/llms/xai/test_xai_chat_transformation.py @@ -16,7 +16,7 @@ from litellm.types.utils import ( class TestXAIReasoningTokenFolding: - """``_fold_reasoning_tokens_into_completion`` re-aligns xAI Usage to the OpenAI invariant.""" + """``fold_reasoning_tokens_into_completion`` re-aligns xAI Usage to the OpenAI invariant.""" @staticmethod def _make_response( @@ -45,7 +45,7 @@ class TestXAIReasoningTokenFolding: reasoning_tokens=312, ) - XAIChatConfig._fold_reasoning_tokens_into_completion(response) + XAIChatConfig.fold_reasoning_tokens_into_completion(response) usage = response.usage assert usage.completion_tokens == 322 @@ -59,7 +59,7 @@ class TestXAIReasoningTokenFolding: reasoning_tokens=312, ) - XAIChatConfig._fold_reasoning_tokens_into_completion(response) + XAIChatConfig.fold_reasoning_tokens_into_completion(response) assert response.usage.completion_tokens == 322 @@ -71,7 +71,7 @@ class TestXAIReasoningTokenFolding: reasoning_tokens=0, ) - XAIChatConfig._fold_reasoning_tokens_into_completion(response) + XAIChatConfig.fold_reasoning_tokens_into_completion(response) assert response.usage.completion_tokens == 10 @@ -84,7 +84,7 @@ class TestXAIReasoningTokenFolding: reasoning_tokens=312, ) - XAIChatConfig._fold_reasoning_tokens_into_completion(response) + XAIChatConfig.fold_reasoning_tokens_into_completion(response) assert response.usage.completion_tokens == 10 assert response.usage.total_tokens == 999 diff --git a/tests/unit/proxy/common_utils/test_validation_error_body.py b/tests/unit/proxy/common_utils/test_validation_error_body.py new file mode 100644 index 00000000000..a86f17b7461 --- /dev/null +++ b/tests/unit/proxy/common_utils/test_validation_error_body.py @@ -0,0 +1,46 @@ +from typing import Final + +from litellm.proxy.common_utils.validation_error_body import public_validation_errors + +_PASSWORD: Final = "hunter2-Sup3rSecret!" + + +def test_public_validation_errors_drops_input_ctx_and_url(): + errors: Final = ( + { + "type": "missing", + "loc": ("body", "user_id"), + "msg": "Field required", + "input": {"invitation_link": "abc", "password": _PASSWORD}, + "url": "https://errors.pydantic.dev/2/v/missing", + }, + { + "type": "value_error", + "loc": ("body", "password"), + "msg": "Value error, password cannot be set here", + "input": _PASSWORD, + "ctx": {"error": ValueError(_PASSWORD)}, + }, + ) + + public: Final = public_validation_errors(errors) + + assert public == ( + {"type": "missing", "loc": ("body", "user_id"), "msg": "Field required"}, + {"type": "value_error", "loc": ("body", "password"), "msg": "Value error, password cannot be set here"}, + ) + assert _PASSWORD not in repr(public) + + +def test_public_validation_errors_keeps_type_loc_and_msg_verbatim_in_order(): + errors: Final = ( + {"type": "int_parsing", "loc": ("body", "litellm_params", "rpm"), "msg": "Input should be a valid integer"}, + {"type": "extra_forbidden", "loc": ("body", "users", 0, "user_emial"), "msg": "Extra inputs are not permitted"}, + {"type": "too_short", "loc": ("body", "users"), "msg": "List should have at least 1 item"}, + ) + + assert public_validation_errors(errors) == errors + + +def test_public_validation_errors_empty_in_empty_out(): + assert public_validation_errors(()) == () diff --git a/tests/unit/responses/litellm_completion_transformation/__init__.py b/tests/unit/responses/litellm_completion_transformation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py b/tests/unit/responses/litellm_completion_transformation/test_function_call_output_normalization.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py rename to tests/unit/responses/litellm_completion_transformation/test_function_call_output_normalization.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_handler.py b/tests/unit/responses/litellm_completion_transformation/test_handler.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_handler.py rename to tests/unit/responses/litellm_completion_transformation/test_handler.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_image_generation_output.py b/tests/unit/responses/litellm_completion_transformation/test_image_generation_output.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_image_generation_output.py rename to tests/unit/responses/litellm_completion_transformation/test_image_generation_output.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py rename to tests/unit/responses/litellm_completion_transformation/test_litellm_completion_responses.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_reasoning_input_item_preservation.py b/tests/unit/responses/litellm_completion_transformation/test_reasoning_input_item_preservation.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_reasoning_input_item_preservation.py rename to tests/unit/responses/litellm_completion_transformation/test_reasoning_input_item_preservation.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py b/tests/unit/responses/litellm_completion_transformation/test_session_handler.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py rename to tests/unit/responses/litellm_completion_transformation/test_session_handler.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler_with_cold_storage.py b/tests/unit/responses/litellm_completion_transformation/test_session_handler_with_cold_storage.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_session_handler_with_cold_storage.py rename to tests/unit/responses/litellm_completion_transformation/test_session_handler_with_cold_storage.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py b/tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py rename to tests/unit/responses/litellm_completion_transformation/test_streaming_iterator_transformation.py diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py b/tests/unit/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py similarity index 100% rename from tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py rename to tests/unit/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py diff --git a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py b/tests/unit/responses/mcp/test_chat_completions_handler.py similarity index 100% rename from tests/test_litellm/responses/mcp/test_chat_completions_handler.py rename to tests/unit/responses/mcp/test_chat_completions_handler.py diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py similarity index 100% rename from tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py rename to tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/unit/responses/mcp/test_mcp_streaming_iterator.py similarity index 100% rename from tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py rename to tests/unit/responses/mcp/test_mcp_streaming_iterator.py diff --git a/tests/test_litellm/responses/test_additional_tools.py b/tests/unit/responses/test_additional_tools.py similarity index 100% rename from tests/test_litellm/responses/test_additional_tools.py rename to tests/unit/responses/test_additional_tools.py diff --git a/tests/test_litellm/responses/test_custom_tool_call.py b/tests/unit/responses/test_custom_tool_call.py similarity index 100% rename from tests/test_litellm/responses/test_custom_tool_call.py rename to tests/unit/responses/test_custom_tool_call.py diff --git a/tests/test_litellm/responses/test_dispatch.py b/tests/unit/responses/test_dispatch.py similarity index 100% rename from tests/test_litellm/responses/test_dispatch.py rename to tests/unit/responses/test_dispatch.py diff --git a/tests/test_litellm/responses/test_metadata_codex_callback.py b/tests/unit/responses/test_metadata_codex_callback.py similarity index 100% rename from tests/test_litellm/responses/test_metadata_codex_callback.py rename to tests/unit/responses/test_metadata_codex_callback.py diff --git a/tests/test_litellm/responses/test_no_duplicate_spend_logs.py b/tests/unit/responses/test_no_duplicate_spend_logs.py similarity index 76% rename from tests/test_litellm/responses/test_no_duplicate_spend_logs.py rename to tests/unit/responses/test_no_duplicate_spend_logs.py index c98b519ae67..7e4bef5812c 100644 --- a/tests/test_litellm/responses/test_no_duplicate_spend_logs.py +++ b/tests/unit/responses/test_no_duplicate_spend_logs.py @@ -15,35 +15,6 @@ import litellm from litellm.integrations.custom_logger import CustomLogger -def test_logging_object_not_popped(): - """ - Test that litellm_logging_obj is not popped from kwargs. - - This is a regression test for issue #15740. The bug was using - kwargs.pop() which removed the logging object, causing duplicate - spend logs for non-OpenAI providers. - """ - import inspect - - from litellm.responses import main as responses_module - - # Get the source code of the responses function - source = inspect.getsource(responses_module.responses) - - # Check that .pop("litellm_logging_obj") is NOT used - # The bug was using kwargs.pop("litellm_logging_obj") which removes it - assert 'kwargs.pop("litellm_logging_obj")' not in source, ( - "FAIL: Found kwargs.pop('litellm_logging_obj') in responses() function. " - "This causes duplicate spend logs. Use kwargs.get('litellm_logging_obj') instead." - ) - - # Check that .get("litellm_logging_obj") IS used - assert 'kwargs.get("litellm_logging_obj")' in source, ( - "FAIL: Expected kwargs.get('litellm_logging_obj') but not found. " - "The logging object must be accessed with .get() not .pop() to prevent duplication." - ) - - @pytest.mark.asyncio async def test_async_no_duplicate_spend_logs(): """ diff --git a/tests/test_litellm/responses/test_null_test_fix.py b/tests/unit/responses/test_null_test_fix.py similarity index 100% rename from tests/test_litellm/responses/test_null_test_fix.py rename to tests/unit/responses/test_null_test_fix.py diff --git a/tests/test_litellm/responses/test_responses_api_bridge_flag.py b/tests/unit/responses/test_responses_api_bridge_flag.py similarity index 100% rename from tests/test_litellm/responses/test_responses_api_bridge_flag.py rename to tests/unit/responses/test_responses_api_bridge_flag.py diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/unit/responses/test_responses_api_request_body.py similarity index 99% rename from tests/test_litellm/responses/test_responses_api_request_body.py rename to tests/unit/responses/test_responses_api_request_body.py index 98e74955c6f..b27401d693a 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/unit/responses/test_responses_api_request_body.py @@ -20,7 +20,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler def _expected_dir() -> Path: - """Path to expected_responses_api_request folder (sibling of test_litellm/responses).""" + """Path to expected_responses_api_request folder (sibling of tests/unit/responses).""" return Path(__file__).resolve().parent.parent / "expected_responses_api_request" diff --git a/tests/test_litellm/responses/test_responses_prompt_management.py b/tests/unit/responses/test_responses_prompt_management.py similarity index 100% rename from tests/test_litellm/responses/test_responses_prompt_management.py rename to tests/unit/responses/test_responses_prompt_management.py diff --git a/tests/test_litellm/responses/test_responses_router_cooldown.py b/tests/unit/responses/test_responses_router_cooldown.py similarity index 100% rename from tests/test_litellm/responses/test_responses_router_cooldown.py rename to tests/unit/responses/test_responses_router_cooldown.py diff --git a/tests/test_litellm/responses/test_responses_streaming_iterator.py b/tests/unit/responses/test_responses_streaming_iterator.py similarity index 100% rename from tests/test_litellm/responses/test_responses_streaming_iterator.py rename to tests/unit/responses/test_responses_streaming_iterator.py diff --git a/tests/test_litellm/responses/test_responses_supported_endpoints_passthrough.py b/tests/unit/responses/test_responses_supported_endpoints_passthrough.py similarity index 100% rename from tests/test_litellm/responses/test_responses_supported_endpoints_passthrough.py rename to tests/unit/responses/test_responses_supported_endpoints_passthrough.py diff --git a/tests/test_litellm/responses/test_responses_utils.py b/tests/unit/responses/test_responses_utils.py similarity index 100% rename from tests/test_litellm/responses/test_responses_utils.py rename to tests/unit/responses/test_responses_utils.py diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/unit/responses/test_responses_websocket_all_providers.py similarity index 97% rename from tests/test_litellm/responses/test_responses_websocket_all_providers.py rename to tests/unit/responses/test_responses_websocket_all_providers.py index 3888a84fb5d..6f346a25d9c 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -2718,97 +2718,6 @@ class TestWebSocketChunkTypes: assert "response.reasoning_content.done" in serialized assert "Complete reasoning" in serialized - def test_extract_output_messages_preserves_multiple_messages(self): - """Test that multiple output messages are all preserved""" - from litellm.responses.streaming_iterator import ( - ManagedResponsesWebSocketHandler, - ) - - completed_event = { - "type": "response.completed", - "response": { - "id": "resp_123", - "output": [ - { - "type": "message", - "role": "assistant", - "content": [{"type": "output_text", "text": "First message"}], - }, - { - "type": "function_call", - "id": "call_123", - "name": "get_weather", - "arguments": "{}", - }, - { - "type": "message", - "role": "assistant", - "content": [{"type": "output_text", "text": "Second message"}], - }, - ], - }, - } - - messages = ManagedResponsesWebSocketHandler._extract_output_messages( - completed_event - ) - assert len(messages) == 3 - assert messages[0]["content"][0]["text"] == "First message" - assert messages[1]["type"] == "function_call" - assert messages[2]["content"][0]["text"] == "Second message" - - def test_input_to_messages_with_mixed_content_types(self): - """Test input conversion with mixed content types""" - from litellm.responses.streaming_iterator import ( - ManagedResponsesWebSocketHandler, - ) - - input_list = [ - { - "type": "message", - "role": "user", - "content": [ - {"type": "input_text", "text": "Question"}, - {"type": "input_image", "image_url": "https://example.com/img.png"}, - ], - } - ] - - messages = ManagedResponsesWebSocketHandler._input_to_messages(input_list) - assert len(messages) == 1 - assert len(messages[0]["content"]) == 2 - assert messages[0]["content"][0]["type"] == "input_text" - assert messages[0]["content"][1]["type"] == "input_image" - - def test_extract_output_messages_with_mixed_text_types(self): - """Test that both 'output_text' and 'text' types are extracted""" - from litellm.responses.streaming_iterator import ( - ManagedResponsesWebSocketHandler, - ) - - completed_event = { - "type": "response.completed", - "response": { - "id": "resp_123", - "output": [ - { - "type": "message", - "role": "assistant", - "content": [ - {"type": "output_text", "text": "Part 1"}, - {"type": "text", "text": "Part 2"}, - ], - } - ], - }, - } - - messages = ManagedResponsesWebSocketHandler._extract_output_messages( - completed_event - ) - assert len(messages) == 1 - assert messages[0]["content"][0]["text"] == "Part 1Part 2" - class TestNativeWebSocketUrlConstruction: """Test that native WebSocket URLs include the model query parameter. diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/unit/responses/test_rust_bridge_websocket.py similarity index 100% rename from tests/test_litellm/responses/test_rust_bridge_websocket.py rename to tests/unit/responses/test_rust_bridge_websocket.py diff --git a/tests/test_litellm/responses/test_sse_output_recovery.py b/tests/unit/responses/test_sse_output_recovery.py similarity index 100% rename from tests/test_litellm/responses/test_sse_output_recovery.py rename to tests/unit/responses/test_sse_output_recovery.py diff --git a/tests/test_litellm/responses/test_streaming_iterator.py b/tests/unit/responses/test_streaming_iterator.py similarity index 82% rename from tests/test_litellm/responses/test_streaming_iterator.py rename to tests/unit/responses/test_streaming_iterator.py index dbf54ec3b9b..9dbbc20591e 100644 --- a/tests/test_litellm/responses/test_streaming_iterator.py +++ b/tests/unit/responses/test_streaming_iterator.py @@ -6,13 +6,14 @@ completion_start_time = end_time.""" import json from datetime import datetime from typing import Final, Optional -from unittest.mock import Mock, patch +from unittest.mock import AsyncMock, Mock, patch import httpx import pytest from pydantic_core import PydanticSerializationError import litellm +from litellm.exceptions import MidStreamFallbackError from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.streaming_iterator import ( @@ -251,6 +252,171 @@ def test_sync_transport_error_before_completed_event_raises(): pass +_DONE_MARKER: Final = b"data: [DONE]\n\n" +_CREATED_EVENT: Final = _sse_event({"type": "response.created"}) +_IN_PROGRESS_EVENT: Final = _sse_event({"type": "response.in_progress"}) +_PARTIAL_OUTPUT_EVENTS: Final = _COMPLETE_STREAM_EVENTS[:-1] +_PRE_OUTPUT_PREFIXES: Final = [ + pytest.param([], True, id="nothing-yielded"), + pytest.param([_CREATED_EVENT], False, id="created"), + pytest.param([_CREATED_EVENT, _IN_PROGRESS_EVENT], False, id="created-and-in-progress"), +] + + +def _failure_tracking_logging_obj() -> Mock: + logging_obj: Final = _logging_obj_stub() + logging_obj.async_failure_handler = AsyncMock() + return logging_obj + + +def _assert_failure_logged_once(logging_obj: Mock, exception: Exception) -> None: + assert logging_obj.async_failure_handler.await_count == 1 + assert logging_obj.async_failure_handler.await_args.kwargs["exception"] is exception + + +@pytest.mark.asyncio +@pytest.mark.parametrize("prefix, pre_first_chunk", _PRE_OUTPUT_PREFIXES) +@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type) +async def test_transport_error_before_any_output_raises_fallback_error(prefix, pre_first_chunk, trailing_error): + """A connection lost while only lifecycle events (response.created / response.in_progress) + have streamed is fallback-eligible, so it must surface as the MidStreamFallbackError the + router re-routes, carrying the raw transport error and no generated content.""" + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=prefix, logging_obj=logging_obj, trailing_error=trailing_error) + + with pytest.raises(MidStreamFallbackError) as exc_info: + async for _ in iterator: + pass + + assert exc_info.value.original_exception is trailing_error + assert exc_info.value.is_pre_first_chunk is pre_first_chunk + assert exc_info.value.generated_content == "" + _assert_failure_logged_once(logging_obj, trailing_error) + + +@pytest.mark.asyncio +async def test_transport_error_after_output_started_is_not_fallback_eligible(): + logging_obj: Final = _failure_tracking_logging_obj() + trailing_error: Final = httpx.ReadError("Response payload is not completed") + iterator: Final = _make_iterator( + sse_events=_PARTIAL_OUTPUT_EVENTS, logging_obj=logging_obj, trailing_error=trailing_error + ) + + with pytest.raises(httpx.ReadError) as exc_info: + async for _ in iterator: + pass + + assert exc_info.value is trailing_error + _assert_failure_logged_once(logging_obj, trailing_error) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +async def test_stream_ending_after_partial_output_without_terminal_event_raises(trailer): + """A clean EOF or `[DONE]` after output text but with no response.completed / + response.incomplete / response.failed is a truncated answer: the partial events still + reach the caller, then an explicit error follows instead of a normal end of stream.""" + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=[*_PARTIAL_OUTPUT_EVENTS, *trailer], logging_obj=logging_obj) + + created: Final = await iterator.__anext__() + delta: Final = await iterator.__anext__() + with pytest.raises(litellm.APIConnectionError) as exc_info: + await iterator.__anext__() + + assert (created.type, delta.type) == ("response.created", "response.output_text.delta") + assert not isinstance(exc_info.value, MidStreamFallbackError) + assert exc_info.value.llm_provider == "openai" + _assert_failure_logged_once(logging_obj, exc_info.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("prefix, pre_first_chunk", _PRE_OUTPUT_PREFIXES) +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +async def test_stream_ending_before_any_output_raises_fallback_error(prefix, pre_first_chunk, trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=[*prefix, *trailer], logging_obj=logging_obj) + + with pytest.raises(MidStreamFallbackError) as exc_info: + async for _ in iterator: + pass + + assert isinstance(exc_info.value.original_exception, litellm.APIConnectionError) + assert exc_info.value.is_pre_first_chunk is pre_first_chunk + assert exc_info.value.generated_content == "" + _assert_failure_logged_once(logging_obj, exc_info.value.original_exception) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +async def test_complete_stream_still_ends_normally(trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=[*_COMPLETE_STREAM_EVENTS, *trailer], logging_obj=logging_obj) + + seen: Final = [event.type async for event in iterator] + + assert seen[-1] == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.async_failure_handler.await_count == 0 + + +@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type) +def test_sync_transport_error_before_any_output_raises_fallback_error(trailing_error): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator( + sse_events=[_CREATED_EVENT, _IN_PROGRESS_EVENT], + logging_obj=logging_obj, + trailing_error=trailing_error, + ) + + with pytest.raises(MidStreamFallbackError) as exc_info: + for _ in iterator: + pass + + assert exc_info.value.original_exception is trailing_error + assert exc_info.value.is_pre_first_chunk is False + assert exc_info.value.generated_content == "" + _assert_failure_logged_once(logging_obj, trailing_error) + + +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +def test_sync_stream_ending_after_partial_output_without_terminal_event_raises(trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator(sse_events=[*_PARTIAL_OUTPUT_EVENTS, *trailer], logging_obj=logging_obj) + + created: Final = next(iterator) + delta: Final = next(iterator) + with pytest.raises(litellm.APIConnectionError) as exc_info: + next(iterator) + + assert (created.type, delta.type) == ("response.created", "response.output_text.delta") + assert not isinstance(exc_info.value, MidStreamFallbackError) + _assert_failure_logged_once(logging_obj, exc_info.value) + + +def test_sync_stream_ending_before_any_output_raises_fallback_error(): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator(sse_events=[_CREATED_EVENT], logging_obj=logging_obj) + + with pytest.raises(MidStreamFallbackError) as exc_info: + for _ in iterator: + pass + + assert isinstance(exc_info.value.original_exception, litellm.APIConnectionError) + assert exc_info.value.is_pre_first_chunk is False + _assert_failure_logged_once(logging_obj, exc_info.value.original_exception) + + +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +def test_sync_complete_stream_still_ends_normally(trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator(sse_events=[*_COMPLETE_STREAM_EVENTS, *trailer], logging_obj=logging_obj) + + seen: Final = [event.type for event in iterator] + + assert seen[-1] == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.async_failure_handler.await_count == 0 + + def test_stream_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch): """ Regression test for LIT-6184 on the /v1/responses streaming surface: the diff --git a/tests/test_litellm/responses/test_streaming_iterator_error_events.py b/tests/unit/responses/test_streaming_iterator_error_events.py similarity index 100% rename from tests/test_litellm/responses/test_streaming_iterator_error_events.py rename to tests/unit/responses/test_streaming_iterator_error_events.py diff --git a/tests/test_litellm/responses/test_text_format_conversion.py b/tests/unit/responses/test_text_format_conversion.py similarity index 100% rename from tests/test_litellm/responses/test_text_format_conversion.py rename to tests/unit/responses/test_text_format_conversion.py diff --git a/tests/unit/router_strategy/adaptive_router/__init__.py b/tests/unit/router_strategy/adaptive_router/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/router_strategy/adaptive_router/fixtures/__init__.py b/tests/unit/router_strategy/adaptive_router/fixtures/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/clean_no_signals.json b/tests/unit/router_strategy/adaptive_router/fixtures/clean_no_signals.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/clean_no_signals.json rename to tests/unit/router_strategy/adaptive_router/fixtures/clean_no_signals.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/clean_satisfaction.json b/tests/unit/router_strategy/adaptive_router/fixtures/clean_satisfaction.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/clean_satisfaction.json rename to tests/unit/router_strategy/adaptive_router/fixtures/clean_satisfaction.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/disengagement_giveup.json b/tests/unit/router_strategy/adaptive_router/fixtures/disengagement_giveup.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/disengagement_giveup.json rename to tests/unit/router_strategy/adaptive_router/fixtures/disengagement_giveup.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/exhaustion_429.json b/tests/unit/router_strategy/adaptive_router/fixtures/exhaustion_429.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/exhaustion_429.json rename to tests/unit/router_strategy/adaptive_router/fixtures/exhaustion_429.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/exhaustion_context_overflow.json b/tests/unit/router_strategy/adaptive_router/fixtures/exhaustion_context_overflow.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/exhaustion_context_overflow.json rename to tests/unit/router_strategy/adaptive_router/fixtures/exhaustion_context_overflow.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/failure_tool_error.json b/tests/unit/router_strategy/adaptive_router/fixtures/failure_tool_error.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/failure_tool_error.json rename to tests/unit/router_strategy/adaptive_router/fixtures/failure_tool_error.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/loop_same_tool.json b/tests/unit/router_strategy/adaptive_router/fixtures/loop_same_tool.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/loop_same_tool.json rename to tests/unit/router_strategy/adaptive_router/fixtures/loop_same_tool.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/misalignment_rephrase.json b/tests/unit/router_strategy/adaptive_router/fixtures/misalignment_rephrase.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/misalignment_rephrase.json rename to tests/unit/router_strategy/adaptive_router/fixtures/misalignment_rephrase.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/mixed_failure_then_satisfaction.json b/tests/unit/router_strategy/adaptive_router/fixtures/mixed_failure_then_satisfaction.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/mixed_failure_then_satisfaction.json rename to tests/unit/router_strategy/adaptive_router/fixtures/mixed_failure_then_satisfaction.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/fixtures/stagnation_repeat.json b/tests/unit/router_strategy/adaptive_router/fixtures/stagnation_repeat.json similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/fixtures/stagnation_repeat.json rename to tests/unit/router_strategy/adaptive_router/fixtures/stagnation_repeat.json diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py b/tests/unit/router_strategy/adaptive_router/test_adaptive_router.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py rename to tests/unit/router_strategy/adaptive_router/test_adaptive_router.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_async_pre_routing.py b/tests/unit/router_strategy/adaptive_router/test_async_pre_routing.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_async_pre_routing.py rename to tests/unit/router_strategy/adaptive_router/test_async_pre_routing.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_bandit.py b/tests/unit/router_strategy/adaptive_router/test_bandit.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_bandit.py rename to tests/unit/router_strategy/adaptive_router/test_bandit.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_classifier.py b/tests/unit/router_strategy/adaptive_router/test_classifier.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_classifier.py rename to tests/unit/router_strategy/adaptive_router/test_classifier.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_config.py b/tests/unit/router_strategy/adaptive_router/test_config.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_config.py rename to tests/unit/router_strategy/adaptive_router/test_config.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py b/tests/unit/router_strategy/adaptive_router/test_e2e_adaptive_router.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py rename to tests/unit/router_strategy/adaptive_router/test_e2e_adaptive_router.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_hooks.py b/tests/unit/router_strategy/adaptive_router/test_hooks.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_hooks.py rename to tests/unit/router_strategy/adaptive_router/test_hooks.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_router_dispatch.py b/tests/unit/router_strategy/adaptive_router/test_router_dispatch.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_router_dispatch.py rename to tests/unit/router_strategy/adaptive_router/test_router_dispatch.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_signals.py b/tests/unit/router_strategy/adaptive_router/test_signals.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_signals.py rename to tests/unit/router_strategy/adaptive_router/test_signals.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py b/tests/unit/router_strategy/adaptive_router/test_state_endpoint.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_state_endpoint.py rename to tests/unit/router_strategy/adaptive_router/test_state_endpoint.py diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_update_queue.py b/tests/unit/router_strategy/adaptive_router/test_update_queue.py similarity index 100% rename from tests/test_litellm/router_strategy/adaptive_router/test_update_queue.py rename to tests/unit/router_strategy/adaptive_router/test_update_queue.py diff --git a/tests/test_litellm/router_strategy/complexity_router/test_context_compaction.py b/tests/unit/router_strategy/complexity_router/test_context_compaction.py similarity index 100% rename from tests/test_litellm/router_strategy/complexity_router/test_context_compaction.py rename to tests/unit/router_strategy/complexity_router/test_context_compaction.py diff --git a/tests/test_litellm/router_strategy/test_auto_router.py b/tests/unit/router_strategy/test_auto_router.py similarity index 100% rename from tests/test_litellm/router_strategy/test_auto_router.py rename to tests/unit/router_strategy/test_auto_router.py diff --git a/tests/test_litellm/router_strategy/test_base_routing_strategy.py b/tests/unit/router_strategy/test_base_routing_strategy.py similarity index 100% rename from tests/test_litellm/router_strategy/test_base_routing_strategy.py rename to tests/unit/router_strategy/test_base_routing_strategy.py diff --git a/tests/unit/router_strategy/test_budget_limiter.py b/tests/unit/router_strategy/test_budget_limiter.py new file mode 100644 index 00000000000..62de1586fdd --- /dev/null +++ b/tests/unit/router_strategy/test_budget_limiter.py @@ -0,0 +1,137 @@ +""" +Spend tracking in RouterBudgetLimiting.async_log_success_event. + +Only chat completions puts custom_llm_provider into litellm_params. The responses, +anthropic_messages, embedding and rerank surfaces leave it unset, which used to make +the callback raise before any spend was recorded, so those budgets never moved. +""" + +from typing import Final + +import pytest + +from litellm.caching.caching import DualCache +from litellm.router_strategy.budget_limiter import RouterBudgetLimiting + + +@pytest.fixture +def disable_budget_sync(monkeypatch): + async def noop(*args, **kwargs): + return None + + monkeypatch.setattr( + "litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis", + noop, + ) + + +def _success_kwargs( + *, + provider_in_litellm_params: str | None, + provider_in_payload: str | None, + call_type: str = "aresponses", + response_cost: float = 0.25, + model_id: str = "deployment-1", +) -> dict[str, object]: + provider_params: Final[dict[str, str]] = ( + {} if provider_in_litellm_params is None else {"custom_llm_provider": provider_in_litellm_params} + ) + litellm_params: Final[dict[str, str]] = {"model": "openai/gpt-4o", **provider_params} + + return { + "call_type": call_type, + "litellm_params": litellm_params, + "standard_logging_object": { + "response_cost": response_cost, + "model_id": model_id, + "custom_llm_provider": provider_in_payload, + }, + } + + +async def _log_success(limiter: RouterBudgetLimiting, kwargs: dict[str, object]) -> None: + await limiter.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=None, end_time=None) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("call_type", ["aresponses", "anthropic_messages", "aembedding", "arerank"]) +async def test_provider_spend_tracked_when_litellm_params_omits_provider(disable_budget_sync, call_type): + """Non-chat surfaces carry the provider only on the standard logging payload.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": {"budget_limit": 10.0, "time_period": "1d"}}, + ) + + await _log_success( + limiter, + _success_kwargs( + provider_in_litellm_params=None, + provider_in_payload="openai", + call_type=call_type, + ), + ) + + assert await limiter.dual_cache.async_get_cache("provider_spend:openai:1d") == 0.25 + + +@pytest.mark.asyncio +async def test_chat_completions_spend_still_tracked(disable_budget_sync): + """Chat completions fills in both sources and must keep accumulating.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": {"budget_limit": 10.0, "time_period": "1d"}}, + ) + + await _log_success( + limiter, + _success_kwargs( + provider_in_litellm_params="openai", + provider_in_payload="openai", + call_type="acompletion", + ), + ) + + assert await limiter.dual_cache.async_get_cache("provider_spend:openai:1d") == 0.25 + + +@pytest.mark.asyncio +async def test_budget_of_other_provider_is_untouched(disable_budget_sync): + """A provider without its own budget must not bleed into a configured one.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": {"budget_limit": 10.0, "time_period": "1d"}}, + ) + + await _log_success( + limiter, + _success_kwargs(provider_in_litellm_params=None, provider_in_payload="anthropic"), + ) + + assert await limiter.dual_cache.async_get_cache("provider_spend:openai:1d") in (None, 0.0) + + +@pytest.mark.asyncio +async def test_deployment_budget_tracked_when_provider_is_unresolvable(disable_budget_sync): + """An unresolvable provider must not abort the deployment and tag budgets that follow it.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config=None, + model_list=[ + { + "model_name": "some-model", + "litellm_params": { + "model": "openai/gpt-4o", + "max_budget": 10.0, + "budget_duration": "1d", + }, + "model_info": {"id": "deployment-1"}, + } + ], + ) + + await _log_success( + limiter, + _success_kwargs(provider_in_litellm_params=None, provider_in_payload=None), + ) + + assert await limiter.dual_cache.async_get_cache("deployment_spend:deployment-1:1d") == 0.25 diff --git a/tests/test_litellm/router_strategy/test_budget_limiter_hotpath.py b/tests/unit/router_strategy/test_budget_limiter_hotpath.py similarity index 100% rename from tests/test_litellm/router_strategy/test_budget_limiter_hotpath.py rename to tests/unit/router_strategy/test_budget_limiter_hotpath.py diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/unit/router_strategy/test_complexity_router.py similarity index 99% rename from tests/test_litellm/router_strategy/test_complexity_router.py rename to tests/unit/router_strategy/test_complexity_router.py index 4e146b59b61..3401a335b2f 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/unit/router_strategy/test_complexity_router.py @@ -2679,6 +2679,25 @@ def _llm_response(content: str, response_cost: float | None = None): return response +_REPLY_SHAPES: Final = ("fenced", "fenced-with-language", "prose-before", "prose-after", "fenced-then-prose") + + +def _wrapped_reply(shape: str, verdict: str) -> str: + match shape: + case "fenced": + return f" ```\n{verdict}\n``` " + case "fenced-with-language": + return f"```json\n{verdict}\n```" + case "prose-before": + return f"Sure {{here}} is the verdict you asked for:\n\n{verdict}" + case "prose-after": + return f"{verdict}\n\nThe efficient solver should handle this {{well}}." + case "fenced-then-prose": + return f"```json\n{verdict}\n```\n\n## Reasoning\n\nThe task is coupled, so the forecasts differ." + case _: + raise AssertionError(shape) + + @pytest.fixture def llm_classifier_config() -> Dict: """Config with an LLM-based classifier wired to a 'haiku-classifier' model.""" @@ -3132,12 +3151,40 @@ class TestCapabilityClassifier: assert outcome.capability_forecast.threshold == pytest.approx(expected_threshold) @pytest.mark.asyncio - async def test_fenced_json_verdict_is_accepted(self, mock_router_instance): - reply = _capability_reply(p_solve=0.8) - mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(f"```json\n{reply}\n```")) + @pytest.mark.parametrize("shape", _REPLY_SHAPES) + async def test_verdict_wrapped_in_fence_or_prose_is_accepted(self, mock_router_instance, shape: str): + reply = _wrapped_reply(shape, _capability_reply(p_solve=0.8)) + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(reply)) outcome = await self._router(mock_router_instance).aclassify("do the task") assert outcome.tier == ComplexityTier.SIMPLE assert outcome.cause == "capability_classifier" + assert outcome.capability_forecast is not None + assert outcome.capability_forecast.p_solve == 0.8 + + @pytest.mark.asyncio + @pytest.mark.parametrize("message_logging_off", (False, True)) + async def test_unparseable_reply_is_logged_with_its_text_unless_message_logging_is_off( + self, mock_router_instance, caplog: pytest.LogCaptureFixture, message_logging_off: bool + ): + reply = "The task text is too {vague} for a forecast, sorry." + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(reply)) + outcome = await self._router(mock_router_instance).aclassify( + "do the task", request_kwargs={"turn_off_message_logging": message_logging_off} + ) + assert outcome.cause == "capability_classifier_fallback" + assert "capability classifier failed (ValidationError)" in caplog.text + assert "classifier verdict rejected (" in caplog.text + assert ("raw reply withheld" in caplog.text) is message_logging_off + assert (reply in caplog.text) is not message_logging_off + + @pytest.mark.asyncio + async def test_call_failure_reason_names_the_exception_type( + self, mock_router_instance, caplog: pytest.LogCaptureFixture + ): + mock_router_instance.acompletion = AsyncMock(side_effect=TimeoutError()) + outcome = await self._router(mock_router_instance).aclassify("do the task") + assert outcome.cause == "capability_classifier_fallback" + assert "capability classifier failed (TimeoutError)" in caplog.text @pytest.mark.asyncio async def test_decimal_rounding_does_not_break_inclusive_threshold(self, mock_router_instance): @@ -3954,6 +4001,34 @@ class TestLLMClassifier: assert call_kwargs["model"] == "haiku-classifier" assert call_kwargs["timeout"] == 0.4 + @pytest.mark.asyncio + @pytest.mark.parametrize("shape", _REPLY_SHAPES) + async def test_aclassify_llm_verdict_wrapped_in_fence_or_prose_still_decides_the_tier( + self, llm_complexity_router, mock_router_instance, shape: str + ): + reply = _wrapped_reply(shape, '{"tier": "COMPLEX"}') + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(reply)) + outcome = await llm_complexity_router.aclassify("hi") + assert outcome.tier == ComplexityTier.COMPLEX + assert outcome.cause == "llm_classifier" + assert "llm-classifier:COMPLEX" in outcome.signals + + @pytest.mark.asyncio + @pytest.mark.parametrize("message_logging_off", (False, True)) + async def test_aclassify_llm_unparseable_reply_is_logged_with_its_text_unless_message_logging_is_off( + self, llm_complexity_router, mock_router_instance, caplog: pytest.LogCaptureFixture, message_logging_off: bool + ): + reply = "I would call this COMPLEX, the {tier} field is implied." + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(reply)) + outcome = await llm_complexity_router.aclassify( + "hi", request_kwargs={"turn_off_message_logging": message_logging_off} + ) + assert outcome.cause != "llm_classifier" + assert "LLM classifier failed (ValidationError)" in caplog.text + assert "classifier verdict rejected (" in caplog.text + assert ("raw reply withheld" in caplog.text) is message_logging_off + assert (reply in caplog.text) is not message_logging_off + @pytest.mark.asyncio async def test_aclassify_llm_success_captures_classifier_cost(self, llm_complexity_router, mock_router_instance): """The classifier call is billed, so its cost must ride the outcome. diff --git a/tests/test_litellm/router_strategy/test_complexity_tier_predictor.py b/tests/unit/router_strategy/test_complexity_tier_predictor.py similarity index 100% rename from tests/test_litellm/router_strategy/test_complexity_tier_predictor.py rename to tests/unit/router_strategy/test_complexity_tier_predictor.py diff --git a/tests/test_litellm/router_strategy/test_fuse_presets.py b/tests/unit/router_strategy/test_fuse_presets.py similarity index 100% rename from tests/test_litellm/router_strategy/test_fuse_presets.py rename to tests/unit/router_strategy/test_fuse_presets.py diff --git a/tests/test_litellm/router_strategy/test_lar1_routing.py b/tests/unit/router_strategy/test_lar1_routing.py similarity index 100% rename from tests/test_litellm/router_strategy/test_lar1_routing.py rename to tests/unit/router_strategy/test_lar1_routing.py diff --git a/tests/test_litellm/router_strategy/test_least_busy.py b/tests/unit/router_strategy/test_least_busy.py similarity index 100% rename from tests/test_litellm/router_strategy/test_least_busy.py rename to tests/unit/router_strategy/test_least_busy.py diff --git a/tests/test_litellm/router_strategy/test_litellm_encoder.py b/tests/unit/router_strategy/test_litellm_encoder.py similarity index 100% rename from tests/test_litellm/router_strategy/test_litellm_encoder.py rename to tests/unit/router_strategy/test_litellm_encoder.py diff --git a/tests/test_litellm/router_strategy/test_llm_v2.py b/tests/unit/router_strategy/test_llm_v2.py similarity index 84% rename from tests/test_litellm/router_strategy/test_llm_v2.py rename to tests/unit/router_strategy/test_llm_v2.py index 6fb6df3265d..3fd2e8808e7 100644 --- a/tests/test_litellm/router_strategy/test_llm_v2.py +++ b/tests/unit/router_strategy/test_llm_v2.py @@ -66,6 +66,25 @@ def _response(content: str) -> ModelResponse: return response +_REPLY_SHAPES: Final = ("fenced", "fenced-with-language", "prose-before", "prose-after", "fenced-then-prose") + + +def _wrapped_reply(shape: str, verdict: str) -> str: + match shape: + case "fenced": + return f" ```\n{verdict}\n``` " + case "fenced-with-language": + return f"```json\n{verdict}\n```" + case "prose-before": + return f"Sure {{here}} is the verdict you asked for:\n\n{verdict}" + case "prose-after": + return f"{verdict}\n\nThe efficient solver should handle this {{well}}." + case "fenced-then-prose": + return f"```json\n{verdict}\n```\n\n## Reasoning\n\nThe task is coupled, so the forecasts differ." + case _: + raise AssertionError(shape) + + def _router(content: str, config: ComplexityRouterConfig | None = None) -> tuple[ComplexityRouter, MagicMock]: client: Final = MagicMock(spec=Router) client.acompletion = AsyncMock(return_value=_response(content)) @@ -334,13 +353,12 @@ async def test_json_object_mode_supplies_schema_in_prompt() -> None: @pytest.mark.asyncio @pytest.mark.parametrize("mode", ("json_schema", "json_object")) -@pytest.mark.parametrize("fence", ("```json", "```")) -async def test_fenced_forecast_routes_by_validated_probabilities(mode: str, fence: str) -> None: +@pytest.mark.parametrize("shape", _REPLY_SHAPES) +async def test_wrapped_forecast_routes_by_validated_probabilities(mode: str, shape: str) -> None: base: Final = _config().llm_v2_config assert base is not None config: Final = _config(llm_v2_config={**base.model_dump(), "response_format": mode}) - content: Final = f" {fence}\n{_verdict().model_dump_json()}\n``` " - router, client = _router(content, config) + router, client = _router(_wrapped_reply(shape, _verdict().model_dump_json()), config) result: Final = await router.async_pre_routing_hook( model="v2-router", messages=[{"role": "user", "content": "Fix nested behavior"}], request_kwargs={} ) @@ -454,6 +472,93 @@ async def test_provider_failure_redacts_prompt_text_from_warning(caplog: pytest. assert "private task text" not in caplog.text +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ("crux", "likely_failure")) +async def test_long_verdict_explanations_still_route_by_validated_probabilities(field: str) -> None: + explanation: Final = "The solver must keep the nested retry behavior intact while it edits. " * 12 + assert len(explanation) > 512 + verdict: Final = _verdict().model_dump() + if field == "crux": + content: Final = json.dumps({**verdict, "crux": explanation}) + else: + forecasts: Final = {**verdict["forecasts"], "efficient": {**verdict["forecasts"]["efficient"], field: explanation}} + content = json.dumps({**verdict, "forecasts": forecasts}) + router, _ = _router(content) + outcome: Final = await router.aclassify("Fix nested behavior") + assert outcome.cause == "llm_v2_classifier" + assert outcome.llm_v2_forecast is not None + assert outcome.llm_v2_forecast.use_efficient + + +@pytest.mark.parametrize("field", ("crux", "likely_failure")) +def test_blank_verdict_explanations_are_still_rejected(field: str) -> None: + verdict: Final = _verdict().model_dump() + blank: Final = ( + {**verdict, "crux": " "} + if field == "crux" + else {**verdict, "forecasts": {**verdict["forecasts"], "capable": {"likely_failure": " ", "p_solve": 0.5}}} + ) + with pytest.raises(ValidationError): + LLMV2Verdict.model_validate(blank) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("message_logging_off", (False, True)) +async def test_unparseable_reply_is_logged_with_its_text_unless_message_logging_is_off( + caplog: pytest.LogCaptureFixture, message_logging_off: bool +) -> None: + reply: Final = "I cannot forecast this one, the task text is too {vague} to score." + router, _ = _router(reply) + outcome: Final = await router.aclassify("hi", request_kwargs={"turn_off_message_logging": message_logging_off}) + assert outcome.cause == "llm_v2_fallback" + assert "classifier verdict rejected (" in caplog.text + assert "Invalid LLM V2 forecast" in caplog.text + assert ("raw reply withheld" in caplog.text) is message_logging_off + assert (reply in caplog.text) is not message_logging_off + + +_MESSAGE_LOGGING_OPT_OUTS: Final = ( + pytest.param({"turn_off_message_logging": "True"}, False, id="key-logging-settings-string"), + pytest.param({"metadata": {"headers": {"x-litellm-enable-message-redaction": "true"}}}, False, id="redaction-header"), + pytest.param({}, True, id="global-setting"), + pytest.param({"metadata": {"headers": None}}, False, id="undecidable-headers-fail-closed"), +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("request_kwargs", "global_off"), _MESSAGE_LOGGING_OPT_OUTS) +async def test_unparseable_reply_text_is_withheld_under_every_message_logging_opt_out( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + request_kwargs: dict[str, object], + global_off: bool, +) -> None: + monkeypatch.setattr(litellm, "turn_off_message_logging", global_off) + reply: Final = "I cannot forecast this one, the task text is too {vague} to score." + router, _ = _router(reply) + outcome: Final = await router.aclassify("hi", request_kwargs=request_kwargs) + assert outcome.cause == "llm_v2_fallback" + assert "raw reply withheld" in caplog.text + assert reply not in caplog.text + + +_REPLIES_THE_JSON_SCANNER_CANNOT_DECODE: Final = ( + pytest.param('{"a":' * 3000, id="deeply-nested"), + pytest.param('{"capability_p": ' + "9" * 5000 + "}", id="integer-over-the-digit-limit"), +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reply", _REPLIES_THE_JSON_SCANNER_CANNOT_DECODE) +async def test_undecodable_reply_is_rejected_as_an_invalid_forecast( + caplog: pytest.LogCaptureFixture, reply: str +) -> None: + router, _ = _router(reply) + outcome: Final = await router.aclassify("hi", request_kwargs={}) + assert outcome.cause == "llm_v2_fallback" + assert "Invalid LLM V2 forecast" in caplog.text + + def test_response_schema_requires_both_model_forecasts() -> None: with pytest.raises(ValidationError): LLMV2Verdict.model_validate( diff --git a/tests/test_litellm/router_strategy/test_lowest_cost.py b/tests/unit/router_strategy/test_lowest_cost.py similarity index 100% rename from tests/test_litellm/router_strategy/test_lowest_cost.py rename to tests/unit/router_strategy/test_lowest_cost.py diff --git a/tests/test_litellm/router_strategy/test_lowest_latency.py b/tests/unit/router_strategy/test_lowest_latency.py similarity index 100% rename from tests/test_litellm/router_strategy/test_lowest_latency.py rename to tests/unit/router_strategy/test_lowest_latency.py diff --git a/tests/test_litellm/router_strategy/test_lowest_tpm_rpm.py b/tests/unit/router_strategy/test_lowest_tpm_rpm.py similarity index 100% rename from tests/test_litellm/router_strategy/test_lowest_tpm_rpm.py rename to tests/unit/router_strategy/test_lowest_tpm_rpm.py diff --git a/tests/test_litellm/router_strategy/test_quality_router.py b/tests/unit/router_strategy/test_quality_router.py similarity index 100% rename from tests/test_litellm/router_strategy/test_quality_router.py rename to tests/unit/router_strategy/test_quality_router.py diff --git a/tests/test_litellm/router_strategy/test_router_routing_groups.py b/tests/unit/router_strategy/test_router_routing_groups.py similarity index 100% rename from tests/test_litellm/router_strategy/test_router_routing_groups.py rename to tests/unit/router_strategy/test_router_routing_groups.py diff --git a/tests/test_litellm/router_strategy/test_router_routing_plugins.py b/tests/unit/router_strategy/test_router_routing_plugins.py similarity index 100% rename from tests/test_litellm/router_strategy/test_router_routing_plugins.py rename to tests/unit/router_strategy/test_router_routing_plugins.py diff --git a/tests/test_litellm/router_strategy/test_router_tag_regex_routing.py b/tests/unit/router_strategy/test_router_tag_regex_routing.py similarity index 100% rename from tests/test_litellm/router_strategy/test_router_tag_regex_routing.py rename to tests/unit/router_strategy/test_router_tag_regex_routing.py diff --git a/tests/test_litellm/router_strategy/test_router_tag_routing.py b/tests/unit/router_strategy/test_router_tag_routing.py similarity index 100% rename from tests/test_litellm/router_strategy/test_router_tag_routing.py rename to tests/unit/router_strategy/test_router_tag_routing.py diff --git a/tests/test_litellm/router_strategy/test_savings_baseline.py b/tests/unit/router_strategy/test_savings_baseline.py similarity index 100% rename from tests/test_litellm/router_strategy/test_savings_baseline.py rename to tests/unit/router_strategy/test_savings_baseline.py diff --git a/tests/test_litellm/router_strategy/test_simple_shuffle.py b/tests/unit/router_strategy/test_simple_shuffle.py similarity index 100% rename from tests/test_litellm/router_strategy/test_simple_shuffle.py rename to tests/unit/router_strategy/test_simple_shuffle.py diff --git a/tests/test_litellm/router_strategy/test_stall_detector.py b/tests/unit/router_strategy/test_stall_detector.py similarity index 100% rename from tests/test_litellm/router_strategy/test_stall_detector.py rename to tests/unit/router_strategy/test_stall_detector.py diff --git a/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py b/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py index a7006c62438..00462b65bc2 100644 --- a/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py +++ b/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py @@ -565,7 +565,7 @@ async def test_wildcard_route_resolves_underlying_model_minimum(local_model_cost @pytest.mark.asyncio async def test_async_filter_deployments_counts_the_prompt_off_the_event_loop(): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, @@ -589,7 +589,7 @@ async def test_async_filter_deployments_counts_the_prompt_off_the_event_loop(): @pytest.mark.asyncio async def test_async_log_success_event_counts_the_prompt_off_the_event_loop(): from tests.large_text import text - from tests.test_litellm.litellm_core_utils.event_loop_lag import ( + from tests.unit.litellm_core_utils.event_loop_lag import ( assert_loop_stayed_free, timed_with_loop_lags, warm_tokenizer, diff --git a/tests/test_litellm/router_utils/test_access_windows.py b/tests/unit/router_utils/test_access_windows.py similarity index 100% rename from tests/test_litellm/router_utils/test_access_windows.py rename to tests/unit/router_utils/test_access_windows.py diff --git a/tests/test_litellm/router_utils/test_add_retry_fallback_headers.py b/tests/unit/router_utils/test_add_retry_fallback_headers.py similarity index 100% rename from tests/test_litellm/router_utils/test_add_retry_fallback_headers.py rename to tests/unit/router_utils/test_add_retry_fallback_headers.py diff --git a/tests/test_litellm/router_utils/test_auto_router_model_naming.py b/tests/unit/router_utils/test_auto_router_model_naming.py similarity index 100% rename from tests/test_litellm/router_utils/test_auto_router_model_naming.py rename to tests/unit/router_utils/test_auto_router_model_naming.py diff --git a/tests/test_litellm/router_utils/test_auto_router_tuning_baseline.py b/tests/unit/router_utils/test_auto_router_tuning_baseline.py similarity index 100% rename from tests/test_litellm/router_utils/test_auto_router_tuning_baseline.py rename to tests/unit/router_utils/test_auto_router_tuning_baseline.py diff --git a/tests/test_litellm/router_utils/test_client_initalization_utils.py b/tests/unit/router_utils/test_client_initalization_utils.py similarity index 100% rename from tests/test_litellm/router_utils/test_client_initalization_utils.py rename to tests/unit/router_utils/test_client_initalization_utils.py diff --git a/tests/test_litellm/router_utils/test_cooldown_cache.py b/tests/unit/router_utils/test_cooldown_cache.py similarity index 100% rename from tests/test_litellm/router_utils/test_cooldown_cache.py rename to tests/unit/router_utils/test_cooldown_cache.py diff --git a/tests/test_litellm/router_utils/test_cooldown_handlers.py b/tests/unit/router_utils/test_cooldown_handlers.py similarity index 100% rename from tests/test_litellm/router_utils/test_cooldown_handlers.py rename to tests/unit/router_utils/test_cooldown_handlers.py diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/unit/router_utils/test_fallback_event_handlers.py similarity index 100% rename from tests/test_litellm/router_utils/test_fallback_event_handlers.py rename to tests/unit/router_utils/test_fallback_event_handlers.py diff --git a/tests/test_litellm/router_utils/test_get_retry_from_policy.py b/tests/unit/router_utils/test_get_retry_from_policy.py similarity index 100% rename from tests/test_litellm/router_utils/test_get_retry_from_policy.py rename to tests/unit/router_utils/test_get_retry_from_policy.py diff --git a/tests/test_litellm/router_utils/test_health_check_allowed_fails_integration.py b/tests/unit/router_utils/test_health_check_allowed_fails_integration.py similarity index 100% rename from tests/test_litellm/router_utils/test_health_check_allowed_fails_integration.py rename to tests/unit/router_utils/test_health_check_allowed_fails_integration.py diff --git a/tests/test_litellm/router_utils/test_health_state_cache.py b/tests/unit/router_utils/test_health_state_cache.py similarity index 100% rename from tests/test_litellm/router_utils/test_health_state_cache.py rename to tests/unit/router_utils/test_health_state_cache.py diff --git a/tests/test_litellm/router_utils/test_pattern_match_deployments.py b/tests/unit/router_utils/test_pattern_match_deployments.py similarity index 100% rename from tests/test_litellm/router_utils/test_pattern_match_deployments.py rename to tests/unit/router_utils/test_pattern_match_deployments.py diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/unit/router_utils/test_reasoning_effort_capability.py similarity index 100% rename from tests/test_litellm/router_utils/test_reasoning_effort_capability.py rename to tests/unit/router_utils/test_reasoning_effort_capability.py diff --git a/tests/test_litellm/router_utils/test_router_health_check_routing.py b/tests/unit/router_utils/test_router_health_check_routing.py similarity index 100% rename from tests/test_litellm/router_utils/test_router_health_check_routing.py rename to tests/unit/router_utils/test_router_health_check_routing.py diff --git a/tests/test_litellm/router_utils/test_router_interactions_endpoints.py b/tests/unit/router_utils/test_router_interactions_endpoints.py similarity index 100% rename from tests/test_litellm/router_utils/test_router_interactions_endpoints.py rename to tests/unit/router_utils/test_router_interactions_endpoints.py diff --git a/tests/test_litellm/router_utils/test_router_utils_common_utils.py b/tests/unit/router_utils/test_router_utils_common_utils.py similarity index 100% rename from tests/test_litellm/router_utils/test_router_utils_common_utils.py rename to tests/unit/router_utils/test_router_utils_common_utils.py diff --git a/tests/test_litellm/rust_bridge/AGENTS.md b/tests/unit/rust_bridge/AGENTS.md similarity index 100% rename from tests/test_litellm/rust_bridge/AGENTS.md rename to tests/unit/rust_bridge/AGENTS.md diff --git a/tests/unit/rust_bridge/messages/test_route_host.py b/tests/unit/rust_bridge/messages/test_route_host.py index a880cfe3588..1be42e2249d 100644 --- a/tests/unit/rust_bridge/messages/test_route_host.py +++ b/tests/unit/rust_bridge/messages/test_route_host.py @@ -3,6 +3,10 @@ from typing import Final from litellm.rust_bridge.messages.route_host import arguments, response from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest +from dataclasses import astuple +import pytest +import litellm +from litellm.rust_bridge.messages import route_host def test_response_is_a_detached_public_messages_dict() -> None: @@ -40,3 +44,121 @@ def test_arguments_are_the_public_kwargs_view() -> None: ) assert arguments(request) is kwargs + + +pytestmark = pytest.mark.usefixtures("local_model_cost_map") + + +def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None: + monkeypatch.setitem( + litellm.model_cost, + name, + { + "litellm_provider": "anthropic", + "mode": "chat", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + **flags, + }, + ) + + +def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None: + _flag_model( + monkeypatch, + "claude-test-adaptive", + supports_reasoning=True, + supports_adaptive_thinking=True, + supports_output_config=True, + supports_xhigh_reasoning_effort=True, + supports_sampling_params=False, + ) + + capabilities: Final = route_host.model_capabilities("anthropic/claude-test-adaptive", None) + + assert capabilities.supports_adaptive_thinking + assert capabilities.supports_output_config + assert not capabilities.supports_legacy_thinking + assert not capabilities.supports_sampling_params + assert capabilities.effort_tiers.xhigh + assert not capabilities.effort_tiers.max + + +def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None: + capabilities: Final = route_host.model_capabilities("anthropic/not-a-real-model", None) + + assert capabilities.supports_sampling_params + assert not capabilities.supports_reasoning + assert not capabilities.supports_adaptive_thinking + assert not any(astuple(capabilities.effort_tiers)) + + +@pytest.mark.parametrize( + ("global_flag", "kwargs", "expected"), + [ + (False, {}, False), + (True, {}, True), + (False, {"drop_params": "true"}, True), + (False, {"drop_params": "nonsense"}, False), + (False, {"drop_params": False}, False), + ], +) +def test_drop_params_merges_the_global_flag_with_the_request( + monkeypatch: pytest.MonkeyPatch, global_flag: bool, kwargs: dict[str, object], expected: bool +) -> None: + monkeypatch.setattr(litellm, "drop_params", global_flag) + + assert route_host.shaping("anthropic/not-a-real-model", None, kwargs)["drop_params"] is expected + + +@pytest.mark.parametrize( + ("configured", "expected"), + [ + (["tools[*].input_examples", 3, "metadata.user_id"], ("tools[*].input_examples", "metadata.user_id")), + ("tools", ()), + (None, ()), + ], +) +def test_additional_drop_params_keep_only_string_paths(configured: object, expected: tuple[str, ...]) -> None: + shaping: Final = route_host.shaping("anthropic/not-a-real-model", None, {"additional_drop_params": configured}) + + assert shaping["additional_drop_params"] == expected + + +def test_native_request_rejections_map_to_the_public_400() -> None: + from types import MappingProxyType + + from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest + + request: Final = LiteLLMMessagesRequest( + model="anthropic/claude-sonnet-5", + messages=(), + max_tokens=8, + stream=None, + api_key=None, + api_base=None, + custom_llm_provider=None, + kwargs=MappingProxyType({}), + ) + rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5") + rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets + + mapped: Final = route_host.map_failure(rejected, request, "anthropic") + + assert isinstance(mapped, litellm.BadRequestError) + assert mapped.status_code == 400 + assert "does not support top_k=5" in mapped.message + assert mapped.model == "claude-sonnet-5" + assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError) + + +def test_stream_hidden_params_projects_upstream_headers_the_way_the_python_handler_does() -> None: + hidden: Final = route_host.stream_hidden_params( + (("request-id", "req_upstream_123"), ("x-ratelimit-remaining-requests", "41")) + ) + + additional: Final = hidden["additional_headers"] + assert isinstance(additional, dict) + assert additional["llm_provider-request-id"] == "req_upstream_123" + assert additional["x-ratelimit-remaining-requests"] == "41" + assert "request-id" not in additional diff --git a/tests/test_litellm/rust_bridge/messages/test_secrets.py b/tests/unit/rust_bridge/messages/test_secrets.py similarity index 100% rename from tests/test_litellm/rust_bridge/messages/test_secrets.py rename to tests/unit/rust_bridge/messages/test_secrets.py diff --git a/tests/test_litellm/rust_bridge/native_route_wheel_test.py b/tests/unit/rust_bridge/native_route_wheel_test.py similarity index 100% rename from tests/test_litellm/rust_bridge/native_route_wheel_test.py rename to tests/unit/rust_bridge/native_route_wheel_test.py diff --git a/tests/test_litellm/rust_bridge/ocr/test_secrets.py b/tests/unit/rust_bridge/ocr/test_secrets.py similarity index 100% rename from tests/test_litellm/rust_bridge/ocr/test_secrets.py rename to tests/unit/rust_bridge/ocr/test_secrets.py diff --git a/tests/test_litellm/rust_bridge/stubtest.ini b/tests/unit/rust_bridge/stubtest.ini similarity index 100% rename from tests/test_litellm/rust_bridge/stubtest.ini rename to tests/unit/rust_bridge/stubtest.ini diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/unit/rust_bridge/test_bindings.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_bindings.py rename to tests/unit/rust_bridge/test_bindings.py diff --git a/tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py b/tests/unit/rust_bridge/test_callbacks_legacy_python.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py rename to tests/unit/rust_bridge/test_callbacks_legacy_python.py diff --git a/tests/test_litellm/rust_bridge/test_catalog.py b/tests/unit/rust_bridge/test_catalog.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_catalog.py rename to tests/unit/rust_bridge/test_catalog.py diff --git a/tests/test_litellm/rust_bridge/test_configuration.py b/tests/unit/rust_bridge/test_configuration.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_configuration.py rename to tests/unit/rust_bridge/test_configuration.py diff --git a/tests/test_litellm/rust_bridge/test_dispatch.py b/tests/unit/rust_bridge/test_dispatch.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_dispatch.py rename to tests/unit/rust_bridge/test_dispatch.py diff --git a/tests/test_litellm/rust_bridge/test_failures.py b/tests/unit/rust_bridge/test_failures.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_failures.py rename to tests/unit/rust_bridge/test_failures.py diff --git a/tests/test_litellm/rust_bridge/test_fork_guard.py b/tests/unit/rust_bridge/test_fork_guard.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_fork_guard.py rename to tests/unit/rust_bridge/test_fork_guard.py diff --git a/tests/test_litellm/rust_bridge/test_lifecycle.py b/tests/unit/rust_bridge/test_lifecycle.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_lifecycle.py rename to tests/unit/rust_bridge/test_lifecycle.py diff --git a/tests/test_litellm/rust_bridge/test_logger.py b/tests/unit/rust_bridge/test_logger.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_logger.py rename to tests/unit/rust_bridge/test_logger.py diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/unit/rust_bridge/test_runtime.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_runtime.py rename to tests/unit/rust_bridge/test_runtime.py diff --git a/tests/test_litellm/rust_bridge/test_secret_manager.py b/tests/unit/rust_bridge/test_secret_manager.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_secret_manager.py rename to tests/unit/rust_bridge/test_secret_manager.py diff --git a/tests/test_litellm/rust_bridge/test_settings.py b/tests/unit/rust_bridge/test_settings.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_settings.py rename to tests/unit/rust_bridge/test_settings.py diff --git a/tests/test_litellm/rust_bridge/test_token_counter.py b/tests/unit/rust_bridge/test_token_counter.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_token_counter.py rename to tests/unit/rust_bridge/test_token_counter.py diff --git a/tests/test_litellm/rust_bridge/test_tokenizer.py b/tests/unit/rust_bridge/test_tokenizer.py similarity index 95% rename from tests/test_litellm/rust_bridge/test_tokenizer.py rename to tests/unit/rust_bridge/test_tokenizer.py index 188aa81093f..c5093cdb0ce 100644 --- a/tests/test_litellm/rust_bridge/test_tokenizer.py +++ b/tests/unit/rust_bridge/test_tokenizer.py @@ -7,7 +7,7 @@ from tokenizers import Tokenizer from litellm.litellm_core_utils.tokenizer import HuggingFaceTokenizer, OpenAIEncoding from litellm.rust_bridge import tokenizer from litellm.utils import claude_json_str -from tests.test_litellm.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON +from tests.unit.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON TEXTS: Final = ("hello <|endoftext|> world", "café 漢字 🙂", " def f():\n return 1\n", "hello again") diff --git a/tests/test_litellm/rust_bridge/test_verify_linux_native_wheel.py b/tests/unit/rust_bridge/test_verify_linux_native_wheel.py similarity index 100% rename from tests/test_litellm/rust_bridge/test_verify_linux_native_wheel.py rename to tests/unit/rust_bridge/test_verify_linux_native_wheel.py diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 99dea6366f9..62ef9f11c2e 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -4581,6 +4581,55 @@ def test_every_openai_entry_with_a_long_context_rate_and_a_batch_rate_declares_t assert undeclared == [] +@pytest.mark.parametrize("prefix", _BATCH_RATE_PREFIXES) +def test_every_xai_entry_with_a_long_context_rate_and_a_batch_rate_declares_the_batch_tier( + _local_model_cost_map: None, prefix: str +) -> None: + undeclared: Final = [ + name + for name, entry in litellm.model_cost.items() + if isinstance(entry, dict) + and entry.get("litellm_provider") == "xai" + and entry.get(f"{prefix}_above_200k_tokens") is not None + and entry.get(f"{prefix}_batches") is not None + and entry.get(f"{prefix}_above_200k_tokens_batches") is None + ] + + assert undeclared == [] + + +_XAI_TIERED_BATCH_MODEL: Final = "xai/grok-4.3" + + +def test_xai_batch_tier_discounts_the_long_context_rate_like_the_flat_batch_rate(_local_model_cost_map: None) -> None: + info: Final = litellm.get_model_info(_XAI_TIERED_BATCH_MODEL, custom_llm_provider="xai") + flat_discount: Final = info["input_cost_per_token_batches"] / info["input_cost_per_token"] + + for prefix in ("input_cost_per_token", "output_cost_per_token", "cache_read_input_token_cost"): + tier_discount = info[f"{prefix}_above_200k_tokens_batches"] / info[f"{prefix}_above_200k_tokens"] + assert tier_discount == pytest.approx(flat_discount) + assert info[f"{prefix}_above_200k_tokens_batches"] < info[f"{prefix}_above_200k_tokens"] + + +@pytest.mark.parametrize( + ("prompt_tokens", "tier"), [(200_000, "_above_200k_tokens_batches"), (199_999, "_batches")] +) +def test_xai_batch_cost_calculator_bills_the_200k_batch_tier_inclusively( + _local_model_cost_map: None, prompt_tokens: int, tier: str +) -> None: + from litellm.cost_calculator import batch_cost_calculator + + info: Final = litellm.get_model_info(_XAI_TIERED_BATCH_MODEL, custom_llm_provider="xai") + usage: Final = Usage(prompt_tokens=prompt_tokens, completion_tokens=64, total_tokens=prompt_tokens + 64) + + prompt_cost, completion_cost_value = batch_cost_calculator( + usage=usage, model=_XAI_TIERED_BATCH_MODEL, custom_llm_provider="xai" + ) + + assert prompt_cost == pytest.approx(prompt_tokens * info[f"input_cost_per_token{tier}"]) + assert completion_cost_value == pytest.approx(64 * info[f"output_cost_per_token{tier}"]) + + def test_batch_cost_calculator_ignores_malformed_batch_tier_keys(): from litellm.cost_calculator import batch_cost_calculator diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index effc038f85b..c06216e4f4e 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -626,6 +626,29 @@ def test_return_raw_request_does_not_call_provider(respx_mock: respx.MockRouter) ] +def test_return_raw_request_ignores_turn_off_message_logging( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + from litellm.types.utils import CallTypes + from litellm.utils import return_raw_request + + model: Final = "gpt-4o" + messages: Final = [{"role": "user", "content": "PRIVATE-PHRASE"}] + route: Final = respx_mock.post("https://api.openai.com/v1/chat/completions").mock( + return_value=_mocked_openai_chat_response(model) + ) + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + + request: Final = return_raw_request( + endpoint=CallTypes.completion, + kwargs={"model": model, "messages": messages}, + ) + + assert route.call_count == 0 + assert request.get("error") is None + assert request["raw_request_body"]["messages"] == messages + + def test_completion_forwards_verbosity_in_raw_request(respx_mock: respx.MockRouter): """Regression test: completion() must forward the verbosity param to the provider request body.""" from litellm.types.utils import CallTypes diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 80131534183..3393c2f0d3c 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -4486,6 +4486,107 @@ async def test_aresponses_streaming_iterator_pre_first_chunk_skips_continuation( assert fbk["input"] == "Hello" # original input, no continuation messages +def _make_native_responses_iterator(*, sse_payloads: tuple[dict[str, str], ...], trailing_error: Exception | None): + """A real ResponsesAPIStreamingIterator over canned SSE bytes, so the router test covers the + iterator's own transport-error classification instead of a hand-built MidStreamFallbackError.""" + from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig + from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator + + async def aiter_bytes(): + for payload in sse_payloads: + yield f"data: {json.dumps(payload)}\n\n".encode() + if trailing_error is not None: + raise trailing_error + + def transform(model, parsed_chunk, logging_obj): + return MagicMock(type=parsed_chunk["type"]) + + response: Final = MagicMock() + response.headers = {} + response.aiter_bytes = aiter_bytes + config: Final = MagicMock(spec=BaseResponsesAPIConfig) + config.transform_streaming_response.side_effect = transform + logging_obj: Final = MagicMock(spec=LiteLLMLogging) + logging_obj.completion_start_time = None + logging_obj.model_call_details = {"litellm_params": {}} + return ResponsesAPIStreamingIterator( + response=response, + model="gpt-4", + responses_api_provider_config=config, + logging_obj=logging_obj, + litellm_metadata={}, + custom_llm_provider="openai", + ) + + +_RESPONSES_LIFECYCLE_PAYLOADS: Final = ({"type": "response.created"}, {"type": "response.in_progress"}) + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_falls_back_on_transport_drop_before_output(): + """A connection lost after response.created but before any output item is re-routed to the + fallback with the original input, the same as a provider error event would be.""" + router: Final = _make_router_with_fallback() + src: Final = _make_native_responses_iterator( + sse_payloads=_RESPONSES_LIFECYCLE_PAYLOADS, + trailing_error=httpx.ReadError("Response payload is not completed"), + ) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=_AsyncList([MagicMock(type="response.completed")]), + ) as mock_fallback_utils: + wrapped: Final = await router._aresponses_streaming_iterator( + response=src, + initial_kwargs={ + "model": "gpt-4", + "stream": True, + "input": "Hello", + "original_generic_function": litellm.aresponses, + }, + ) + seen: Final = [chunk.type async for chunk in wrapped] + + assert seen == ["response.created", "response.in_progress", "response.completed"] + assert isinstance(mock_fallback_utils.call_args.kwargs["e"], MidStreamFallbackError) + assert mock_fallback_utils.call_args.kwargs["kwargs"]["input"] == "Hello" + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_surfaces_transport_drop_when_no_fallback_lands(): + transport_error: Final = httpx.ReadError("Response payload is not completed") + router: Final = _make_router_with_fallback() + src: Final = _make_native_responses_iterator( + sse_payloads=_RESPONSES_LIFECYCLE_PAYLOADS, trailing_error=transport_error + ) + + async def reraise_trigger(**kwargs): + raise kwargs["e"] + + with patch.object( + router, "async_function_with_fallbacks_common_utils", new=AsyncMock(side_effect=reraise_trigger) + ) as mock_fallback_utils: + wrapped: Final = await router._aresponses_streaming_iterator( + response=src, + initial_kwargs={ + "model": "gpt-4", + "stream": True, + "input": "Hello", + "original_generic_function": litellm.aresponses, + }, + ) + with pytest.raises(httpx.ReadError) as exc_info: + async for _ in wrapped: + pass + + assert exc_info.value is transport_error + assert mock_fallback_utils.await_count == 1 + trigger: Final = mock_fallback_utils.await_args.kwargs["e"] + assert isinstance(trigger, MidStreamFallbackError) + assert trigger.original_exception is transport_error + + @pytest.mark.asyncio async def test_aresponses_streaming_iterator_partial_content_injects_continuation(): """Mid-stream error: input is rewritten to include user prompt + @@ -6090,6 +6191,32 @@ def test_update_kwargs_with_deployment_passthrough_router_stream_timeout_sources assert _passthrough_timeout(default_router, default_router.model_list[0], stream=False) == 120.0 +def test_update_kwargs_with_deployment_passthrough_honors_global_request_timeout(monkeypatch: pytest.MonkeyPatch): + """litellm_settings.request_timeout must bound the native responses route when neither the + deployment nor the router carries a timeout, while a deployment timeout keeps winning.""" + monkeypatch.setattr("litellm.request_timeout", 44.0, raising=False) + monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "responses-global-timeout", + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "fake-key"}, + }, + { + "model_name": "responses-deployment-timeout", + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "fake-key", "timeout": 3}, + }, + ], + ) + global_only, per_deployment = router.model_list + + with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}): + assert _passthrough_timeout(router, global_only, stream=True) == 44.0 + assert _passthrough_timeout(router, global_only, stream=False) == 44.0 + assert _passthrough_timeout(router, per_deployment, stream=True) == 3.0 + assert _passthrough_timeout(router, per_deployment, stream=False) == 3.0 + + @pytest.mark.asyncio async def test_router_acompletion_with_unknown_model_and_default_fallback(): """ @@ -6343,6 +6470,51 @@ def test_get_deployment_credentials_with_provider_includes_bucket_name(): assert credentials["custom_llm_provider"] == "vertex_ai" +def test_get_deployment_credentials_with_provider_keeps_legacy_bucket_name(): + router = litellm.Router( + model_list=[ + { + "model_name": "vertex-gemini", + "litellm_params": { + "model": "vertex_ai/gemini-3.5-flash", + "vertex_project": "my-project", + "vertex_location": "global", + "bucket_name": "my-legacy-bucket", + }, + } + ], + ) + + credentials = router.get_deployment_credentials_with_provider(model_id="vertex-gemini") + + assert credentials is not None + assert credentials["bucket_name"] == "my-legacy-bucket" + assert "gcs_bucket_name" not in credentials + + +def test_get_deployment_credentials_with_provider_keeps_both_bucket_keys(): + router = litellm.Router( + model_list=[ + { + "model_name": "vertex-gemini", + "litellm_params": { + "model": "vertex_ai/gemini-3.5-flash", + "vertex_project": "my-project", + "vertex_location": "global", + "gcs_bucket_name": "new-bucket", + "bucket_name": "legacy-bucket", + }, + } + ], + ) + + credentials = router.get_deployment_credentials_with_provider(model_id="vertex-gemini") + + assert credentials is not None + assert credentials["gcs_bucket_name"] == "new-bucket" + assert credentials["bucket_name"] == "legacy-bucket" + + def test_get_deployment_credentials_with_provider_resolves_credential_name(): """ Test that get_deployment_credentials_with_provider correctly resolves diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index b69d7ac3cd1..75234b34bc3 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -775,6 +775,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_read_input_token_cost_above_32k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_128k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_200k_tokens": {"type": "number"}, + "cache_read_input_token_cost_above_200k_tokens_batches": {"type": "number"}, "cache_read_input_token_cost_above_256k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens": {"type": "number"}, "cache_read_input_token_cost_above_272k_tokens_flex": {"type": "number"}, @@ -798,6 +799,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "input_cost_per_video_token": {"type": "number"}, "input_cost_per_token_above_32k_tokens": {"type": "number"}, "input_cost_per_token_above_200k_tokens": {"type": "number"}, + "input_cost_per_token_above_200k_tokens_batches": {"type": "number"}, "input_cost_per_token_above_256k_tokens": {"type": "number"}, "input_cost_per_token_above_272k_tokens": {"type": "number"}, "input_cost_per_token_above_512k_tokens": {"type": "number"}, @@ -899,6 +901,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "output_cost_per_token_above_32k_tokens": {"type": "number"}, "output_cost_per_token_above_128k_tokens": {"type": "number"}, "output_cost_per_token_above_200k_tokens": {"type": "number"}, + "output_cost_per_token_above_200k_tokens_batches": {"type": "number"}, "output_cost_per_token_above_256k_tokens": {"type": "number"}, "output_cost_per_token_above_272k_tokens": {"type": "number"}, "output_cost_per_token_above_512k_tokens": {"type": "number"}, diff --git a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx index f5138d55b5d..bc9889da724 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 databricksLogo from "../../public/assets/logos/databricks.svg"; interface CallbackConfig { id: string; @@ -181,6 +182,20 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [ }, description: "PointFive Logging Integration", }, + { + id: "zerobus", + displayName: "Databricks Zerobus", + logo: databricksLogo.src, + supports_key_team_logging: false, + dynamic_params: { + ZEROBUS_WORKSPACE_URL: "text", + ZEROBUS_SERVER_ENDPOINT: "text", + ZEROBUS_CLIENT_ID: "text", + ZEROBUS_CLIENT_SECRET: "password", + ZEROBUS_TABLE_NAME: "text", + }, + description: "Databricks Zerobus Ingest Logging Integration", + }, { id: "s3", displayName: "S3", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index edcc9ed1e67..d48b627c54d 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -32945,6 +32945,8 @@ export interface components { azure_username?: string | null; /** Bedrock Tags */ bedrock_tags?: unknown[] | null; + /** Bucket Name */ + bucket_name?: string | null; /** Budget Duration */ budget_duration?: string | null; /** Cache Creation Input Audio Token Cost */ @@ -32979,6 +32981,8 @@ export interface components { cache_read_input_token_cost?: number | null; /** Cache Read Input Token Cost Above 200K Tokens */ cache_read_input_token_cost_above_200k_tokens?: number | null; + /** Cache Read Input Token Cost Above 200K Tokens Batches */ + cache_read_input_token_cost_above_200k_tokens_batches?: number | null; /** Cache Read Input Token Cost Above 200K Tokens Priority */ cache_read_input_token_cost_above_200k_tokens_priority?: number | null; /** Cache Read Input Token Cost Above 272K Tokens */ @@ -33059,6 +33063,8 @@ export interface components { input_cost_per_token_above_128k_tokens?: number | null; /** Input Cost Per Token Above 200K Tokens */ input_cost_per_token_above_200k_tokens?: number | null; + /** Input Cost Per Token Above 200K Tokens Batches */ + input_cost_per_token_above_200k_tokens_batches?: number | null; /** Input Cost Per Token Above 200K Tokens Priority */ input_cost_per_token_above_200k_tokens_priority?: number | null; /** Input Cost Per Token Above 272K Tokens */ @@ -33182,6 +33188,8 @@ export interface components { output_cost_per_token_above_128k_tokens?: number | null; /** Output Cost Per Token Above 200K Tokens */ output_cost_per_token_above_200k_tokens?: number | null; + /** Output Cost Per Token Above 200K Tokens Batches */ + output_cost_per_token_above_200k_tokens_batches?: number | null; /** Output Cost Per Token Above 200K Tokens Priority */ output_cost_per_token_above_200k_tokens_priority?: number | null; /** Output Cost Per Token Above 272K Tokens */ @@ -46734,6 +46742,8 @@ export interface components { azure_username?: string | null; /** Bedrock Tags */ bedrock_tags?: unknown[] | null; + /** Bucket Name */ + bucket_name?: string | null; /** Budget Duration */ budget_duration?: string | null; /** Cache Creation Input Audio Token Cost */ @@ -46768,6 +46778,8 @@ export interface components { cache_read_input_token_cost?: number | null; /** Cache Read Input Token Cost Above 200K Tokens */ cache_read_input_token_cost_above_200k_tokens?: number | null; + /** Cache Read Input Token Cost Above 200K Tokens Batches */ + cache_read_input_token_cost_above_200k_tokens_batches?: number | null; /** Cache Read Input Token Cost Above 200K Tokens Priority */ cache_read_input_token_cost_above_200k_tokens_priority?: number | null; /** Cache Read Input Token Cost Above 272K Tokens */ @@ -46848,6 +46860,8 @@ export interface components { input_cost_per_token_above_128k_tokens?: number | null; /** Input Cost Per Token Above 200K Tokens */ input_cost_per_token_above_200k_tokens?: number | null; + /** Input Cost Per Token Above 200K Tokens Batches */ + input_cost_per_token_above_200k_tokens_batches?: number | null; /** Input Cost Per Token Above 200K Tokens Priority */ input_cost_per_token_above_200k_tokens_priority?: number | null; /** Input Cost Per Token Above 272K Tokens */ @@ -46971,6 +46985,8 @@ export interface components { output_cost_per_token_above_128k_tokens?: number | null; /** Output Cost Per Token Above 200K Tokens */ output_cost_per_token_above_200k_tokens?: number | null; + /** Output Cost Per Token Above 200K Tokens Batches */ + output_cost_per_token_above_200k_tokens_batches?: number | null; /** Output Cost Per Token Above 200K Tokens Priority */ output_cost_per_token_above_200k_tokens_priority?: number | null; /** Output Cost Per Token Above 272K Tokens */ diff --git a/uv.lock b/uv.lock index c235171ecb2..527f53bd372 100644 --- a/uv.lock +++ b/uv.lock @@ -4743,6 +4743,7 @@ proxy-dev = [ { name = "opentelemetry-sdk" }, { name = "prisma" }, { name = "prometheus-client" }, + { name = "sentry-sdk" }, ] [package.metadata] @@ -4956,6 +4957,7 @@ proxy-dev = [ { name = "opentelemetry-sdk", specifier = "==1.33.1" }, { name = "prisma", specifier = "==0.11.0" }, { name = "prometheus-client", specifier = "==0.20.0" }, + { name = "sentry-sdk", specifier = "==2.21.0" }, ] [[package]]